- 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
372 lines
11 KiB
Go
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)
|
|
}
|
|
}
|
|
}
|