feat(11-04): add flare Web Push with a stdlib VAPID driver
- RFC 8291 aes128gcm encryption from crypto/ecdh, crypto/hkdf and AES-GCM, matching the RFC 8291 Appendix A vector byte for byte - RFC 8292 vapid t=<ES256 JWT>, k=<key> header (aud origin, exp +12h, sub) - Pusher, Subscription, SendOptions, SubscriptionSource, Service and From reading push.* (enabled, keys, subject, ttl, allowed_hosts) - sends only to https endpoints on push.allowed_hosts, checked before dialing, and never follows redirects; 404/410 map to ErrSubscriptionGone - module README and root modules row
This commit is contained in:
@@ -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. |
|
| [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. |
|
| [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. |
|
| [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. |
|
| [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. |
|
| [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. |
|
| [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. |
|
||||||
|
|||||||
127
modules/flare/README.md
Normal file
127
modules/flare/README.md
Normal file
@@ -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=<JWT>, k=<public key>` 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.
|
||||||
163
modules/flare/encrypt.go
Normal file
163
modules/flare/encrypt.go
Normal file
@@ -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)
|
||||||
|
}
|
||||||
112
modules/flare/encrypt_test.go
Normal file
112
modules/flare/encrypt_test.go
Normal file
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
423
modules/flare/flare.go
Normal file
423
modules/flare/flare.go
Normal file
@@ -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()
|
||||||
|
}
|
||||||
371
modules/flare/send_test.go
Normal file
371
modules/flare/send_test.go
Normal file
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
140
modules/flare/vapid.go
Normal file
140
modules/flare/vapid.go
Normal file
@@ -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=<JWT>, k=<public key>". 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
|
||||||
|
}
|
||||||
Reference in New Issue
Block a user