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