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 }