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:
@@ -4,6 +4,7 @@ import (
|
||||
"bytes"
|
||||
"crypto/ecdh"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
@@ -110,3 +111,35 @@ func TestRFC8291AppendixA(t *testing.T) {
|
||||
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
355
modules/flare/flare_test.go
Normal 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
|
||||
}
|
||||
Reference in New Issue
Block a user