test(11-07): cover flare VAPID, allowlist, statuses and config

- TestVAPIDHeader (origin rules, exp, subject), TestVAPIDKeys,
  TestSendAllowlist (T-11-22 host table), TestSendStatuses (2xx, 404/410,
  StatusError without body, disabled, host-only transport errors),
  TestFlareConfig, TestAgo, TestEncryptRejects (coverage 90.0%)
This commit is contained in:
Jakub Zych
2026-09-30 14:22:44 +02:00
parent 33194a1f98
commit 6dadbf6957
2 changed files with 388 additions and 0 deletions

View File

@@ -4,6 +4,7 @@ import (
"bytes" "bytes"
"crypto/ecdh" "crypto/ecdh"
"encoding/base64" "encoding/base64"
"errors"
"strings" "strings"
"testing" "testing"
) )
@@ -110,3 +111,35 @@ func TestRFC8291AppendixA(t *testing.T) {
t.Fatalf("3993-byte payload: %v", err) t.Fatalf("3993-byte payload: %v", err)
} }
} }
// TestEncryptRejects covers the RFC 8291 input checks: a payload over
// 3993 bytes, a p256dh that is not a 65-byte P-256 point, an auth secret
// that is not 16 bytes and non-base64url input are refused.
func TestEncryptRejects(t *testing.T) {
sub := newTestSubscriber(t, "https://fcm.googleapis.com/fcm/send/x").sub
if _, err := Encrypt(make([]byte, MaxPayloadSize+1), sub); !errors.Is(err, ErrPayloadTooLarge) {
t.Fatalf("oversized payload: %v", err)
}
if _, err := Encrypt(make([]byte, MaxPayloadSize), sub); err != nil {
t.Fatalf("payload at the limit: %v", err)
}
short := sub
short.P256dh = base64.RawURLEncoding.EncodeToString(make([]byte, 33))
offCurve := sub
offCurve.P256dh = base64.RawURLEncoding.EncodeToString(append([]byte{4}, make([]byte, 64)...))
badAuth := sub
badAuth.Auth = base64.RawURLEncoding.EncodeToString(make([]byte, 8))
notB64 := sub
notB64.Auth = "***"
for name, s := range map[string]Subscription{"short_key": short, "off_curve": offCurve, "short_auth": badAuth, "auth_not_base64": notB64} {
if _, err := Encrypt([]byte("x"), s); err == nil {
t.Errorf("%s accepted", name)
}
}
if b, err := decodeBase64URL("YWJj"); err != nil || string(b) != "abc" {
t.Fatalf("decode unpadded: %q %v", b, err)
}
if b, err := decodeBase64URL("YQ=="); err != nil || string(b) != "a" {
t.Fatalf("decode padded: %q %v", b, err)
}
}

355
modules/flare/flare_test.go Normal file
View File

