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:
Jakub Zych
2026-09-30 13:43:09 +02:00
parent a8305a015d
commit a9af0d77c7
7 changed files with 1337 additions and 0 deletions

127
modules/flare/README.md Normal file
View 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
View 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)
}

View 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
View 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
View 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
View 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
}