Files
summercms/modules/flare/send_test.go
Jakub Zych a9af0d77c7 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
2026-09-30 13:43:09 +02:00

372 lines
11 KiB
Go

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