@@ -0,0 +1,355 @@
package flare
import (
"context"
"crypto/ecdsa"
"crypto/elliptic"
"encoding/base64"
"errors"
"fmt"
"log/slog"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"git.golem15.com/golem15/summercms/modules/backpack"
"git.golem15.com/golem15/summercms/modules/compass"
"github.com/golang-jwt/jwt/v5"
)
// TestVAPIDHeader covers RFC 8292: aud is the endpoint's origin (default
// port stripped, other ports kept, IPv6 bracketed, lowercase), exp is now
// plus the token lifetime, sub must be mailto: or https:, and invalid keys
// or endpoints sign nothing.
func TestVAPIDHeader(t *testing.T) {
keys, err := GenerateVAPIDKeys()
if err != nil {
t.Fatal(err)
}
now := time.Date(2026, 9, 30, 12, 0, 0, 0, time.UTC)
cases := []struct {
endpoint, aud string
}{
{"https://fcm.googleapis.com/fcm/send/abc", "https://fcm.googleapis.com"},
{"https://FCM.GoogleAPIs.com:443/x", "https://fcm.googleapis.com"},
{"https://push.example.com:8443/x", "https://push.example.com:8443"},
{"http://push.example.com:80/x", "http://push.example.com"},
{"https://[2001:db8::1]:9443/x", "https://[2001:db8::1]:9443"},
}
for _, c := range cases {
h, err := VAPIDHeader(c.endpoint, "mailto:ops@example.com", keys, now)
if err != nil {
t.Fatalf("%s: %v", c.endpoint, err)
}
tok, k, ok := strings.Cut(strings.TrimPrefix(h, "vapid t="), ", k=")
if !ok || !strings.HasPrefix(h, "vapid t=") || k != keys.PublicKey {
t.Fatalf("header = %s", h)
}
pub, err := base64.RawURLEncoding.DecodeString(keys.PublicKey)
if err != nil {
t.Fatal(err)
}
claims := jwt.MapClaims{}
if _, err := jwt.ParseWithClaims(tok, claims, func(*jwt.Token) (any, error) { return ecdsaPublic(t, pub), nil },
jwt.WithValidMethods([]string{"ES256"}), jwt.WithTimeFunc(func() time.Time { return now })); err != nil {
t.Fatalf("%s: token does not verify: %v", c.endpoint, err)
}
if claims["aud"] != c.aud || claims["sub"] != "mailto:ops@example.com" || claims["exp"] != float64(now.Add(VAPIDTokenLifetime).Unix()) {
t.Fatalf("%s: claims = %v, want aud %s", c.endpoint, claims, c.aud)
}
}
if VAPIDTokenLifetime > 24*time.Hour {
t.Fatalf("token lifetime %s exceeds RFC 8292's 24h", VAPIDTokenLifetime)
}
if _, err := VAPIDHeader("https://fcm.googleapis.com/x", "https://example.com/contact", keys, now); err != nil {
t.Fatalf("https subject refused: %v", err)
}
for _, sub := range []string{"", "ops@example.com", "http://example.com", "MAILTO:ops@example.com"} {
if _, err := VAPIDHeader("https://fcm.googleapis.com/x", sub, keys, now); !errors.Is(err, ErrInvalidSubject) {
t.Errorf("subject %q: %v, want ErrInvalidSubject", sub, err)
}
}
for _, ep := range []string{"", "/relative", "fcm.googleapis.com/x", "://x"} {
if _, err := VAPIDHeader(ep, "mailto:a@b.c", keys, now); !errors.Is(err, ErrEndpointNotAllowed) {
t.Errorf("endpoint %q: %v, want ErrEndpointNotAllowed", ep, err)
}
}
if _, err := VAPIDHeader("https://fcm.googleapis.com/x", "mailto:a@b.c", VAPIDKeys{PublicKey: keys.PublicKey}, now); !errors.Is(err, ErrInvalidVAPIDKeys) {
t.Fatalf("missing private key: %v", err)
}
}
// TestVAPIDKeys covers key parsing: padded input, wrong lengths, a scalar
// outside the curve order and a public key of another pair.
func TestVAPIDKeys(t *testing.T) {
a, err := GenerateVAPIDKeys()
if err != nil {
t.Fatal(err)
}
b, err := GenerateVAPIDKeys()
if err != nil {
t.Fatal(err)
}
pad := func(s string) string { return s + strings.Repeat("=", (4-len(s)%4)%4) }
if _, err := ParseVAPIDKeys(pad(a.PublicKey), pad(a.PrivateKey)); err != nil {
t.Fatalf("padded keys refused: %v", err)
}
allFF := base64.RawURLEncoding.EncodeToString([]byte(strings.Repeat("\xff", 32)))
for name, pair := range map[string][2]string{
"mismatched": {b.PublicKey, a.PrivateKey},
"short_private": {a.PublicKey, a.PrivateKey[:20]},
"short_public": {a.PublicKey[:40], a.PrivateKey},
"not_base64": {a.PublicKey, "!!!!" + a.PrivateKey[4:]},
"scalar_too_big": {a.PublicKey, allFF},
"empty": {"", ""},
"public_not_url64": {"%%%%", a.PrivateKey},
} {
if _, err := ParseVAPIDKeys(pair[0], pair[1]); !errors.Is(err, ErrInvalidVAPIDKeys) {
t.Errorf("%s: %v, want ErrInvalidVAPIDKeys", name, err)
}
}
if s := fmt.Sprintf("%v|%#v|%s", a, a, Config{PrivateKey: a.PrivateKey}); strings.Contains(s, a.PrivateKey) {
t.Fatal("a formatted key pair shows the private key")
}
if v := (Config{PrivateKey: a.PrivateKey}).LogValue(); strings.Contains(v.String(), a.PrivateKey) {
t.Fatal("LogValue shows the private key")
}
if !strings.Contains((Config{}).GoString(), "PrivateKey: unset") {
t.Fatal("an empty private key is not reported as unset")
}
}
// TestSendAllowlist covers T-11-22's host rule: exact hosts, "*." entries
// for subdomains only, case and a trailing dot ignored.
func TestSendAllowlist(t *testing.T) {
allowed := []string{"fcm.googleapis.com", " *.push.apple.com ", "*.", ""}
cases := []struct {
host string
want bool
}{
{"fcm.googleapis.com", true},
{"FCM.GoogleAPIs.com.", true},
{"evil-fcm.googleapis.com", false},
{"fcm.googleapis.com.evil.test", false},
{"api.push.apple.com", true},
{"a.b.push.apple.com", true},
{"push.apple.com", false},
{"push.apple.com.evil.test", false},
{"xpush.apple.com", false},
{"", false},
{".", false},
}
for _, c := range cases {
if got := HostAllowed(c.host, allowed); got != c.want {
t.Errorf("HostAllowed(%q) = %v, want %v", c.host, got, c.want)
}
}
for _, h := range DefaultAllowedHosts() {
if !strings.Contains(h, ".") {
t.Errorf("default host %q", h)
}
}
p := NewVAPIDPusher(Config{Enabled: true, AllowedHosts: allowed}, nil)
for _, ep := range []string{"http://fcm.googleapis.com/x", "https://user@fcm.googleapis.com/x", "https://push.apple.com/x", "mailto:x", "https:///x"} {
if err := p.Send(context.Background(), Subscription{Endpoint: ep}, []byte("{}"), SendOptions{}); !errors.Is(err, ErrEndpointNotAllowed) {
t.Errorf("Send(%q) = %v, want ErrEndpointNotAllowed", ep, err)
}
}
}
// TestSendStatuses covers the answers of a push service: 2xx success, 404
// and 410 ErrSubscriptionGone, any other status a StatusError without the
// body; a disabled or nil pusher sends nothing; a transport failure names
// the host, never the endpoint path.
func TestSendStatuses(t *testing.T) {
keys, err := GenerateVAPIDKeys()
if err != nil {
t.Fatal(err)
}
var hits int
srv := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
hits++
switch r.URL.Path {
case "/ok":
w.WriteHeader(http.StatusCreated)
case "/accepted":
w.WriteHeader(http.StatusAccepted)
case "/missing":
w.WriteHeader(http.StatusNotFound)
case "/gone":
w.WriteHeader(http.StatusGone)
case "/limited":
w.WriteHeader(http.StatusTooManyRequests)
default:
w.WriteHeader(http.StatusInternalServerError)
_, _ = w.Write([]byte("internal detail"))
}
}))
defer srv.Close()
p := NewVAPIDPusher(testConfig(keys, "127.0.0.1"), srv.Client())
send := func(path string) error {
return p.Send(context.Background(), newTestSubscriber(t, srv.URL+path).sub, []byte(`{"t":1}`), SendOptions{TTL: time.Minute})
}
for path, want := range map[string]error{"/ok": nil, "/accepted": nil, "/missing": ErrSubscriptionGone, "/gone": ErrSubscriptionGone} {
if err := send(path); !errors.Is(err, want) {
t.Errorf("%s: %v, want %v", path, err, want)
}
}
for path, code := range map[string]int{"/boom": 500, "/limited": 429} {
err := send(path)
var se *StatusError
if !errors.As(err, &se) || se.Code != code || strings.Contains(err.Error(), "internal detail") {
t.Errorf("%s: %v, want a StatusError %d without the body", path, err, code)
}
}
before := hits
off := NewVAPIDPusher(Config{AllowedHosts: []string{"127.0.0.1"}}, srv.Client())
if err := off.Send(context.Background(), newTestSubscriber(t, srv.URL+"/ok").sub, []byte("{}"), SendOptions{}); !errors.Is(err, ErrPushDisabled) {
t.Fatalf("disabled: %v", err)
}
var nilPusher *VAPIDPusher
if err := nilPusher.Send(context.Background(), Subscription{}, nil, SendOptions{}); !errors.Is(err, ErrPushDisabled) {
t.Fatalf("nil pusher: %v", err)
}
if hits != before {
t.Fatal("a disabled pusher sent a request")
}
dead := httptest.NewTLSServer(http.NotFoundHandler())
url := dead.URL
dead.Close()
err = p.Send(nil, newTestSubscriber(t, url+"/secret-capability-path").sub, []byte("{}"), SendOptions{})
if err == nil || strings.Contains(err.Error(), "secret-capability-path") || !strings.Contains(err.Error(), "127.0.0.1") {
t.Fatalf("transport error = %v, want the host without the path", err)
}
if hostOf("://bad") != "endpoint" {
t.Fatal("hostOf of an unparsable endpoint")
}
bad := newTestSubscriber(t, srv.URL+"/ok").sub
bad.P256dh = "not-a-key"
if err := p.Send(context.Background(), bad, []byte("{}"), SendOptions{}); err == nil {
t.Fatal("a subscription with a bad key was sent")
}
noSubject := testConfig(keys, "127.0.0.1")
noSubject.Subject = ""
if err := NewVAPIDPusher(noSubject, srv.Client()).Send(context.Background(), newTestSubscriber(t, srv.URL+"/ok").sub, []byte("{}"), SendOptions{}); !errors.Is(err, ErrInvalidSubject) {
t.Fatalf("empty subject: %v", err)
}
}
// TestFlareConfig covers push.* parsing and the Service accessors.
func TestFlareConfig(t *testing.T) {
def := LoadConfig(nil)
if def.Enabled || def.TTL != DefaultTTL || len(def.AllowedHosts) != len(DefaultAllowedHosts()) {
t.Fatalf("defaults = %v", def)
}
app := func(t *testing.T, kv map[string]any) *backpack.App {
t.Helper()
cfg, err := compass.Open(compass.Options{Dir: t.TempDir(), Env: "testing", Environ: []string{}})
if err != nil {
t.Fatal(err)
}
for k, v := range kv {
if err := cfg.Set(k, v); err != nil {
t.Fatal(err)
}
}
return backpack.New(cfg)
}
cases := []struct {
name string
kv map[string]any
ttl time.Duration
hosts string
}{
{"seconds", map[string]any{"push.ttl": 60}, time.Minute, strings.Join(DefaultAllowedHosts(), ",")},
{"duration", map[string]any{"push.ttl": "90m"}, 90 * time.Minute, strings.Join(DefaultAllowedHosts(), ",")},
{"invalid_ttl", map[string]any{"push.ttl": "soon"}, DefaultTTL, strings.Join(DefaultAllowedHosts(), ",")},
{"comma_hosts", map[string]any{"push.allowed_hosts": " A.example.com, ,*.b.example.com "}, DefaultTTL, "a.example.com,*.b.example.com"},
{"list_hosts", map[string]any{"push.allowed_hosts": []any{"c.example.com", 3}}, DefaultTTL, "c.example.com"},
{"empty_hosts_keep_default", map[string]any{"push.allowed_hosts": ""}, DefaultTTL, strings.Join(DefaultAllowedHosts(), ",")},
}
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
got := LoadConfig(app(t, c.kv).Config)
if got.TTL != c.ttl || strings.Join(got.AllowedHosts, ",") != c.hosts {
t.Fatalf("config = %v", got)
}
})
}
if got := stringList([]string{" X ", ""}); strings.Join(got, ",") != "x" {
t.Fatalf("stringList = %v", got)
}
if got := stringList(42); len(got) != 0 {
t.Fatalf("stringList(42) = %v", got)
}
a := app(t, map[string]any{"push.enabled": true, "push.subject": "mailto:ops@example.com"})
logger := slog.New(slog.NewTextHandler(new(strings.Builder), nil))
if err := a.Publish(logger); err != nil {
t.Fatal(err)
}
svc, err := From(a)
if err != nil {
t.Fatal(err)
}
if again, err := From(a); err != nil || again != svc {
t.Fatal("From is not idempotent")
}
if !svc.Enabled() || svc.Config().Subject != "mailto:ops@example.com" || svc.Logger() != logger || svc.Pusher() == nil {
t.Fatal("service accessors")
}
svc.SetHTTPClient(&http.Client{Timeout: time.Second})
if svc.Pusher().(*VAPIDPusher).hc.Timeout != time.Second || svc.Pusher().(*VAPIDPusher).hc.CheckRedirect == nil {
t.Fatal("SetHTTPClient did not keep redirects refused")
}
var nilSvc *Service
nilSvc.SetHTTPClient(nil)
if nilSvc.Enabled() || nilSvc.Pusher() == nil || nilSvc.Logger() == nil || nilSvc.Config().TTL != DefaultTTL {
t.Fatal("nil service accessors")
}
if _, err := From(nil); err == nil {
t.Fatal("From(nil) succeeded")
}
}
// TestAgo covers the relative time the test-push command prints.
func TestAgo(t *testing.T) {
now := time.Date(2026, 9, 30, 12, 0, 0, 0, time.UTC)
at := func(d time.Duration) *time.Time { v := now.Add(-d); return &v }
cases := map[string]*time.Time{
"unknown": nil,
"just now": at(-time.Minute),
"1 second ago": at(time.Second),
"5 minutes ago": at(5 * time.Minute),
"1 hour ago": at(time.Hour),
"3 days ago": at(72 * time.Hour),
"2 weeks ago": at(15 * 24 * time.Hour),
"2 months ago": at(65 * 24 * time.Hour),
"1 year ago": at(400 * 24 * time.Hour),
}
for want, ts := range cases {
if got := ago(ts, now); got != want {
t.Errorf("ago = %q, want %q", got, want)
}
}
zero := time.Time{}
if ago(&zero, now) != "unknown" {
t.Fatal("zero time")
}
if truncateKey("short") != "..." || truncateKey("abcdefghijklmnopqrstuvwxyz") != "abcdefgh...wxyz" {
t.Fatal("truncateKey")
}
}
func ecdsaPublic(t *testing.T, raw []byte) *ecdsa.PublicKey {
t.Helper()
pub, err := ecdsa.ParseUncompressedPublicKey(elliptic.P256(), raw)
if err != nil {
t.Fatal(err)
}
return pub
}