diff --git a/README.md b/README.md index 9a490de..a4c8e76 100644 --- a/README.md +++ b/README.md @@ -92,6 +92,7 @@ The `migrate`, `migrate:status`, `migrate:rollback`, `serve` and admin commands | [conga](modules/conga/README.md) | Background jobs on River over the shared Postgres pool: transactional dispatch, a `summer_jobs` progress record, in-process or dedicated workers, and a wall-clock scheduler. | | [festival](modules/festival/README.md) | Typed, synchronous event bus with listener priorities, payload collection and stop-when-handled dispatch. | | [fetchguard](modules/fetchguard/README.md) | Guarded outbound HTTPS fetcher that blocks private and reserved addresses and enforces host, size and timeout limits. | +| [flare](modules/flare/README.md) | Web Push delivery with VAPID (RFC 8292) and aes128gcm payload encryption (RFC 8291) behind a small Pusher interface. | | [lagoon](modules/lagoon/README.md) | Postgres data layer: the shared GORM connection, per-plugin migrations, model helpers and file attachments. | | [lighthouse](modules/lighthouse/README.md) | Transport-neutral realtime: a publisher interface with pluggable drivers, subscribe-time channel authorization, and model broadcasts enqueued in the write transaction. | | [pact](modules/pact/README.md) | Capability interfaces that compiled plugins implement to contribute routes, config, migrations, middleware, commands, admin screens, translations, mail templates and jobs. | diff --git a/modules/flare/README.md b/modules/flare/README.md new file mode 100644 index 0000000..fdeb835 --- /dev/null +++ b/modules/flare/README.md @@ -0,0 +1,127 @@ +# flare + +Web Push delivery with VAPID (RFC 8292) and aes128gcm payload encryption (RFC 8291) behind a small Pusher interface. + +`import "git.golem15.com/golem15/summercms/modules/flare"` + +## Overview + +flare sends browser push notifications. Push is a separate channel from realtime: `lighthouse` publishes to clients that hold an open connection, while flare hands a message to the browser vendor's push service, which wakes the browser even when no page is open. + +The application owns the subscriptions. When a browser subscribes, the frontend posts its `PushSubscription` (the endpoint URL and the `p256dh` and `auth` keys) to the application, which stores it. flare never reads a database. Code that sends a push passes a `flare.Subscription` to a `flare.Pusher`, and operator tooling reads stored subscriptions through a `flare.SubscriptionSource` that the application publishes on the app. + +`flare.From` builds the app-scoped `flare.Service` from `push.*` on first use. Its `flare.Service.Pusher` is the VAPID driver, `flare.VAPIDPusher`, which talks to push services directly with the standard library: + +- The payload is encrypted for the subscriber with `flare.Encrypt`: an ephemeral P-256 key agreement (`crypto/ecdh`) with the subscription's `p256dh` key, mixed with its `auth` secret through HKDF-SHA-256 (`crypto/hkdf`), then one AES-128-GCM record in the `aes128gcm` content coding. The implementation reproduces the RFC 8291 Appendix A test vector byte for byte. +- Every request carries `Authorization: vapid t=, k=` from `flare.VAPIDHeader`. The ES256 token's `aud` is the endpoint's origin, `exp` lies `flare.VAPIDTokenLifetime` (12 hours) ahead and `sub` is `push.subject`. + +Endpoints come from browsers, so they are untrusted URLs. The driver only sends to `https` endpoints whose host is in `push.allowed_hosts`, checks this before it opens a connection, and never follows a redirect. + +## Features + +- `flare.Pusher` with one method, `Send(ctx, sub, payload, opts)`. `flare.SendOptions` sets the `TTL` (default `push.ttl`), `Urgency` and `Topic` headers. +- The VAPID driver POSTs the encrypted body with `TTL`, `Content-Encoding: aes128gcm`, `Content-Type: application/octet-stream`, the optional `Urgency` and `Topic` and the VAPID `Authorization` header. A 2xx answer is success. 404 and 410 return `flare.ErrSubscriptionGone`, so the caller can delete the subscription. Any other status returns a `*flare.StatusError` with the code and without the response body. Requests time out after `flare.DefaultTimeout` (10 s). +- Nothing is sent while `push.enabled` is false: `Send` returns `flare.ErrPushDisabled`. +- Endpoint allowlist: `flare.HostAllowed` matches a host against `push.allowed_hosts`, where `*.example.com` matches any subdomain of `example.com` (not `example.com` itself). The defaults, `flare.DefaultAllowedHosts`, are Firebase Cloud Messaging, Mozilla autopush, Apple and Windows push. A refused endpoint returns `flare.ErrEndpointNotAllowed`, and the error names the host, never the endpoint path. +- Payloads up to `flare.MaxPayloadSize` (3993 bytes, the RFC 8291 limit for a 4096-byte body). A larger one returns `flare.ErrPayloadTooLarge`. +- VAPID keys: `flare.GenerateVAPIDKeys` returns a P-256 pair as unpadded base64url (`flare.PublicKeyLength`, 87 characters, and `flare.PrivateKeyLength`, 43 characters). `flare.ParseVAPIDKeys` accepts padded or unpadded input and checks that the public key belongs to the private key; a bad pair returns `flare.ErrInvalidVAPIDKeys`. A subject that is not `mailto:` or `https:` returns `flare.ErrInvalidSubject`. +- The private key never reaches logs or formatted output: `flare.VAPIDKeys` and `flare.Config` redact it in `String`, `GoString` and (for the config) `LogValue`, and no error carries key material. + +## Usage + +Send a push to one stored subscription: + +```go +svc, err := flare.From(app) +if err != nil { + return err +} +payload := []byte(`{"title":"New comment","body":"Someone replied to your post"}`) +err = svc.Pusher().Send(ctx, flare.Subscription{ + Endpoint: row.Endpoint, + P256dh: row.P256dh, + Auth: row.Auth, +}, payload, flare.SendOptions{Urgency: "normal"}) +if errors.Is(err, flare.ErrSubscriptionGone) { + // The browser unsubscribed: delete the row. +} +``` + +Publish a subscription source from a plugin's `Boot`, so operator commands can read the stored subscriptions of a user: + +```go +type blogSubscriptions struct{ db *gorm.DB } + +func (s blogSubscriptions) Subscriptions(ctx context.Context, userID uint) ([]flare.SubscriptionInfo, error) { + var user models.User + if err := s.db.WithContext(ctx).First(&user, userID).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, flare.ErrUserNotFound + } + return nil, err + } + var rows []models.PushSubscription + if err := s.db.WithContext(ctx).Where("user_id = ?", userID).Find(&rows).Error; err != nil { + return nil, err + } + out := make([]flare.SubscriptionInfo, 0, len(rows)) + for _, r := range rows { + out = append(out, flare.SubscriptionInfo{ + Subscription: flare.Subscription{Endpoint: r.Endpoint, P256dh: r.P256dh, Auth: r.Auth}, + ID: r.ID, + UserAgent: r.UserAgent, + SubscribedAt: &r.CreatedAt, + }) + } + return out, nil +} + +// In Boot: +if err := app.Publish[flare.SubscriptionSource](blogSubscriptions{db: gdb}); err != nil { + return err +} +``` + +## API reference + +| Identifier | Description | +|------------|-------------| +| `flare.From(app)` | The app's `*flare.Service`, built from `push.*` and published on first use. | +| `flare.Service` | `Config`, `Enabled`, `Pusher`, `SetHTTPClient` (replace the driver's HTTP client, for example in tests) and `Logger`. | +| `flare.Pusher` | `Send(ctx, sub, payload, opts) error`. | +| `flare.VAPIDPusher`, `flare.NewVAPIDPusher(cfg, hc)` | The VAPID driver. `hc` may be nil; a given client is copied and never follows redirects. | +| `flare.Subscription` | `Endpoint`, `P256dh`, `Auth`, as `PushSubscription.toJSON` returns them. | +| `flare.SendOptions` | `TTL`, `Urgency`, `Topic`. | +| `flare.SubscriptionSource`, `flare.SubscriptionInfo` | The application's subscription store: `Subscriptions(ctx, userID)` returns the subscriptions with `ID`, `UserAgent`, `SubscribedAt` and `LastUsedAt`. | +| `flare.Config`, `flare.LoadConfig` | The `push.*` settings with their defaults; `Keys` returns the key pair. | +| `flare.VAPIDKeys`, `flare.GenerateVAPIDKeys`, `flare.ParseVAPIDKeys` | VAPID key pairs as unpadded base64url. | +| `flare.VAPIDHeader(endpoint, subject, keys, now)` | The RFC 8292 `Authorization` header value. | +| `flare.Encrypt(payload, sub)` | The RFC 8291 `aes128gcm` request body. | +| `flare.HostAllowed(host, allowed)`, `flare.DefaultAllowedHosts` | The endpoint host allowlist and its default. | +| `flare.ErrPushDisabled`, `flare.ErrEndpointNotAllowed`, `flare.ErrSubscriptionGone`, `flare.ErrUserNotFound`, `flare.ErrPayloadTooLarge`, `flare.ErrInvalidVAPIDKeys`, `flare.ErrInvalidSubject`, `flare.StatusError` | Errors. | +| `flare.ContentEncoding`, `flare.MaxPayloadSize`, `flare.DefaultTTL`, `flare.DefaultTimeout`, `flare.VAPIDTokenLifetime`, `flare.PublicKeyLength`, `flare.PrivateKeyLength` | Constants. | + +## Configuration + +| Key | Default | Description | +|-----|---------|-------------| +| `push.enabled` | `false` | Nothing is sent while false. | +| `push.public_key` | `""` | VAPID public key, base64url (87 characters unpadded). | +| `push.private_key` | `""` | VAPID private key, base64url (43 characters). Keep it out of committed files; set `SUMMER_PUSH__PRIVATE_KEY`. | +| `push.subject` | `""` | VAPID `sub` claim: a `mailto:` or `https:` contact URI. | +| `push.ttl` | `2419200` | Default `TTL` header, in seconds or as a duration string. | +| `push.allowed_hosts` | FCM, Mozilla autopush, `*.push.apple.com`, `*.notify.windows.com` | Push service hosts an endpoint may point at, as a list or a comma-separated string. | + +## Dependencies + +- `backpack` and `compass` from this repository. +- `github.com/golang-jwt/jwt/v5` (the ES256 VAPID token). +- Everything else is the standard library: `crypto/ecdh`, `crypto/ecdsa`, `crypto/hkdf`, `crypto/aes`, `crypto/cipher` and `net/http`. No Web Push library is used. + +## Testing + +```bash +go test ./modules/flare/... +``` + +`TestRFC8291AppendixA` fixes the RFC's salt and application server key and compares the output with the RFC's published bytes. The send tests run an `httptest.NewTLSServer` push service with `127.0.0.1` in the allowlist; it verifies the VAPID token with the key from `k=` and decrypts the body as a browser would. Pass the test server's client to `flare.NewVAPIDPusher` or `flare.Service.SetHTTPClient` so it trusts the test certificate. diff --git a/modules/flare/encrypt.go b/modules/flare/encrypt.go new file mode 100644 index 0000000..b16a981 --- /dev/null +++ b/modules/flare/encrypt.go @@ -0,0 +1,163 @@ +package flare + +import ( + "crypto/aes" + "crypto/cipher" + "crypto/ecdh" + "crypto/hkdf" + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "encoding/binary" + "errors" + "fmt" + "strings" +) + +// ContentEncoding is the only content coding a push message may use +// (RFC 8291 Section 4). +const ContentEncoding = "aes128gcm" + +// MaxPayloadSize is the largest plaintext Encrypt accepts: a push service +// need not accept more than 4096 octets of body, and the aes128gcm header +// (86 octets), the padding delimiter (1) and the GCM tag (16) take the rest. +const MaxPayloadSize = 3993 + +// recordSize is the "rs" field of the aes128gcm header. A push message is a +// single record, so any value above the record length works; 4096 is the +// value RFC 8291 uses. +const recordSize = 4096 + +const ( + saltSize = 16 + authSize = 16 + publicKeySize = 65 + headerSize = saltSize + 4 + 1 + publicKeySize +) + +// Key derivation labels from RFC 8291 Section 3.4 (the RFC 8188 labels end +// with a zero octet). +const ( + webPushInfo = "WebPush: info\x00" + cekInfo = "Content-Encoding: aes128gcm\x00" + nonceInfo = "Content-Encoding: nonce\x00" +) + +// ErrPayloadTooLarge is returned by Encrypt for a payload above +// MaxPayloadSize. +var ErrPayloadTooLarge = errors.New("flare: payload exceeds 3993 bytes") + +// Encrypt encrypts payload for the user agent that owns sub, following +// RFC 8291: an ephemeral P-256 key agreement with sub.P256dh, mixed with the +// sub.Auth secret, then one aes128gcm record (RFC 8188). The result is the +// request body: salt || record size || key id length || the ephemeral +// public key || ciphertext. +func Encrypt(payload []byte, sub Subscription) ([]byte, error) { + asKey, err := ecdh.P256().GenerateKey(rand.Reader) + if err != nil { + return nil, fmt.Errorf("flare: generate ephemeral key: %w", err) + } + salt := make([]byte, saltSize) + if _, err := rand.Read(salt); err != nil { + return nil, fmt.Errorf("flare: generate salt: %w", err) + } + return encrypt(payload, sub, asKey, salt) +} + +// encrypt is Encrypt with the application server key and the salt given, +// so the RFC 8291 Appendix A vector can fix both. +func encrypt(payload []byte, sub Subscription, asKey *ecdh.PrivateKey, salt []byte) ([]byte, error) { + if len(payload) > MaxPayloadSize { + return nil, ErrPayloadTooLarge + } + if len(salt) != saltSize { + return nil, fmt.Errorf("flare: salt must be %d bytes", saltSize) + } + uaRaw, err := decodeBase64URL(sub.P256dh) + if err != nil { + return nil, fmt.Errorf("flare: subscription p256dh is not base64url") + } + uaPublic, err := ecdh.P256().NewPublicKey(uaRaw) + if err != nil { + return nil, fmt.Errorf("flare: subscription p256dh is not an uncompressed P-256 point") + } + authSecret, err := decodeBase64URL(sub.Auth) + if err != nil { + return nil, fmt.Errorf("flare: subscription auth is not base64url") + } + if len(authSecret) != authSize { + return nil, fmt.Errorf("flare: subscription auth must be %d bytes, got %d", authSize, len(authSecret)) + } + + ecdhSecret, err := asKey.ECDH(uaPublic) + if err != nil { + return nil, fmt.Errorf("flare: key agreement: %w", err) + } + asPublic := asKey.PublicKey().Bytes() + cek, nonce, err := deriveKeys(ecdhSecret, authSecret, salt, uaPublic.Bytes(), asPublic) + if err != nil { + return nil, err + } + + block, err := aes.NewCipher(cek) + if err != nil { + return nil, fmt.Errorf("flare: %w", err) + } + gcm, err := cipher.NewGCM(block) + if err != nil { + return nil, fmt.Errorf("flare: %w", err) + } + // A single, final record: the plaintext followed by the 0x02 padding + // delimiter and no padding. + plain := make([]byte, 0, len(payload)+1) + plain = append(plain, payload...) + plain = append(plain, 0x02) + + out := make([]byte, headerSize, headerSize+len(plain)+gcm.Overhead()) + copy(out, salt) + binary.BigEndian.PutUint32(out[saltSize:], recordSize) + out[saltSize+4] = byte(len(asPublic)) + copy(out[saltSize+5:], asPublic) + return gcm.Seal(out, nonce, plain, nil), nil +} + +// deriveKeys runs the RFC 8291 Section 3.4 key schedule and returns the +// content encryption key and the nonce. +func deriveKeys(ecdhSecret, authSecret, salt, uaPublic, asPublic []byte) (cek, nonce []byte, err error) { + prkKey, err := hkdf.Extract(sha256.New, ecdhSecret, authSecret) + if err != nil { + return nil, nil, fmt.Errorf("flare: %w", err) + } + keyInfo := make([]byte, 0, len(webPushInfo)+len(uaPublic)+len(asPublic)) + keyInfo = append(keyInfo, webPushInfo...) + keyInfo = append(keyInfo, uaPublic...) + keyInfo = append(keyInfo, asPublic...) + ikm, err := hkdf.Expand(sha256.New, prkKey, string(keyInfo), 32) + if err != nil { + return nil, nil, fmt.Errorf("flare: %w", err) + } + prk, err := hkdf.Extract(sha256.New, ikm, salt) + if err != nil { + return nil, nil, fmt.Errorf("flare: %w", err) + } + cek, err = hkdf.Expand(sha256.New, prk, cekInfo, 16) + if err != nil { + return nil, nil, fmt.Errorf("flare: %w", err) + } + nonce, err = hkdf.Expand(sha256.New, prk, nonceInfo, 12) + if err != nil { + return nil, nil, fmt.Errorf("flare: %w", err) + } + return cek, nonce, nil +} + +// decodeBase64URL decodes base64url with or without padding. Browsers hand +// out unpadded keys; stored copies sometimes carry padding or use the +// standard alphabet, which is accepted too. +func decodeBase64URL(s string) ([]byte, error) { + s = strings.TrimRight(strings.TrimSpace(s), "=") + if b, err := base64.RawURLEncoding.DecodeString(s); err == nil { + return b, nil + } + return base64.RawStdEncoding.DecodeString(s) +} diff --git a/modules/flare/encrypt_test.go b/modules/flare/encrypt_test.go new file mode 100644 index 0000000..d386966 --- /dev/null +++ b/modules/flare/encrypt_test.go @@ -0,0 +1,112 @@ +package flare + +import ( + "bytes" + "crypto/ecdh" + "encoding/base64" + "strings" + "testing" +) + +// Values copied verbatim from RFC 8291 Section 5 and Appendix A +// (https://www.rfc-editor.org/rfc/rfc8291), with the line-wrapping +// whitespace removed. +const ( + rfcPlaintext = "V2hlbiBJIGdyb3cgdXAsIEkgd2FudCB0byBiZSBhIHdhdGVybWVsb24" + rfcASPublic = "BP4z9KsN6nGRTbVYI_c7VJSPQTBtkgcy27mlmlMoZIIgDll6e3vCYLocInmYWAmS6TlzAC8wEqKK6PBru3jl7A8" + rfcASPrivate = "yfWPiYE-n46HLnH0KqZOF1fJJU3MYrct3AELtAQ-oRw" + rfcUAPublic = "BCVxsr7N_eNgVRqvHtD0zTZsEc6-VV-JvLexhqUzORcxaOzi6-AYWXvTBHm4bjyPjs7Vd8pZGH6SRpkNtoIAiw4" + rfcUAPrivate = "q1dXpw3UpT5VOmu_cf_v6ih07Aems3njxI-JWgLcM94" + rfcSalt = "DGv6ra1nlYgDCS1FRnbzlw" + rfcAuthSecret = "BTBZMqHH6r4Tts7J_aSIgg" + rfcCEK = "oIhVW04MRdy2XN9CiKLxTg" + rfcNonce = "4h_95klXJ5E_qnoN" + rfcHeader = "DGv6ra1nlYgDCS1FRnbzlwAAEABBBP4z9KsN6nGRTbVYI_c7VJSPQTBtkgcy27mlmlMoZIIgDll6e3vCYLocInmYWAmS6TlzAC8wEqKK6PBru3jl7A8" + rfcCiphertext = "8pfeW0KbunFT06SuDKoJH9Ql87S1QUrdirN6GcG7sFz1y1sqLgVi1VhjVkHsUoEsbI_0LpXMuGvnzQ" + // The request body of Section 5: header || ciphertext. + rfcBody = "DGv6ra1nlYgDCS1FRnbzlwAAEABBBP4z9KsN6nGRTbVYI_c7VJSPQTBtkgcy27ml" + + "mlMoZIIgDll6e3vCYLocInmYWAmS6TlzAC8wEqKK6PBru3jl7A_yl95bQpu6cVPT" + + "pK4Mqgkf1CXztLVBSt2Ks3oZwbuwXPXLWyouBWLVWGNWQexSgSxsj_Qulcy4a-fN" +) + +func b64(t *testing.T, s string) []byte { + t.Helper() + b, err := base64.RawURLEncoding.DecodeString(s) + if err != nil { + t.Fatalf("decode %q: %v", s, err) + } + return b +} + +func TestRFC8291AppendixA(t *testing.T) { + plaintext := b64(t, rfcPlaintext) + if string(plaintext) != "When I grow up, I want to be a watermelon" { + t.Fatalf("plaintext = %q", plaintext) + } + asKey, err := ecdh.P256().NewPrivateKey(b64(t, rfcASPrivate)) + if err != nil { + t.Fatalf("as_private: %v", err) + } + if !bytes.Equal(asKey.PublicKey().Bytes(), b64(t, rfcASPublic)) { + t.Fatal("as_private does not produce the RFC as_public") + } + sub := Subscription{Endpoint: "https://push.example.net/push/JzLQ3raZJfFBR0aqvOMsLrt54w4rJUsV", P256dh: rfcUAPublic, Auth: rfcAuthSecret} + + got, err := encrypt(plaintext, sub, asKey, b64(t, rfcSalt)) + if err != nil { + t.Fatalf("encrypt: %v", err) + } + want := b64(t, rfcBody) + if !bytes.Equal(got, want) { + t.Fatalf("body mismatch\n got %s\nwant %s", base64.RawURLEncoding.EncodeToString(got), rfcBody) + } + header := b64(t, rfcHeader) + if len(header) != 86 || !bytes.Equal(got[:86], header) { + t.Fatalf("header mismatch: %s", base64.RawURLEncoding.EncodeToString(got[:86])) + } + if !bytes.Equal(got[86:], b64(t, rfcCiphertext)) { + t.Fatalf("ciphertext mismatch: %s", base64.RawURLEncoding.EncodeToString(got[86:])) + } + + // The intermediate CEK and nonce, derived from the receiver side. + uaKey, err := ecdh.P256().NewPrivateKey(b64(t, rfcUAPrivate)) + if err != nil { + t.Fatalf("ua_private: %v", err) + } + secret, err := uaKey.ECDH(asKey.PublicKey()) + if err != nil { + t.Fatal(err) + } + cek, nonce, err := deriveKeys(secret, b64(t, rfcAuthSecret), b64(t, rfcSalt), b64(t, rfcUAPublic), b64(t, rfcASPublic)) + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(cek, b64(t, rfcCEK)) || !bytes.Equal(nonce, b64(t, rfcNonce)) { + t.Fatalf("CEK/NONCE = %x/%x", cek, nonce) + } + + // The receiver decrypts the RFC body back to the plaintext. + back, err := decryptForTest(want, rfcUAPrivate, rfcAuthSecret) + if err != nil { + t.Fatalf("decrypt: %v", err) + } + if !bytes.Equal(back, plaintext) { + t.Fatalf("decrypted %q", back) + } + + // Encrypt with fresh randomness also round-trips, and refuses a + // payload above 3993 bytes. + body, err := Encrypt(plaintext, sub) + if err != nil { + t.Fatal(err) + } + if back, err := decryptForTest(body, rfcUAPrivate, rfcAuthSecret); err != nil || !bytes.Equal(back, plaintext) { + t.Fatalf("round trip: %q, %v", back, err) + } + if _, err := Encrypt([]byte(strings.Repeat("x", MaxPayloadSize+1)), sub); err != ErrPayloadTooLarge { + t.Fatalf("oversized payload: %v", err) + } + if _, err := Encrypt([]byte(strings.Repeat("x", MaxPayloadSize)), sub); err != nil { + t.Fatalf("3993-byte payload: %v", err) + } +} diff --git a/modules/flare/flare.go b/modules/flare/flare.go new file mode 100644 index 0000000..2590e4f --- /dev/null +++ b/modules/flare/flare.go @@ -0,0 +1,423 @@ +// Package flare delivers Web Push notifications: a small Pusher interface +// and a VAPID driver that encrypts every payload with aes128gcm (RFC 8291) +// and signs every request with an ES256 VAPID token (RFC 8292). Push is a +// separate channel from realtime; the application owns the subscription +// store and hands subscriptions to flare through a SubscriptionSource. +package flare + +import ( + "bytes" + "context" + "errors" + "fmt" + "io" + "log/slog" + "net/http" + "net/url" + "strconv" + "strings" + "sync" + "time" + + "git.golem15.com/golem15/summercms/modules/backpack" + "git.golem15.com/golem15/summercms/modules/compass" +) + +// DefaultTTL is how long a push service keeps an undelivered message when +// push.ttl is not set: four weeks. +const DefaultTTL = 2419200 * time.Second + +// DefaultTimeout bounds every request to a push service. +const DefaultTimeout = 10 * time.Second + +// Errors returned by Pusher.Send and SubscriptionSource. +var ( + // ErrPushDisabled is returned while push.enabled is false. + ErrPushDisabled = errors.New("flare: push is disabled (push.enabled is false)") + // ErrEndpointNotAllowed is returned, before any connection is made, + // for an endpoint that is not https or whose host is not in + // push.allowed_hosts. + ErrEndpointNotAllowed = errors.New("flare: push endpoint not allowed") + // ErrSubscriptionGone is returned when the push service answers 404 or + // 410: the subscription expired and should be deleted. + ErrSubscriptionGone = errors.New("flare: push subscription is gone") + // ErrUserNotFound is returned by a SubscriptionSource for an unknown + // user. + ErrUserNotFound = errors.New("flare: user not found") +) + +// StatusError is a push service answer other than 2xx, 404 or 410. It +// carries the status code, never the response body. +type StatusError struct { + Code int +} + +func (e *StatusError) Error() string { + return "flare: push service answered HTTP " + strconv.Itoa(e.Code) +} + +// Subscription is a browser push subscription as PushSubscription.toJSON +// returns it: the endpoint URL and the user agent's P-256 public key and +// 16-byte authentication secret, both base64url. +type Subscription struct { + Endpoint string `json:"endpoint"` + P256dh string `json:"p256dh"` + Auth string `json:"auth"` +} + +// SendOptions tune one push. A zero TTL uses push.ttl; Urgency ("very-low", +// "low", "normal" or "high") and Topic are sent only when set. +type SendOptions struct { + TTL time.Duration + Urgency string + Topic string +} + +// Pusher sends one encrypted push message to one subscription. +type Pusher interface { + Send(ctx context.Context, sub Subscription, payload []byte, opts SendOptions) error +} + +// SubscriptionInfo is a stored subscription with the details the +// websockets:test-push command reports. +type SubscriptionInfo struct { + Subscription + ID uint + UserAgent string + SubscribedAt *time.Time + LastUsedAt *time.Time +} + +// SubscriptionSource is implemented by the application that stores push +// subscriptions, and published on the app with backpack.App.Publish as a +// flare.SubscriptionSource. Subscriptions returns ErrUserNotFound for an +// unknown user and an empty list for a user without subscriptions. +type SubscriptionSource interface { + Subscriptions(ctx context.Context, userID uint) ([]SubscriptionInfo, error) +} + +// Config is the push.* configuration. String and GoString never print +// PrivateKey. +type Config struct { + // Enabled is push.enabled; nothing is sent while it is false. + Enabled bool + // PublicKey and PrivateKey are the VAPID key pair, base64url. + PublicKey string + PrivateKey string + // Subject is the VAPID sub claim, a mailto: or https: URI. + Subject string + // TTL is the default TTL header. + TTL time.Duration + // AllowedHosts are the push service hosts an endpoint may point at. + // "*.example.com" matches any subdomain of example.com. + AllowedHosts []string +} + +// String prints the configuration without the private key. +func (c Config) String() string { + priv := "unset" + if c.PrivateKey != "" { + priv = "[redacted]" + } + return fmt.Sprintf("flare.Config{Enabled: %t, PublicKey: %s, PrivateKey: %s, Subject: %s, TTL: %s, AllowedHosts: %v}", + c.Enabled, c.PublicKey, priv, c.Subject, c.TTL, c.AllowedHosts) +} + +// GoString is String. +func (c Config) GoString() string { return c.String() } + +// LogValue keeps the private key out of structured logs. +func (c Config) LogValue() slog.Value { return slog.StringValue(c.String()) } + +// Keys returns the configured VAPID key pair. +func (c Config) Keys() VAPIDKeys { + return VAPIDKeys{PublicKey: c.PublicKey, PrivateKey: c.PrivateKey} +} + +// DefaultAllowedHosts are the push services endpoints may point at when +// push.allowed_hosts is not set: Firebase Cloud Messaging (Chrome, Edge on +// Android), Mozilla autopush (Firefox), Apple and Windows push. +func DefaultAllowedHosts() []string { + return []string{ + "fcm.googleapis.com", + "updates.push.services.mozilla.com", + "*.push.apple.com", + "*.notify.windows.com", + } +} + +// LoadConfig reads push.* from c, filling the defaults. push.ttl is an +// integer number of seconds or a duration string; push.allowed_hosts is a +// list or a comma-separated string. +func LoadConfig(c *compass.Config) Config { + cfg := Config{TTL: DefaultTTL, AllowedHosts: DefaultAllowedHosts()} + if c == nil { + return cfg + } + cfg.Enabled = c.Bool("push.enabled") + cfg.PublicKey = strings.TrimSpace(c.String("push.public_key")) + cfg.PrivateKey = strings.TrimSpace(c.String("push.private_key")) + cfg.Subject = strings.TrimSpace(c.String("push.subject")) + if d := durationSetting(c, "push.ttl"); d > 0 { + cfg.TTL = d + } + if v, ok := c.Lookup("push.allowed_hosts"); ok { + if hosts := stringList(v); len(hosts) > 0 { + cfg.AllowedHosts = hosts + } + } + return cfg +} + +func durationSetting(c *compass.Config, key string) time.Duration { + raw := strings.TrimSpace(c.String(key)) + if raw == "" { + return 0 + } + if n, err := strconv.ParseInt(raw, 10, 64); err == nil { + return time.Duration(n) * time.Second + } + if d, err := time.ParseDuration(raw); err == nil { + return d + } + return 0 +} + +func stringList(v any) []string { + var parts []string + switch t := v.(type) { + case string: + parts = strings.Split(t, ",") + case []string: + parts = t + case []any: + for _, e := range t { + if s, ok := e.(string); ok { + parts = append(parts, s) + } + } + } + out := make([]string, 0, len(parts)) + for _, p := range parts { + if p = strings.ToLower(strings.TrimSpace(p)); p != "" { + out = append(out, p) + } + } + return out +} + +// Service is the app-scoped push service. Get it with From. +type Service struct { + cfg Config + log *slog.Logger + + mu sync.RWMutex + pusher *VAPIDPusher +} + +// From returns the app's Service, building it from push.* and publishing it +// on first use. +func From(app *backpack.App) (*Service, error) { + if app == nil { + return nil, fmt.Errorf("flare: app is nil") + } + if s, ok := app.Lookup[*Service](); ok && s != nil { + return s, nil + } + cfg := LoadConfig(app.Config) + svc := &Service{cfg: cfg, log: loggerFromApp(app), pusher: NewVAPIDPusher(cfg, nil)} + if err := app.Publish(svc); err != nil { + if existing, ok := app.Lookup[*Service](); ok && existing != nil { + return existing, nil + } + return nil, fmt.Errorf("flare: %w", err) + } + return svc, nil +} + +// Config returns the push configuration the service was built with. +func (s *Service) Config() Config { + if s == nil { + return LoadConfig(nil) + } + return s.cfg +} + +// Enabled reports push.enabled. +func (s *Service) Enabled() bool { return s != nil && s.cfg.Enabled } + +// Pusher returns the VAPID driver. +func (s *Service) Pusher() Pusher { + if s == nil { + return NewVAPIDPusher(LoadConfig(nil), nil) + } + s.mu.RLock() + defer s.mu.RUnlock() + return s.pusher +} + +// SetHTTPClient replaces the HTTP client of the VAPID driver, for example +// with an httptest TLS client in tests. Redirects stay refused. +func (s *Service) SetHTTPClient(hc *http.Client) { + if s == nil { + return + } + s.mu.Lock() + defer s.mu.Unlock() + s.pusher = NewVAPIDPusher(s.cfg, hc) +} + +// Logger returns the app logger. +func (s *Service) Logger() *slog.Logger { + if s == nil || s.log == nil { + return slog.Default() + } + return s.log +} + +// VAPIDPusher is the Pusher that talks to push services directly with +// VAPID authentication. It is safe for concurrent use. +type VAPIDPusher struct { + cfg Config + hc *http.Client + now func() time.Time +} + +var _ Pusher = (*VAPIDPusher)(nil) + +// NewVAPIDPusher returns a driver for cfg. hc may be nil for a client with +// a DefaultTimeout timeout; a given client is copied. Either way the driver +// never follows redirects, so a push service cannot send it to a host +// outside push.allowed_hosts. +func NewVAPIDPusher(cfg Config, hc *http.Client) *VAPIDPusher { + var client http.Client + if hc != nil { + client = *hc + } + if client.Timeout == 0 { + client.Timeout = DefaultTimeout + } + client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } + return &VAPIDPusher{cfg: cfg, hc: &client, now: time.Now} +} + +// Send encrypts payload for sub and POSTs it to sub.Endpoint. It returns +// ErrPushDisabled while push is disabled and ErrEndpointNotAllowed, before +// dialing, for an endpoint outside the allowlist. A 2xx answer is success, +// 404 and 410 are ErrSubscriptionGone, any other status a *StatusError. +func (p *VAPIDPusher) Send(ctx context.Context, sub Subscription, payload []byte, opts SendOptions) error { + if p == nil || !p.cfg.Enabled { + return ErrPushDisabled + } + if err := p.checkEndpoint(sub.Endpoint); err != nil { + return err + } + body, err := Encrypt(payload, sub) + if err != nil { + return err + } + auth, err := VAPIDHeader(sub.Endpoint, p.cfg.Subject, p.cfg.Keys(), p.now()) + if err != nil { + return err + } + if ctx == nil { + ctx = context.Background() + } + req, err := http.NewRequestWithContext(ctx, http.MethodPost, sub.Endpoint, bytes.NewReader(body)) + if err != nil { + return fmt.Errorf("flare: build request: %w", err) + } + ttl := opts.TTL + if ttl <= 0 { + ttl = p.cfg.TTL + } + req.Header.Set("TTL", strconv.FormatInt(int64(ttl/time.Second), 10)) + req.Header.Set("Content-Encoding", ContentEncoding) + req.Header.Set("Content-Type", "application/octet-stream") + req.Header.Set("Authorization", auth) + if opts.Urgency != "" { + req.Header.Set("Urgency", opts.Urgency) + } + if opts.Topic != "" { + req.Header.Set("Topic", opts.Topic) + } + resp, err := p.hc.Do(req) + if err != nil { + // *url.Error repeats the full endpoint; keep only the cause. + var uerr *url.Error + if errors.As(err, &uerr) && uerr.Err != nil { + err = uerr.Err + } + return fmt.Errorf("flare: push to %s: %w", hostOf(sub.Endpoint), err) + } + defer resp.Body.Close() + _, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 64<<10)) + switch { + case resp.StatusCode >= 200 && resp.StatusCode <= 299: + return nil + case resp.StatusCode == http.StatusNotFound || resp.StatusCode == http.StatusGone: + return ErrSubscriptionGone + default: + return &StatusError{Code: resp.StatusCode} + } +} + +// checkEndpoint allows only https URLs, without user info, whose host is in +// push.allowed_hosts. The error names the host, never the full endpoint, +// whose path is a capability. +func (p *VAPIDPusher) checkEndpoint(endpoint string) error { + u, err := url.Parse(endpoint) + if err != nil || u.Host == "" { + return fmt.Errorf("%w: not an absolute URL", ErrEndpointNotAllowed) + } + if !strings.EqualFold(u.Scheme, "https") { + return fmt.Errorf("%w: %s is not https", ErrEndpointNotAllowed, strings.ToLower(u.Scheme)) + } + if u.User != nil { + return fmt.Errorf("%w: user info in URL", ErrEndpointNotAllowed) + } + host := strings.ToLower(u.Hostname()) + if !HostAllowed(host, p.cfg.AllowedHosts) { + return fmt.Errorf("%w: host %s is not in push.allowed_hosts", ErrEndpointNotAllowed, host) + } + return nil +} + +// HostAllowed reports whether host matches an entry of allowed. An entry +// "*.example.com" matches any subdomain of example.com but not example.com +// itself; any other entry must equal host. Matching is case-insensitive. +func HostAllowed(host string, allowed []string) bool { + host = strings.TrimSuffix(strings.ToLower(host), ".") + if host == "" { + return false + } + for _, entry := range allowed { + entry = strings.ToLower(strings.TrimSpace(entry)) + if suffix, ok := strings.CutPrefix(entry, "*."); ok { + if suffix != "" && strings.HasSuffix(host, "."+suffix) { + return true + } + continue + } + if entry != "" && host == entry { + return true + } + } + return false +} + +func hostOf(endpoint string) string { + if u, err := url.Parse(endpoint); err == nil { + return u.Host + } + return "endpoint" +} + +func loggerFromApp(app *backpack.App) *slog.Logger { + if app != nil { + if log, ok := app.Lookup[*slog.Logger](); ok && log != nil { + return log + } + } + return slog.Default() +} diff --git a/modules/flare/send_test.go b/modules/flare/send_test.go new file mode 100644 index 0000000..e065bfd --- /dev/null +++ b/modules/flare/send_test.go @@ -0,0 +1,371 @@ +package flare + +import ( + "context" + "crypto/aes" + "crypto/cipher" + "crypto/ecdh" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "encoding/base64" + "encoding/binary" + "errors" + "fmt" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync/atomic" + "testing" + "time" + + "github.com/golang-jwt/jwt/v5" +) + +// decryptForTest is the receiving user agent's side of RFC 8291: it parses +// the aes128gcm header, agrees on the key with the sender's key id, derives +// CEK and nonce, opens the record and checks the 0x02 padding delimiter. +func decryptForTest(body []byte, uaPrivate, authSecret string) ([]byte, error) { + if len(body) < 21 { + return nil, errors.New("body shorter than the header") + } + salt := body[:16] + if rs := binary.BigEndian.Uint32(body[16:20]); rs < 18 { + return nil, fmt.Errorf("record size %d", rs) + } + idlen := int(body[20]) + if len(body) < 21+idlen { + return nil, errors.New("truncated key id") + } + keyID := body[21 : 21+idlen] + ciphertext := body[21+idlen:] + privRaw, err := base64.RawURLEncoding.DecodeString(uaPrivate) + if err != nil { + return nil, err + } + uaKey, err := ecdh.P256().NewPrivateKey(privRaw) + if err != nil { + return nil, err + } + asPublic, err := ecdh.P256().NewPublicKey(keyID) + if err != nil { + return nil, fmt.Errorf("key id is not a P-256 point: %w", err) + } + secret, err := uaKey.ECDH(asPublic) + if err != nil { + return nil, err + } + auth, err := base64.RawURLEncoding.DecodeString(authSecret) + if err != nil { + return nil, err + } + cek, nonce, err := deriveKeys(secret, auth, salt, uaKey.PublicKey().Bytes(), keyID) + if err != nil { + return nil, err + } + block, err := aes.NewCipher(cek) + if err != nil { + return nil, err + } + gcm, err := cipher.NewGCM(block) + if err != nil { + return nil, err + } + plain, err := gcm.Open(nil, nonce, ciphertext, nil) + if err != nil { + return nil, err + } + i := len(plain) - 1 + for i >= 0 && plain[i] == 0 { + i-- + } + if i < 0 || plain[i] != 0x02 { + return nil, errors.New("missing 0x02 padding delimiter") + } + return plain[:i], nil +} + +// testSubscriber is a browser-side subscription with its private key. +type testSubscriber struct { + sub Subscription + private string +} + +func newTestSubscriber(t *testing.T, endpoint string) testSubscriber { + t.Helper() + key, err := ecdh.P256().GenerateKey(rand.Reader) + if err != nil { + t.Fatal(err) + } + auth := make([]byte, 16) + if _, err := rand.Read(auth); err != nil { + t.Fatal(err) + } + return testSubscriber{ + sub: Subscription{ + Endpoint: endpoint, + P256dh: base64.RawURLEncoding.EncodeToString(key.PublicKey().Bytes()), + Auth: base64.RawURLEncoding.EncodeToString(auth), + }, + private: base64.RawURLEncoding.EncodeToString(key.Bytes()), + } +} + +// receivedPush is what the fake push service saw and verified. +type receivedPush struct { + header http.Header + claims jwt.MapClaims + jwtHeader map[string]any + k string + payload []byte + err error +} + +// verifyVAPIDRequest checks the Authorization header the way a push service +// does: the JWT must verify with the ES256 key from k=. +func verifyVAPIDRequest(r *http.Request, sub testSubscriber) receivedPush { + got := receivedPush{header: r.Header.Clone()} + auth := r.Header.Get("Authorization") + rest, ok := strings.CutPrefix(auth, "vapid t=") + if !ok { + got.err = fmt.Errorf("authorization %q is not vapid t=", auth) + return got + } + token, k, ok := strings.Cut(rest, ", k=") + if !ok { + got.err = errors.New("authorization lacks k=") + return got + } + got.k = k + rawKey, err := base64.RawURLEncoding.DecodeString(k) + if err != nil { + got.err = fmt.Errorf("k is not unpadded base64url: %w", err) + return got + } + pub, err := ecdsa.ParseUncompressedPublicKey(elliptic.P256(), rawKey) + if err != nil { + got.err = fmt.Errorf("k is not a P-256 point: %w", err) + return got + } + claims := jwt.MapClaims{} + parsed, err := jwt.ParseWithClaims(token, claims, func(*jwt.Token) (any, error) { return pub, nil }, + jwt.WithValidMethods([]string{"ES256"}), jwt.WithExpirationRequired()) + if err != nil { + got.err = fmt.Errorf("JWT does not verify: %w", err) + return got + } + got.claims = claims + got.jwtHeader = parsed.Header + body, err := io.ReadAll(r.Body) + if err != nil { + got.err = err + return got + } + got.payload, got.err = decryptForTest(body, sub.private, sub.sub.Auth) + return got +} + +func testConfig(keys VAPIDKeys, hosts ...string) Config { + return Config{ + Enabled: true, + PublicKey: keys.PublicKey, + PrivateKey: keys.PrivateKey, + Subject: "mailto:ops@example.com", + TTL: DefaultTTL, + AllowedHosts: hosts, + } +} + +func TestVAPIDSendRoundTrip(t *testing.T) { + keys, err := GenerateVAPIDKeys() + if err != nil { + t.Fatal(err) + } + if len(keys.PublicKey) != PublicKeyLength || len(keys.PrivateKey) != PrivateKeyLength { + t.Fatalf("key lengths %d/%d", len(keys.PublicKey), len(keys.PrivateKey)) + } + if strings.Contains(fmt.Sprintf("%v %#v %s", keys, keys, testConfig(keys)), keys.PrivateKey) { + t.Fatal("formatting VAPIDKeys or Config printed the private key") + } + + var sub testSubscriber + received := make(chan receivedPush, 4) + srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + switch r.URL.Path { + case "/gone": + w.WriteHeader(http.StatusGone) + return + case "/missing": + w.WriteHeader(http.StatusNotFound) + return + case "/boom": + w.WriteHeader(http.StatusInternalServerError) + _, _ = io.WriteString(w, "internal detail") + return + } + received <- verifyVAPIDRequest(r, sub) + w.WriteHeader(http.StatusCreated) + })) + defer srv.Close() + sub = newTestSubscriber(t, srv.URL+"/push/JzLQ3raZJfFBR0aqvOMsLrt54w4rJUsV") + + pusher := NewVAPIDPusher(testConfig(keys, "127.0.0.1"), srv.Client()) + now := time.Now() + payload := []byte(`{"title":"acme test","body":"hello"}`) + if err := pusher.Send(context.Background(), sub.sub, payload, SendOptions{Urgency: "high", Topic: "acme"}); err != nil { + t.Fatalf("Send: %v", err) + } + got := <-received + if got.err != nil { + t.Fatalf("push service rejected the request: %v", got.err) + } + if string(got.payload) != string(payload) { + t.Fatalf("decrypted payload %q", got.payload) + } + for name, want := range map[string]string{ + "Content-Encoding": "aes128gcm", + "Content-Type": "application/octet-stream", + "TTL": "2419200", + "Urgency": "high", + "Topic": "acme", + } { + if v := got.header.Get(name); v != want { + t.Errorf("header %s = %q, want %q", name, v, want) + } + } + if got.k != keys.PublicKey { + t.Errorf("k = %q, want the configured public key", got.k) + } + if got.jwtHeader["typ"] != "JWT" || got.jwtHeader["alg"] != "ES256" { + t.Errorf("JWT header %v", got.jwtHeader) + } + if got.claims["aud"] != srv.URL { + t.Errorf("aud = %v, want %s", got.claims["aud"], srv.URL) + } + if got.claims["sub"] != "mailto:ops@example.com" { + t.Errorf("sub = %v", got.claims["sub"]) + } + exp, err := got.claims.GetExpirationTime() + if err != nil || exp == nil { + t.Fatalf("exp: %v", err) + } + if ahead := exp.Sub(now); ahead <= 0 || ahead > 24*time.Hour { + t.Errorf("exp is %s ahead, want within 24h", ahead) + } + + // A per-send TTL replaces push.ttl; Urgency and Topic are omitted + // when unset. + if err := pusher.Send(context.Background(), sub.sub, []byte("x"), SendOptions{TTL: 90 * time.Second}); err != nil { + t.Fatalf("Send: %v", err) + } + got = <-received + if got.err != nil || got.header.Get("TTL") != "90" || got.header.Get("Urgency") != "" || got.header.Get("Topic") != "" { + t.Fatalf("second push: err=%v headers=%v", got.err, got.header) + } + + // Status mapping. + for path, check := range map[string]func(error) bool{ + "/gone": func(err error) bool { return errors.Is(err, ErrSubscriptionGone) }, + "/missing": func(err error) bool { return errors.Is(err, ErrSubscriptionGone) }, + "/boom": func(err error) bool { + var se *StatusError + return errors.As(err, &se) && se.Code == 500 && !strings.Contains(err.Error(), "internal detail") + }, + } { + s := sub.sub + s.Endpoint = srv.URL + path + if err := pusher.Send(context.Background(), s, []byte("x"), SendOptions{}); !check(err) { + t.Errorf("%s: unexpected error %v", path, err) + } + } +} + +func TestSendRefusesDisallowedEndpoint(t *testing.T) { + keys, err := GenerateVAPIDKeys() + if err != nil { + t.Fatal(err) + } + var hits atomic.Int32 + count := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + hits.Add(1) + w.WriteHeader(http.StatusCreated) + }) + plain := httptest.NewServer(count) + defer plain.Close() + tlsSrv := httptest.NewTLSServer(count) + defer tlsSrv.Close() + // A push service that redirects elsewhere: the redirect is not + // followed. + redirect := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + http.Redirect(w, r, tlsSrv.URL+"/internal", http.StatusFound) + })) + defer redirect.Close() + + sub := newTestSubscriber(t, "") + pusher := NewVAPIDPusher(testConfig(keys, "127.0.0.1"), tlsSrv.Client()) + port := tlsSrv.URL[strings.LastIndex(tlsSrv.URL, ":")+1:] + for name, endpoint := range map[string]string{ + "plain http": plain.URL + "/push/abc", + "unlisted host": "https://localhost:" + port + "/push/abc", + "foreign host": "https://evil.example.com/push/abc", + "user info": "https://user:secret@127.0.0.1:" + port + "/push/abc", + "relative URL": "/push/abc", + "other scheme": "ftp://127.0.0.1:" + port + "/push/abc", + "no host": "https:///push/abc", + "host suffix hit": "https://127.0.0.1.evil.example.com/push/abc", + } { + s := sub.sub + s.Endpoint = endpoint + err := pusher.Send(context.Background(), s, []byte("x"), SendOptions{}) + if !errors.Is(err, ErrEndpointNotAllowed) { + t.Errorf("%s: err = %v, want ErrEndpointNotAllowed", name, err) + } + if err != nil && strings.Contains(err.Error(), "/push/abc") { + t.Errorf("%s: error exposes the endpoint path: %v", name, err) + } + } + if n := hits.Load(); n != 0 { + t.Fatalf("%d requests reached a server for refused endpoints", n) + } + + s := sub.sub + s.Endpoint = redirect.URL + "/push/abc" + var se *StatusError + if err := pusher.Send(context.Background(), s, []byte("x"), SendOptions{}); !errors.As(err, &se) || se.Code != http.StatusFound { + t.Fatalf("redirect: err = %v, want StatusError 302", err) + } + if n := hits.Load(); n != 0 { + t.Fatal("the redirect was followed") + } + + disabled := testConfig(keys, "127.0.0.1") + disabled.Enabled = false + s.Endpoint = tlsSrv.URL + "/push/abc" + if err := NewVAPIDPusher(disabled, tlsSrv.Client()).Send(context.Background(), s, []byte("x"), SendOptions{}); !errors.Is(err, ErrPushDisabled) { + t.Fatalf("disabled: err = %v", err) + } + if n := hits.Load(); n != 0 { + t.Fatal("a disabled pusher sent a request") + } + + defaults := DefaultAllowedHosts() + for host, want := range map[string]bool{ + "fcm.googleapis.com": true, + "FCM.googleapis.com.": true, + "updates.push.services.mozilla.com": true, + "web.push.apple.com": true, + "wns2-by3p.notify.windows.com": true, + "push.apple.com": false, + "fcm.googleapis.com.evil.example": false, + "evilpush.apple.com": false, + "notify.windows.com": false, + "android.googleapis.com": false, + "127.0.0.1": false, + "updates.push.services.mozilla.com.evil": false, + } { + if got := HostAllowed(host, defaults); got != want { + t.Errorf("HostAllowed(%q) = %t, want %t", host, got, want) + } + } +} diff --git a/modules/flare/vapid.go b/modules/flare/vapid.go new file mode 100644 index 0000000..1b64014 --- /dev/null +++ b/modules/flare/vapid.go @@ -0,0 +1,140 @@ +package flare + +import ( + "crypto/ecdh" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/subtle" + "encoding/base64" + "errors" + "fmt" + "net/url" + "strings" + "time" + + "github.com/golang-jwt/jwt/v5" +) + +// VAPIDTokenLifetime is how far ahead of the request the exp claim of a +// VAPID token lies. RFC 8292 allows at most 24 hours. +const VAPIDTokenLifetime = 12 * time.Hour + +// Lengths of unpadded base64url VAPID keys: a 65-byte uncompressed P-256 +// point and a 32-byte scalar. +const ( + PublicKeyLength = 87 + PrivateKeyLength = 43 +) + +// ErrInvalidVAPIDKeys is returned when the configured key pair cannot be +// parsed or does not match. It never carries key material. +var ErrInvalidVAPIDKeys = errors.New("flare: invalid VAPID key pair") + +// ErrInvalidSubject is returned when the VAPID subject is not a mailto: or +// https: URI. +var ErrInvalidSubject = errors.New("flare: VAPID subject must start with mailto: or https:") + +// VAPIDKeys is an application server key pair (RFC 8292) as unpadded +// base64url: PublicKey is the 65-byte uncompressed P-256 point, PrivateKey +// the 32-byte private scalar. String and GoString never print PrivateKey. +type VAPIDKeys struct { + PublicKey string + PrivateKey string +} + +// String prints the public key and hides the private key. +func (k VAPIDKeys) String() string { + return "flare.VAPIDKeys{PublicKey: " + k.PublicKey + ", PrivateKey: [redacted]}" +} + +// GoString is String, so %#v does not print the private key either. +func (k VAPIDKeys) GoString() string { return k.String() } + +// GenerateVAPIDKeys returns a fresh P-256 key pair. +func GenerateVAPIDKeys() (VAPIDKeys, error) { + priv, err := ecdh.P256().GenerateKey(rand.Reader) + if err != nil { + return VAPIDKeys{}, fmt.Errorf("flare: generate VAPID keys: %w", err) + } + return VAPIDKeys{ + PublicKey: base64.RawURLEncoding.EncodeToString(priv.PublicKey().Bytes()), + PrivateKey: base64.RawURLEncoding.EncodeToString(priv.Bytes()), + }, nil +} + +// ParseVAPIDKeys decodes a key pair given as base64url, with or without +// padding, and checks that the public key belongs to the private key. +func ParseVAPIDKeys(public, private string) (*ecdh.PrivateKey, error) { + privRaw, err := decodeBase64URL(private) + if err != nil || len(privRaw) != 32 { + return nil, fmt.Errorf("%w: the private key must be 32 bytes of base64url", ErrInvalidVAPIDKeys) + } + priv, err := ecdh.P256().NewPrivateKey(privRaw) + if err != nil { + return nil, fmt.Errorf("%w: the private key is not a P-256 scalar", ErrInvalidVAPIDKeys) + } + pubRaw, err := decodeBase64URL(public) + if err != nil || len(pubRaw) != publicKeySize { + return nil, fmt.Errorf("%w: the public key must be 65 bytes of base64url", ErrInvalidVAPIDKeys) + } + if subtle.ConstantTimeCompare(pubRaw, priv.PublicKey().Bytes()) != 1 { + return nil, fmt.Errorf("%w: the public key does not match the private key", ErrInvalidVAPIDKeys) + } + return priv, nil +} + +// VAPIDHeader returns the RFC 8292 Authorization header value for a push to +// endpoint: "vapid t=, k=". The ES256 token carries aud +// (the endpoint's origin), exp (now + VAPIDTokenLifetime) and sub. +func VAPIDHeader(endpoint, subject string, keys VAPIDKeys, now time.Time) (string, error) { + if !validSubject(subject) { + return "", ErrInvalidSubject + } + aud, err := origin(endpoint) + if err != nil { + return "", err + } + priv, err := ParseVAPIDKeys(keys.PublicKey, keys.PrivateKey) + if err != nil { + return "", err + } + signer, err := ecdsa.ParseRawPrivateKey(elliptic.P256(), priv.Bytes()) + if err != nil { + return "", ErrInvalidVAPIDKeys + } + token := jwt.NewWithClaims(jwt.SigningMethodES256, jwt.MapClaims{ + "aud": aud, + "exp": now.Add(VAPIDTokenLifetime).Unix(), + "sub": subject, + }) + signed, err := token.SignedString(signer) + if err != nil { + return "", fmt.Errorf("flare: sign VAPID token: %w", err) + } + k := base64.RawURLEncoding.EncodeToString(priv.PublicKey().Bytes()) + return "vapid t=" + signed + ", k=" + k, nil +} + +func validSubject(subject string) bool { + return strings.HasPrefix(subject, "mailto:") || strings.HasPrefix(subject, "https:") +} + +// origin is the RFC 6454 serialization of the endpoint's origin: scheme and +// host, with the port only when it is not the scheme's default. +func origin(endpoint string) (string, error) { + u, err := url.Parse(endpoint) + if err != nil || u.Scheme == "" || u.Host == "" { + return "", fmt.Errorf("%w: not an absolute URL", ErrEndpointNotAllowed) + } + scheme := strings.ToLower(u.Scheme) + host := strings.ToLower(u.Hostname()) + if strings.Contains(host, ":") { + host = "[" + host + "]" + } + port := u.Port() + if port != "" && !(scheme == "https" && port == "443") && !(scheme == "http" && port == "80") { + host += ":" + port + } + return scheme + "://" + host, nil +}