From 33194a1f98fac824e6f98444be10a904d7b266ac Mon Sep 17 00:00:00 2001 From: Jakub Zych Date: Wed, 30 Sep 2026 14:21:03 +0200 Subject: [PATCH] test(11-07): cover lighthouse realtime and the Centrifugo driver - lighthouse: TestSuppression (Widget silenced, Gadget not, nesting, stale outer ctx), TestBulkEmitsOnce, TestBroadcastEdges (zero-key batch, update actor, id-only delete, method contract, multi-channel, savepoint), TestBroadcastPublishFailure, TestFromSelectsDriver, TestMountSurfaces, TestRegistry under -race, drivers, args JSON, Bind (coverage 91.7%) - centrifugo: TestTokenClaims, TestTokenHandler, TestClientRequests, TestClientLoadConfig and a TestProxy table porting the WinterCMS WS-005, WS-007 and WS-013 cases (coverage 92.4%) --- modules/lighthouse/broadcast_test.go | 424 +++++++++++++++++++ modules/lighthouse/centrifugo/client_test.go | 260 ++++++++++++ modules/lighthouse/centrifugo/proxy_test.go | 222 ++++++++++ modules/lighthouse/centrifugo/token_test.go | 210 +++++++++ modules/lighthouse/channel_test.go | 45 +- modules/lighthouse/lighthouse_test.go | 248 +++++++++++ modules/lighthouse/postgres_test.go | 2 +- modules/lighthouse/registry_test.go | 73 ++++ modules/lighthouse/route_test.go | 142 +++++++ 9 files changed, 1593 insertions(+), 33 deletions(-) create mode 100644 modules/lighthouse/centrifugo/client_test.go create mode 100644 modules/lighthouse/centrifugo/proxy_test.go create mode 100644 modules/lighthouse/centrifugo/token_test.go create mode 100644 modules/lighthouse/lighthouse_test.go create mode 100644 modules/lighthouse/registry_test.go create mode 100644 modules/lighthouse/route_test.go diff --git a/modules/lighthouse/broadcast_test.go b/modules/lighthouse/broadcast_test.go index 371e715..72e1c33 100644 --- a/modules/lighthouse/broadcast_test.go +++ b/modules/lighthouse/broadcast_test.go @@ -7,12 +7,14 @@ import ( "errors" "fmt" "log/slog" + "reflect" "strings" "sync" "testing" "time" "git.golem15.com/golem15/summercms/modules/backpack" + "git.golem15.com/golem15/summercms/modules/bouncer" "git.golem15.com/golem15/summercms/modules/compass" "git.golem15.com/golem15/summercms/modules/conga" "git.golem15.com/golem15/summercms/modules/lagoon" @@ -360,3 +362,425 @@ func TestBroadcastTx(t *testing.T) { } }) } + +// Sprocket broadcasts through every optional method contract. +type Sprocket struct { + ID uint `gorm:"column:id;primaryKey"` + Name string `gorm:"column:name"` + Channels string `gorm:"column:channels"` +} + +func (Sprocket) TableName() string { return "acme_sprockets" } + +func (s *Sprocket) BroadcastChannels(_ context.Context, tx *gorm.DB) ([]string, error) { + if s.Channels == "" { + return nil, nil + } + if s.Channels == "fail" { + return nil, errors.New("channels failed") + } + if s.Channels == "abort-tx" { + // A failed statement aborts a Postgres transaction unless it runs + // inside a savepoint. + return nil, tx.Exec(`SELECT * FROM acme_missing_table`).Error + } + return strings.Split(s.Channels, ","), nil +} + +func (Sprocket) BroadcastAlias() string { return "acme.sprocket" } +func (Sprocket) BroadcastTTL() int { return 15 } + +func (s *Sprocket) BroadcastPayload(_ context.Context, _ *gorm.DB, ev Event) (any, error) { + if s.Name == "bad-payload" { + return nil, errors.New("payload failed") + } + return struct { + ID uint `json:"id"` + Action Action `json:"action"` + TTL int `json:"ttl"` + Actor Actor `json:"actor"` + }{s.ID, ev.Action, ev.TTL, ev.Actor}, nil +} + +func (s *Sprocket) ShouldBroadcast(a Action) bool { return a != ActionUpdated } + +// TestSuppression covers D-08: WithoutBroadcasting[T] silences one model +// type for the ctx it hands out and nothing else. Another type still +// broadcasts, nesting adds types, a pointer type argument names the +// struct, and a write through a stale outer ctx is not suppressed. +func TestSuppression(t *testing.T) { + env := newLHEnv(t, nil, func(t *testing.T, svc *Service) { + if err := Bind[Gadget](svc, Binding[Gadget]{Channels: func(_ context.Context, _ *gorm.DB, m *Gadget) ([]string, error) { + return []string{fmt.Sprintf("gadgets:%d", m.OwnerID)}, nil + }}); err != nil { + t.Fatal(err) + } + }) + mem := env.memory(t) + ctx := t.Context() + events := func(pubs []Publication) []string { + var out []string + for _, p := range pubs { + out = append(out, p.Event+"@"+strings.Join(p.Channels, ",")) + } + return out + } + + t.Run("silences_widget_not_gadget", func(t *testing.T) { + from := len(mem.Publications()) + err := WithoutBroadcasting[Widget](ctx, func(ctx context.Context) error { + return lagoon.Transaction(ctx, env.gdb, func(ctx context.Context, tx *gorm.DB) error { + if err := tx.WithContext(ctx).Create(&Widget{Name: "quiet", OwnerID: 1}).Error; err != nil { + return err + } + return tx.WithContext(ctx).Create(&Gadget{Name: "loud", OwnerID: 1}).Error + }) + }) + if err != nil { + t.Fatal(err) + } + got := events(waitPublications(t, mem, from, 1)) + if len(got) != 1 || got[0] != "created.lighthouse.gadget@gadgets:1" { + t.Fatalf("publications = %v, want only the gadget", got) + } + }) + + t.Run("nesting_and_pointer_type_argument", func(t *testing.T) { + before := env.broadcastJobs(t) + err := WithoutBroadcasting[*Widget](ctx, func(ctx context.Context) error { + return WithoutBroadcasting[Gadget](ctx, func(ctx context.Context) error { + if err := env.gdb.WithContext(ctx).Create(&Widget{Name: "nested", OwnerID: 2}).Error; err != nil { + return err + } + return env.gdb.WithContext(ctx).Create(&Gadget{Name: "nested", OwnerID: 2}).Error + }) + }) + if err != nil { + t.Fatal(err) + } + if after := env.broadcastJobs(t); after != before { + t.Fatalf("broadcast jobs %d -> %d with both types suppressed", before, after) + } + }) + + t.Run("stale_outer_ctx_is_not_suppressed", func(t *testing.T) { + from := len(mem.Publications()) + outer := ctx + err := WithoutBroadcasting[Widget](ctx, func(context.Context) error { + return env.gdb.WithContext(outer).Create(&Widget{Name: "outer", OwnerID: 3}).Error + }) + if err != nil { + t.Fatal(err) + } + got := events(waitPublications(t, mem, from, 1)) + if got[0] != "created.lighthouse.widget@widgets:3" { + t.Fatalf("publications = %v", got) + } + }) + + t.Run("fn_error_is_returned", func(t *testing.T) { + boom := errors.New("boom") + if err := WithoutBroadcasting[Widget](nil, func(ctx context.Context) error { + if ctx == nil { + t.Error("nil ctx was not replaced") + } + return boom + }); !errors.Is(err, boom) { + t.Fatalf("err = %v", err) + } + if suppressed(nil, reflect.TypeFor[Widget]()) { + t.Fatal("a nil ctx suppresses") + } + }) +} + +type bulkPayload struct { + Reason string `json:"reason"` + Count int `json:"count"` +} + +// TestBulkEmitsOnce covers D-08 and RT-03: N suppressed creates plus one +// Emit in the same transaction publish exactly one summary event, with the +// payload bytes in declaration order, and nothing when the transaction +// rolls back. +func TestBulkEmitsOnce(t *testing.T) { + env := newLHEnv(t, map[string]any{"realtime.broadcast_namespace": "Acme"}, nil) + mem := env.memory(t) + ctx := t.Context() + + from := len(mem.Publications()) + err := WithoutBroadcasting[Widget](ctx, func(ctx context.Context) error { + return lagoon.Transaction(ctx, env.gdb, func(ctx context.Context, tx *gorm.DB) error { + for i := range 3 { + if err := tx.WithContext(ctx).Create(&Widget{Name: fmt.Sprintf("bulk-%d", i), OwnerID: 9}).Error; err != nil { + return err + } + } + return env.svc.Emit(ctx, tx, Broadcast{Channels: []string{"Widgets:9"}, Event: "collection.bulk_updated", Payload: bulkPayload{Reason: "bulk_create", Count: 3}}) + }) + }) + if err != nil { + t.Fatal(err) + } + p := waitPublications(t, mem, from, 1)[0] + if p.Event != "collection.bulk_updated" || p.Channels[0] != "acme:widgets:9" || p.Method != "publish" { + t.Fatalf("publication = %+v", p) + } + if string(p.Payload) != `{"reason":"bulk_create","count":3}` { + t.Fatalf("payload = %s, want the ordered bytes", p.Payload) + } + + before := env.broadcastJobs(t) + rollback := errors.New("rollback") + err = lagoon.Transaction(ctx, env.gdb, func(ctx context.Context, tx *gorm.DB) error { + if err := env.svc.Emit(ctx, tx, Broadcast{Channels: []string{"widgets:9"}, Event: "collection.bulk_updated", Payload: bulkPayload{Reason: "bulk_create", Count: 1}}); err != nil { + return err + } + return rollback + }) + if !errors.Is(err, rollback) { + t.Fatalf("err = %v", err) + } + if after := env.broadcastJobs(t); after != before { + t.Fatalf("an Emit in a rolled-back transaction left %d job(s)", after-before) + } + + // Emit without a transaction enqueues on its own; several channels are + // one broadcast. + from = len(mem.Publications()) + if err := env.svc.Emit(nil, nil, Broadcast{Channels: []string{"a:1", "b:2"}, Event: "acme.pinged", Payload: nil}); err != nil { + t.Fatal(err) + } + p = waitPublications(t, mem, from, 1)[0] + if p.Method != "broadcast" || strings.Join(p.Channels, ",") != "acme:a:1,acme:b:2" || string(p.Payload) != "null" { + t.Fatalf("publication = %+v payload %s", p, p.Payload) + } + + if err := env.svc.Emit(ctx, env.gdb, Broadcast{Event: "x"}); err != nil { + t.Fatalf("Emit without channels = %v, want nil", err) + } + if err := env.svc.Emit(ctx, env.gdb, Broadcast{Channels: []string{"a:1"}}); err == nil { + t.Fatal("Emit without an event name succeeded") + } + if err := env.svc.Emit(ctx, env.gdb, Broadcast{Channels: []string{"a:1"}, Event: "x", Payload: make(chan int)}); err == nil { + t.Fatal("Emit with an unencodable payload succeeded") + } + var nilSvc *Service + if err := nilSvc.Emit(ctx, env.gdb, Broadcast{Channels: []string{"a:1"}, Event: "x"}); err == nil { + t.Fatal("Emit on a nil service succeeded") + } +} + +// failingDriver fails every publication. +type failingDriver struct{ calls *int32 } + +func (failingDriver) Name() string { return "acme-failing" } +func (failingDriver) Routes() []Route { return nil } +func (d failingDriver) Publish(context.Context, string, string, json.RawMessage) error { + return errors.New("realtime server is down") +} +func (d failingDriver) Broadcast(context.Context, []string, string, json.RawMessage) error { + return errors.New("realtime server is down") +} + +func init() { + RegisterDriver("acme-failing", func(*backpack.App, *Service) (Driver, error) { return failingDriver{}, nil }) +} + +// TestBroadcastEdges covers the remaining broadcast rules: a zero-key batch +// update is skipped, updates and deletes broadcast (a delete through an +// id-only model carries the reloaded row), the method-based contract +// (alias, ttl, payload, filter), multi-channel broadcasts, the actor of a +// frontend and a backend principal, failures inside the savepoint that +// never abort the write, and a failed publish that is logged and not +// retried. +func TestBroadcastEdges(t *testing.T) { + env := newLHEnv(t, nil, nil) + name := "Ann" + env.svc.SetUserLookup(func(_ context.Context, id uint) (User, bool, error) { + switch id { + case 7: + return User{ID: 7, Name: &name}, true, nil + case 8: + return User{}, false, errors.New("lookup failed") + } + return User{}, false, nil + }) + mem := env.memory(t) + ctx := t.Context() + + w := Widget{Name: "edge", OwnerID: 4} + if err := env.gdb.WithContext(ctx).Create(&w).Error; err != nil { + t.Fatal(err) + } + waitPublications(t, mem, 0, 1) + + t.Run("zero_key_batch_update_is_skipped", func(t *testing.T) { + before := env.broadcastJobs(t) + if err := env.gdb.WithContext(ctx).Model(&Widget{}).Where("owner_id = ?", 4).Update("name", "batch").Error; err != nil { + t.Fatal(err) + } + if after := env.broadcastJobs(t); after != before { + t.Fatalf("a zero-key batch update enqueued %d job(s)", after-before) + } + }) + + t.Run("update_carries_the_frontend_actor", func(t *testing.T) { + from := len(mem.Publications()) + uctx := bouncer.WithUser(ctx, &bouncer.Principal{ID: 7}) + w.Name = "renamed" + if err := env.gdb.WithContext(uctx).Save(&w).Error; err != nil { + t.Fatal(err) + } + p := waitPublications(t, mem, from, 1)[0] + var body struct { + Actor Actor `json:"actor"` + } + if err := json.Unmarshal(p.Payload, &body); err != nil { + t.Fatal(err) + } + if p.Event != "updated.lighthouse.widget" || body.Actor.UserID == nil || *body.Actor.UserID != 7 || body.Actor.Name == nil || *body.Actor.Name != "Ann" { + t.Fatalf("publication %s payload %s", p.Event, p.Payload) + } + }) + + t.Run("actor_rules", func(t *testing.T) { + if a := env.svc.Actor(bouncer.WithUser(ctx, &bouncer.Principal{ID: 3, Backend: true})); a.UserID != nil || *a.Name != "System" { + t.Fatalf("backend actor = %+v, want System", a) + } + if a := env.svc.Actor(bouncer.WithUser(ctx, &bouncer.Principal{ID: 8})); a.UserID == nil || *a.UserID != 8 || a.Name != nil { + t.Fatalf("lookup-error actor = %+v, want the id and a null name", a) + } + if a := env.svc.Actor(bouncer.WithUser(ctx, &bouncer.Principal{ID: 99})); a.UserID == nil || a.Name != nil { + t.Fatalf("unknown-user actor = %+v", a) + } + raw, _ := json.Marshal(SystemActor()) + if string(raw) != `{"user_id":null,"name":"System"}` { + t.Fatalf("SystemActor = %s", raw) + } + }) + + t.Run("delete_of_an_id_only_model_carries_the_row", func(t *testing.T) { + from := len(mem.Publications()) + if err := env.gdb.WithContext(ctx).Delete(&Widget{ID: w.ID}).Error; err != nil { + t.Fatal(err) + } + p := waitPublications(t, mem, from, 1)[0] + var body struct { + Model Widget `json:"model"` + } + if err := json.Unmarshal(p.Payload, &body); err != nil { + t.Fatal(err) + } + if p.Event != "deleted.lighthouse.widget" || p.Channels[0] != "widgets:4" || body.Model.Name != "renamed" { + t.Fatalf("publication %s %v payload %s", p.Event, p.Channels, p.Payload) + } + }) + + t.Run("method_contract_and_multi_channel", func(t *testing.T) { + from := len(mem.Publications()) + s := Sprocket{Name: "multi", Channels: "Room:1,room:2"} + if err := env.gdb.WithContext(ctx).Create(&s).Error; err != nil { + t.Fatal(err) + } + p := waitPublications(t, mem, from, 1)[0] + if p.Method != "broadcast" || strings.Join(p.Channels, ",") != "room:1,room:2" || p.Event != "created.acme.sprocket" { + t.Fatalf("publication = %+v", p) + } + if !strings.HasPrefix(string(p.Payload), fmt.Sprintf(`{"id":%d,"action":"created","ttl":15,`, s.ID)) { + t.Fatalf("payload = %s", p.Payload) + } + before := env.broadcastJobs(t) + s.Name = "filtered" + if err := env.gdb.WithContext(ctx).Save(&s).Error; err != nil { + t.Fatal(err) + } + if after := env.broadcastJobs(t); after != before { + t.Fatal("ShouldBroadcast(updated) = false did not veto the update") + } + if err := env.gdb.WithContext(ctx).Create(&Sprocket{Name: "silent"}).Error; err != nil { + t.Fatal(err) + } + if after := env.broadcastJobs(t); after != before { + t.Fatal("a model without channels enqueued a job") + } + }) + + t.Run("failures_are_logged_and_never_abort_the_write", func(t *testing.T) { + before := env.broadcastJobs(t) + err := lagoon.Transaction(ctx, env.gdb, func(ctx context.Context, tx *gorm.DB) error { + if err := tx.WithContext(ctx).Create(&Sprocket{Name: "chan-fail", Channels: "fail"}).Error; err != nil { + return err + } + if err := tx.WithContext(ctx).Create(&Sprocket{Name: "bad-payload", Channels: "room:3"}).Error; err != nil { + return err + } + // The transaction is still usable after both failures. + return tx.WithContext(ctx).Create(&Widget{Name: "after-failures", OwnerID: 0}).Error + }) + if err != nil { + t.Fatalf("a broadcast failure aborted the write: %v", err) + } + var n int64 + if err := env.gdb.Model(&Sprocket{}).Where("name IN ?", []string{"chan-fail", "bad-payload"}).Count(&n).Error; err != nil || n != 2 { + t.Fatalf("committed sprockets = %d (err %v), want 2", n, err) + } + if after := env.broadcastJobs(t); after != before { + t.Fatalf("failed broadcasts enqueued %d job(s)", after-before) + } + if env.logs.count("realtime: broadcast skipped") < 2 { + t.Fatalf("no skip warnings:\n%s", env.logs.all()) + } + }) + + t.Run("a_failed_query_inside_the_savepoint_keeps_the_transaction", func(t *testing.T) { + err := lagoon.Transaction(ctx, env.gdb, func(ctx context.Context, tx *gorm.DB) error { + if err := tx.WithContext(ctx).Create(&Sprocket{Name: "abort", Channels: "abort-tx"}).Error; err != nil { + return err + } + return tx.WithContext(ctx).Create(&Widget{Name: "after-abort", OwnerID: 0}).Error + }) + if err != nil { + t.Fatalf("a failed query in the broadcast savepoint aborted the write: %v", err) + } + var n int64 + if err := env.gdb.Model(&Widget{}).Where("name = ?", "after-abort").Count(&n).Error; err != nil || n != 1 { + t.Fatalf("write after the failed broadcast query: %d rows (err %v)", n, err) + } + }) +} + +// TestBroadcastPublishFailure covers D-09: a failed publish is logged at +// Warn without the payload and the one-attempt job is not retried. +func TestBroadcastPublishFailure(t *testing.T) { + env := newLHEnv(t, map[string]any{"realtime.driver": "acme-failing"}, nil) + ctx := t.Context() + if err := env.gdb.WithContext(ctx).Create(&Widget{Name: "secret-payload-value", OwnerID: 6}).Error; err != nil { + t.Fatal(err) + } + deadline := time.Now().Add(10 * time.Second) + for env.logs.count("realtime: broadcast failed") == 0 { + if time.Now().After(deadline) { + t.Fatalf("no failure log:\n%s", env.logs.all()) + } + time.Sleep(20 * time.Millisecond) + } + rec, _ := env.logs.find("realtime: broadcast failed") + if rec.level != slog.LevelWarn || !strings.Contains(rec.attrs["channels"], "widgets:6") || strings.Contains(env.logs.all(), "secret-payload-value") { + t.Fatalf("failure log = %+v", rec) + } + var state string + var attempt int + for time.Now().Before(deadline) { + if err := env.db.QueryRowContext(ctx, `SELECT state::text, attempt FROM river_job WHERE kind = 'summer.broadcast' ORDER BY id DESC LIMIT 1`).Scan(&state, &attempt); err != nil { + t.Fatal(err) + } + if state == "completed" { + break + } + time.Sleep(20 * time.Millisecond) + } + if state != "completed" || attempt != 1 { + t.Fatalf("job state %q attempt %d, want completed after one attempt (never retried)", state, attempt) + } +} diff --git a/modules/lighthouse/centrifugo/client_test.go b/modules/lighthouse/centrifugo/client_test.go new file mode 100644 index 0000000..25cb3f6 --- /dev/null +++ b/modules/lighthouse/centrifugo/client_test.go @@ -0,0 +1,260 @@ +package centrifugo + +import ( + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + "git.golem15.com/golem15/summercms/modules/backpack" + "git.golem15.com/golem15/summercms/modules/compass" + "git.golem15.com/golem15/summercms/modules/lighthouse" +) + +type apiCall struct { + method, path, auth, contentType, body string +} + +type fakeAPI struct { + mu sync.Mutex + calls []apiCall + status int + answer string +} + +func newFakeAPI(t *testing.T) (*fakeAPI, *httptest.Server) { + t.Helper() + f := &fakeAPI{status: http.StatusOK, answer: `{}`} + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + f.mu.Lock() + f.calls = append(f.calls, apiCall{r.Method, r.URL.Path, r.Header.Get("Authorization"), r.Header.Get("Content-Type"), string(body)}) + status, answer := f.status, f.answer + f.mu.Unlock() + w.WriteHeader(status) + _, _ = io.WriteString(w, answer) + })) + t.Cleanup(srv.Close) + return f, srv +} + +func (f *fakeAPI) set(status int, answer string) { + f.mu.Lock() + f.status, f.answer = status, answer + f.mu.Unlock() +} + +func (f *fakeAPI) take() []apiCall { + f.mu.Lock() + defer f.mu.Unlock() + out := f.calls + f.calls = nil + return out +} + +func testClient(url, key string) *Client { + c := NewClient(Config{APIURL: url + "/api", APIKey: key}, nil) + c.now = func() time.Time { return fixedNow } + return c +} + +// TestClientRequests covers RT-01 and T-11-10: exact paths and bodies, +// the apikey header, the Carbon +00:00 timestamp, [] for an empty payload, +// any 2xx as success (including an error body on publish), errors that +// never carry the key, and no request at all without a key. +func TestClientRequests(t *testing.T) { + f, srv := newFakeAPI(t) + c := testClient(srv.URL, testAPIKey) + ctx := context.Background() + + if err := c.Publish(ctx, "collection:5", "created.acme.widget", json.RawMessage(`{"z":1,"a":"x/y"}`)); err != nil { + t.Fatal(err) + } + if err := c.Publish(ctx, "collection:5", "pinged", nil); err != nil { + t.Fatal(err) + } + if err := c.Broadcast(ctx, []string{"a:1", "b:2"}, "bulk", json.RawMessage(`{"count":2}`)); err != nil { + t.Fatal(err) + } + if err := c.Broadcast(ctx, nil, "bulk", nil); err != nil { + t.Fatal(err) + } + if err := c.Unsubscribe(ctx, 7, "collection:5"); err != nil { + t.Fatal(err) + } + f.set(200, `{"result":{"presence":{"c1":{"user":"7"}}}}`) + pres, err := c.Presence(ctx, "presence:room:1") + if err != nil || pres["c1"] == nil { + t.Fatalf("presence = %v, %v", pres, err) + } + f.set(200, `{"result":{"nodes":[{"name":"n1"}]}}`) + info, err := c.Info(ctx) + if err != nil || info["nodes"] == nil { + t.Fatalf("info = %v, %v", info, err) + } + want := []apiCall{ + {"POST", "/api/publish", "apikey " + testAPIKey, "application/json", `{"channel":"collection:5","data":{"event":"created.acme.widget","payload":{"z":1,"a":"x/y"},"timestamp":"2026-09-30T12:00:00+00:00"}}`}, + {"POST", "/api/publish", "apikey " + testAPIKey, "application/json", `{"channel":"collection:5","data":{"event":"pinged","payload":[],"timestamp":"2026-09-30T12:00:00+00:00"}}`}, + {"POST", "/api/broadcast", "apikey " + testAPIKey, "application/json", `{"channels":["a:1","b:2"],"data":{"event":"bulk","payload":{"count":2},"timestamp":"2026-09-30T12:00:00+00:00"}}`}, + {"POST", "/api/unsubscribe", "apikey " + testAPIKey, "application/json", `{"user":"7","channel":"collection:5"}`}, + {"POST", "/api/presence", "apikey " + testAPIKey, "application/json", `{"channel":"presence:room:1"}`}, + {"POST", "/api/info", "apikey " + testAPIKey, "application/json", `{}`}, + } + got := f.take() + if len(got) != len(want) { + t.Fatalf("calls = %+v", got) + } + for i := range want { + if got[i] != want[i] { + t.Errorf("call %d = %+v\nwant %+v", i, got[i], want[i]) + } + } + + t.Run("2xx_error_body_is_success_for_publish", func(t *testing.T) { + f.set(200, `{"error":{"code":102,"message":"unknown channel"}}`) + if err := c.Publish(ctx, "x:1", "e", nil); err != nil { + t.Fatalf("publish with an error body = %v, want success like WinterCMS", err) + } + if _, err := c.Info(ctx); err == nil || !strings.Contains(err.Error(), "error 102: unknown channel") { + t.Fatalf("info with an error body = %v", err) + } + f.set(200, `not json`) + if _, err := c.Info(ctx); err == nil { + t.Fatal("unreadable info answer accepted") + } + if p, err := c.Presence(ctx, "x"); err == nil || p == nil || len(p) != 0 { + t.Fatalf("unreadable presence = %v, %v", p, err) + } + f.set(202, `{}`) + if err := c.Unsubscribe(ctx, 1, "x"); err != nil { + t.Fatalf("202 = %v", err) + } + if p, err := c.Presence(ctx, "x"); err != nil || p == nil { + t.Fatalf("presence without result = %v, %v", p, err) + } + f.take() + }) + + t.Run("non_2xx_is_an_error_without_the_key", func(t *testing.T) { + f.set(500, `{"secret":"`+testAPIKey+`"}`) + for name, call := range map[string]func() error{ + "publish": func() error { return c.Publish(ctx, "x:1", "e", nil) }, + "broadcast": func() error { return c.Broadcast(ctx, []string{"x:1"}, "e", nil) }, + "unsubscribe": func() error { return c.Unsubscribe(ctx, 1, "x:1") }, + "presence": func() error { _, err := c.Presence(ctx, "x:1"); return err }, + "info": func() error { _, err := c.Info(ctx); return err }, + } { + err := call() + if err == nil || !strings.Contains(err.Error(), "HTTP 500") || strings.Contains(err.Error(), testAPIKey) { + t.Errorf("%s: err = %v", name, err) + } + } + f.take() + }) + + t.Run("unreachable_server", func(t *testing.T) { + dead := httptest.NewServer(http.NotFoundHandler()) + url := dead.URL + dead.Close() + err := testClient(url, testAPIKey).Publish(ctx, "x:1", "e", nil) + if err == nil || strings.Contains(err.Error(), testAPIKey) { + t.Fatalf("err = %v", err) + } + if err := testClient("http://[::1", testAPIKey).Publish(ctx, "x:1", "e", nil); err == nil { + t.Fatal("malformed URL accepted") + } + if err := testClient(srv.URL, testAPIKey).Publish(nil, "x:1", "e", json.RawMessage(`{"bad"`)); err == nil { + t.Fatal("invalid payload JSON accepted") + } + }) + + t.Run("empty_key_sends_nothing", func(t *testing.T) { + off := testClient(srv.URL, "") + f.take() + if !errors.Is(off.Publish(ctx, "x", "e", nil), ErrNotConfigured) || + !errors.Is(off.Broadcast(ctx, []string{"x"}, "e", nil), ErrNotConfigured) || + !errors.Is(off.Unsubscribe(ctx, 1, "x"), ErrNotConfigured) { + t.Fatal("publishing without a key did not return ErrNotConfigured") + } + if p, err := off.Presence(ctx, "x"); !errors.Is(err, ErrNotConfigured) || p == nil { + t.Fatalf("presence = %v, %v", p, err) + } + if i, err := off.Info(ctx); !errors.Is(err, ErrNotConfigured) || i == nil { + t.Fatalf("info = %v, %v", i, err) + } + if n := len(f.take()); n != 0 { + t.Fatalf("%d requests sent without an API key", n) + } + if off.Enabled() || off.DebugInfo().APIKeySet || !c.DebugInfo().Enabled { + t.Fatal("Enabled/DebugInfo") + } + var nilClient *Client + if nilClient.Enabled() || nilClient.DebugInfo() != (DebugInfo{}) { + t.Fatal("nil client") + } + raw, _ := json.Marshal(c.DebugInfo()) + if strings.Contains(string(raw), testAPIKey) { + t.Fatal("DebugInfo carries the key") + } + }) +} + +// TestClientLoadConfig covers realtime.centrifugo.* parsing and the +// driver built from it. +func TestClientLoadConfig(t *testing.T) { + def := LoadConfig(nil) + if def.APIURL != DefaultAPIURL || def.TokenTTL != DefaultTokenTTL || def.WSURL != DefaultWSURL || def.TokenPath != DefaultTokenPath || def.SubscribePath != DefaultSubscribePath { + t.Fatalf("defaults = %+v", def) + } + cfg, err := compass.Open(compass.Options{Dir: t.TempDir(), Env: "testing", Environ: []string{}}) + if err != nil { + t.Fatal(err) + } + for k, v := range map[string]any{ + "realtime.driver": "centrifugo", + "realtime.centrifugo.api_url": "http://127.0.0.1:1/api/", + "realtime.centrifugo.api_key": " k ", + "realtime.centrifugo.token_secret": "s", + "realtime.centrifugo.proxy_secret": "p", + "realtime.centrifugo.token_ttl": "2m", + "realtime.centrifugo.ws_url": "wss://rt.example.test/ws", + "realtime.centrifugo.token_path": "/rt/token", + "realtime.centrifugo.subscribe_path": "/rt/sub", + "http.trusted_proxies": []any{"10.0.0.0/8"}, + } { + if err := cfg.Set(k, v); err != nil { + t.Fatal(err) + } + } + got := LoadConfig(cfg) + if got.APIURL != "http://127.0.0.1:1/api" || got.APIKey != "k" || got.TokenSecret != "s" || got.ProxySecret != "p" || + got.TokenTTL != 2*time.Minute || got.WSURL != "wss://rt.example.test/ws" || got.TokenPath != "/rt/token" || got.SubscribePath != "/rt/sub" || len(got.TrustedProxies) != 1 { + t.Fatalf("config = %+v", got) + } + svc, err := lighthouse.From(backpack.New(cfg)) + if err != nil { + t.Fatal(err) + } + d, ok := svc.Driver().(*Driver) + if !ok || d.Name() != DriverName || !d.Enabled() || d.Config().TokenPath != "/rt/token" || d.Client() == nil || !d.Issuer().Configured() { + t.Fatalf("driver = %#v", svc.Driver()) + } + routes := d.Routes() + if len(routes) != 2 || routes[0].Path != "/rt/token" || routes[0].Surface != lighthouse.UserAuth || routes[1].Path != "/rt/sub" || routes[1].Surface != lighthouse.ServerToServer { + t.Fatalf("routes = %+v", routes) + } + // Publish and Broadcast go through the client; an unreachable API is an + // error, not a panic. + if err := d.Publish(context.Background(), "a:1", "e", nil); err == nil { + t.Fatal("publish to an unreachable API succeeded") + } + if err := d.Broadcast(context.Background(), []string{"a:1"}, "e", nil); err == nil { + t.Fatal("broadcast to an unreachable API succeeded") + } +} diff --git a/modules/lighthouse/centrifugo/proxy_test.go b/modules/lighthouse/centrifugo/proxy_test.go new file mode 100644 index 0000000..1ffe43b --- /dev/null +++ b/modules/lighthouse/centrifugo/proxy_test.go @@ -0,0 +1,222 @@ +package centrifugo + +import ( + "bytes" + "context" + "log/slog" + "net/http" + "net/http/httptest" + "strconv" + "strings" + "sync" + "testing" + + "git.golem15.com/golem15/summercms/modules/backpack" + "git.golem15.com/golem15/summercms/modules/lighthouse" +) + +const ( + proxyTestSecret = "proxy-secret-test-only-4b1d" + wrongSecret = "not-the-proxy-secret-9c2e" + denyBody = `{"error":{"code":403,"message":"Access denied"}}` + allowEmptyInfo = `{"result":{"info":[]}}` +) + +type lockedBuffer struct { + mu sync.Mutex + buf bytes.Buffer +} + +func (b *lockedBuffer) Write(p []byte) (int, error) { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.Write(p) +} + +func (b *lockedBuffer) String() string { + b.mu.Lock() + defer b.mu.Unlock() + return b.buf.String() +} + +// acmeAuthorizer allows user 7 on acme:room:1 and records every call. +type acmeAuthorizer struct { + mu sync.Mutex + calls []string +} + +func (a *acmeAuthorizer) Authorize(ctx context.Context, userID uint, channel string) lighthouse.Result { + a.mu.Lock() + a.calls = append(a.calls, strings.Join([]string{uintString(userID), channel, lighthouse.ClientID(ctx)}, "|")) + a.mu.Unlock() + switch { + case userID == 7 && (channel == "acme:room:1" || channel == "presence:acme:room:1"): + return lighthouse.Allowed(nil) + case userID == 7 && channel == "acme:room:info": + return lighthouse.Allowed(map[string]any{"role": "owner"}) + case userID == 7 && channel == "acme:room:bad-info": + return lighthouse.Allowed(map[string]any{"bad": make(chan int)}) + case userID == 7 && channel == "presence:acme:room:caps": + r := lighthouse.Allowed(nil) + r.Capabilities = []string{"prs", "sub"} + r.Overrides = map[string]any{"join_leave": map[string]bool{"value": false}, "zeta": 1, "alpha": "a"} + return r + case channel == "acme:room:silent": + return lighthouse.Result{} + } + return lighthouse.Denied("not a member of " + channel) +} + +func (a *acmeAuthorizer) last() string { + a.mu.Lock() + defer a.mu.Unlock() + if len(a.calls) == 0 { + return "" + } + return a.calls[len(a.calls)-1] +} + +func uintString(v uint) string { return strconv.FormatUint(uint64(v), 10) } + +func proxyService(t *testing.T) (*lighthouse.Service, *acmeAuthorizer, *lockedBuffer) { + t.Helper() + app := backpack.New(nil) + logs := &lockedBuffer{} + if err := app.Publish(slog.New(slog.NewJSONHandler(logs, nil))); err != nil { + t.Fatal(err) + } + svc, err := lighthouse.From(app) + if err != nil { + t.Fatal(err) + } + auth := &acmeAuthorizer{} + if err := svc.Registry().Register("acme", auth); err != nil { + t.Fatal(err) + } + return svc, auth, logs +} + +func proxyCall(h http.HandlerFunc, secret *string, body string) *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodPost, "/api/realtime/subscribe", strings.NewReader(body)) + req.RemoteAddr = "203.0.113.9:4242" + if secret != nil { + req.Header.Set("X-Centrifugo-Secret", *secret) + } + rec := httptest.NewRecorder() + h(rec, req) + return rec +} + +func strp(s string) *string { return &s } + +// TestProxy covers RT-02, T-11-01 and T-11-02, porting the WinterCMS +// websockets security tests (WS-005, WS-007, WS-013): the proxy secret is +// compared in constant time and an empty configured secret denies +// everything; empty, zero and non-scalar users deny; channels are parsed +// with the presence and segment rules and routed byte-exactly to their +// namespace authorizer with the PHP (int) user id and the client id; +// allows answer the exact info, allow and override bytes; every deny is +// the same HTTP 200 body with the reason only in the logs, and no secret +// is ever logged. +func TestProxy(t *testing.T) { + svc, auth, logs := proxyService(t) + h := ProxyHandler(svc, Config{ProxySecret: proxyTestSecret}) + good := strp(proxyTestSecret) + cases := []struct { + name string + secret *string + body string + want string + reason string + authCall string + }{ + {"missing_secret", nil, `{"user":"7","channel":"acme:room:1"}`, denyBody, "Invalid or missing proxy secret", ""}, + {"wrong_secret", strp(wrongSecret), `{"user":"7","channel":"acme:room:1"}`, denyBody, "Invalid or missing proxy secret", ""}, + {"secret_prefix", strp(proxyTestSecret[:10]), `{"user":"7","channel":"acme:room:1"}`, denyBody, "Invalid or missing proxy secret", ""}, + {"malformed_json", good, `{"user":`, denyBody, "Malformed proxy request", ""}, + {"empty_user", good, `{"user":"","channel":"acme:room:1"}`, denyBody, "Authentication required", ""}, + {"zero_user_string", good, `{"user":"0","channel":"acme:room:1"}`, denyBody, "Authentication required", ""}, + {"zero_user_number", good, `{"user":0,"channel":"acme:room:1"}`, denyBody, "Authentication required", ""}, + {"zero_user_float", good, `{"user":0.0,"channel":"acme:room:1"}`, denyBody, "Authentication required", ""}, + {"bool_user", good, `{"user":true,"channel":"acme:room:1"}`, denyBody, "Authentication required", ""}, + {"object_user", good, `{"user":{"id":7},"channel":"acme:room:1"}`, denyBody, "Authentication required", ""}, + {"missing_user", good, `{"channel":"acme:room:1"}`, denyBody, "Authentication required", ""}, + {"missing_channel", good, `{"user":"7"}`, denyBody, "Missing channel", ""}, + {"empty_channel", good, `{"user":"7","channel":""}`, denyBody, "Missing channel", ""}, + {"double_presence_ws005", good, `{"user":"7","channel":"presence:presence:acme:room:1"}`, denyBody, "Unknown channel namespace", ""}, + {"four_segments_ws005", good, `{"user":"7","channel":"acme:room:1:extra"}`, denyBody, "Unknown channel namespace", ""}, + {"leading_colon_ws005", good, `{"user":"7","channel":":acme:room"}`, denyBody, "Unknown channel namespace", ""}, + {"unknown_namespace", good, `{"user":"7","channel":"other:1"}`, denyBody, "Unknown channel namespace", ""}, + {"case_mismatched_namespace", good, `{"user":"7","channel":"ACME:room:1"}`, denyBody, "Unknown channel namespace", ""}, + {"allow_string_user", good, `{"user":"7","channel":"acme:room:1","client":"c-1"}`, allowEmptyInfo, "", "7|acme:room:1|c-1"}, + {"allow_number_user", good, `{"user":7,"channel":"acme:room:1"}`, allowEmptyInfo, "", "7|acme:room:1|"}, + {"php_int_cast_user", good, `{"user":"7abc","channel":"acme:room:1"}`, allowEmptyInfo, "", "7|acme:room:1|"}, + {"negative_user_is_zero", good, `{"user":"-5","channel":"acme:room:1"}`, denyBody, "not a member of acme:room:1", "0|acme:room:1|"}, + {"authorizer_deny", good, `{"user":"8","channel":"acme:room:1"}`, denyBody, "not a member of acme:room:1", "8|acme:room:1|"}, + {"authorizer_deny_without_reason", good, `{"user":"8","channel":"acme:room:silent"}`, denyBody, "Access denied", "8|acme:room:silent|"}, + {"allow_with_info", good, `{"user":"7","channel":"acme:room:info"}`, `{"result":{"info":{"role":"owner"}}}`, "", ""}, + {"unencodable_info_denies", good, `{"user":"7","channel":"acme:room:bad-info"}`, denyBody, "", ""}, + {"presence_defaults_ws013", good, `{"user":"7","channel":"presence:acme:room:1"}`, + `{"result":{"info":[],"allow":["prs"],"override":{"presence":{"value":true},"join_leave":{"value":true},"force_push_join_leave":{"value":false}}}}`, "", "7|presence:acme:room:1|"}, + {"presence_override_merge", good, `{"user":"7","channel":"presence:acme:room:caps"}`, + `{"result":{"info":[],"allow":["prs","sub"],"override":{"presence":{"value":true},"join_leave":{"value":false},"force_push_join_leave":{"value":false},"alpha":"a","zeta":1}}}`, "", ""}, + {"oversized_body", good, `{"user":"7","channel":"acme:room:1","pad":"` + strings.Repeat("x", 64<<10) + `"}`, denyBody, "Malformed proxy request", ""}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + before := auth.last() + rec := proxyCall(h, c.secret, c.body) + if rec.Code != http.StatusOK { + t.Fatalf("status = %d, want 200 (Centrifugo reads non-200 as an internal error)", rec.Code) + } + if got := rec.Body.String(); got != c.want { + t.Fatalf("body = %s\nwant %s", got, c.want) + } + if rec.Header().Get("Content-Type") != "application/json" || rec.Header().Get("Cache-Control") != "no-cache, private" { + t.Fatalf("headers = %v", rec.Header()) + } + if c.reason != "" && !strings.Contains(logs.String(), `"reason":"`+c.reason+`"`) { + t.Fatalf("no deny log with reason %q:\n%s", c.reason, logs.String()) + } + if c.authCall != "" && auth.last() != c.authCall { + t.Fatalf("authorizer call = %q, want %q", auth.last(), c.authCall) + } + if c.authCall == "" && c.want == denyBody && c.reason != "" && !strings.HasPrefix(c.reason, "not a member") && c.reason != "Access denied" && auth.last() != before { + t.Fatalf("authorizer was consulted for a request refused before it: %q", auth.last()) + } + }) + } + if out := logs.String(); strings.Contains(out, proxyTestSecret) || strings.Contains(out, wrongSecret) || strings.Contains(out, proxyTestSecret[:10]) { + t.Fatalf("a secret reached the logs:\n%s", out) + } + if !strings.Contains(logs.String(), `"ip":"203.0.113.9"`) { + t.Fatal("secret failures do not log the client IP") + } + + t.Run("empty_configured_secret_denies_everything", func(t *testing.T) { + off := ProxyHandler(svc, Config{}) + for _, secret := range []*string{nil, strp(""), strp(proxyTestSecret)} { + if rec := proxyCall(off, secret, `{"user":"7","channel":"acme:room:1"}`); rec.Body.String() != denyBody { + t.Fatalf("secret %v: body %s", secret, rec.Body.String()) + } + } + }) + + t.Run("concurrent_subscribes", func(t *testing.T) { + var wg sync.WaitGroup + for i := range 16 { + wg.Add(1) + go func() { + defer wg.Done() + body, want := `{"user":"7","channel":"acme:room:1"}`, allowEmptyInfo + if i%2 == 1 { + body, want = `{"user":"8","channel":"acme:room:1"}`, denyBody + } + if rec := proxyCall(h, good, body); rec.Body.String() != want { + t.Errorf("concurrent %d: %s", i, rec.Body.String()) + } + }() + } + wg.Wait() + }) +} diff --git a/modules/lighthouse/centrifugo/token_test.go b/modules/lighthouse/centrifugo/token_test.go new file mode 100644 index 0000000..df7d880 --- /dev/null +++ b/modules/lighthouse/centrifugo/token_test.go @@ -0,0 +1,210 @@ +package centrifugo + +import ( + "context" + "encoding/base64" + "errors" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + "git.golem15.com/golem15/summercms/modules/backpack" + "git.golem15.com/golem15/summercms/modules/bouncer" + "git.golem15.com/golem15/summercms/modules/lighthouse" + "github.com/golang-jwt/jwt/v5" +) + +const tokenTestSecret = "test-only-centrifugo-token-secret-0123456789" + +var fixedNow = time.Date(2026, 9, 30, 12, 0, 0, 0, time.UTC) + +// jwtParts returns the decoded header and claims segments of token after +// checking its HS256 signature with secret. +func jwtParts(t *testing.T, token, secret string) (string, string) { + t.Helper() + parsed, err := jwt.Parse(token, func(tok *jwt.Token) (any, error) { return []byte(secret), nil }, + jwt.WithValidMethods([]string{"HS256"}), jwt.WithoutClaimsValidation()) + if err != nil || !parsed.Valid { + t.Fatalf("token does not verify: %v", err) + } + seg := strings.Split(token, ".") + if len(seg) != 3 { + t.Fatalf("token has %d segments", len(seg)) + } + dec := func(s string) string { + b, err := base64.RawURLEncoding.DecodeString(s) + if err != nil { + t.Fatal(err) + } + return string(b) + } + return dec(seg[0]), dec(seg[1]) +} + +// TestTokenClaims covers RT-01 and T-11-11: every generator signs HS256 with +// the exact WinterCMS claim set and order, the user token carries only the +// name, and an empty secret signs nothing. +func TestTokenClaims(t *testing.T) { + iss := NewTokenIssuer(tokenTestSecret, time.Hour) + iss.Now = func() time.Time { return fixedNow } + exp := fixedNow.Add(time.Hour).Unix() + name := "Ann" + cases := []struct { + name string + sign func() (string, error) + want string + }{ + {"for_user", func() (string, error) { return iss.ForUser(lighthouse.User{ID: 7, Name: &name}) }, + fmt.Sprintf(`{"sub":"7","exp":%d,"info":{"name":"Ann"}}`, exp)}, + {"for_user_null_name", func() (string, error) { return iss.ForUser(lighthouse.User{ID: 8}) }, + fmt.Sprintf(`{"sub":"8","exp":%d,"info":{"name":null}}`, exp)}, + {"subscription", func() (string, error) { return iss.Subscription(lighthouse.User{ID: 7}, "collection:5") }, + fmt.Sprintf(`{"sub":"7","channel":"collection:5","exp":%d}`, exp)}, + {"anonymous", iss.Anonymous, + fmt.Sprintf(`{"sub":"","exp":%d}`, fixedNow.Add(5*time.Minute).Unix())}, + {"for_identifier_empty_info", func() (string, error) { return iss.ForIdentifier("kiosk-1", nil) }, + fmt.Sprintf(`{"sub":"kiosk-1","exp":%d,"info":[]}`, exp)}, + {"for_identifier_info", func() (string, error) { return iss.ForIdentifier("kiosk-1", map[string]any{"room": "a/b"}) }, + fmt.Sprintf(`{"sub":"kiosk-1","exp":%d,"info":{"room":"a/b"}}`, exp)}, + {"subscription_for_identifier", func() (string, error) { return iss.SubscriptionForIdentifier("kiosk-1", "room:1") }, + fmt.Sprintf(`{"sub":"kiosk-1","channel":"room:1","exp":%d}`, exp)}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + tok, err := c.sign() + if err != nil { + t.Fatal(err) + } + header, claims := jwtParts(t, tok, tokenTestSecret) + if header != `{"alg":"HS256","typ":"JWT"}` { + t.Fatalf("header = %s", header) + } + if claims != c.want { + t.Fatalf("claims = %s, want %s", claims, c.want) + } + }) + } + + empty := NewTokenIssuer("", 0) + if empty.Configured() || empty.ttl != DefaultTokenTTL { + t.Fatalf("empty issuer configured=%v ttl=%s", empty.Configured(), empty.ttl) + } + for name, sign := range map[string]func() (string, error){ + "ForUser": func() (string, error) { return empty.ForUser(lighthouse.User{ID: 1}) }, + "Subscription": func() (string, error) { return empty.Subscription(lighthouse.User{ID: 1}, "a:1") }, + "Anonymous": empty.Anonymous, + "ForIdentifier": func() (string, error) { return empty.ForIdentifier("x", nil) }, + "SubscriptionForIdentifier": func() (string, error) { return empty.SubscriptionForIdentifier("x", "a:1") }, + } { + if tok, err := sign(); !errors.Is(err, ErrNotConfigured) || tok != "" { + t.Errorf("%s with an empty secret = %q, %v; want ErrNotConfigured", name, tok, err) + } + } + var nilIss *TokenIssuer + if nilIss.Configured() { + t.Fatal("nil issuer is configured") + } + if _, err := iss.ForIdentifier("x", map[string]any{"bad": make(chan int)}); err == nil { + t.Fatal("unencodable info accepted") + } + // The default clock is used when Now is nil. + live := NewTokenIssuer(tokenTestSecret, time.Minute) + tok, err := live.Anonymous() + if err != nil { + t.Fatal(err) + } + if _, claims := jwtParts(t, tok, tokenTestSecret); !strings.Contains(claims, `"exp":`) { + t.Fatalf("claims = %s", claims) + } +} + +func tokenService(t *testing.T, lookup lighthouse.UserLookup) *lighthouse.Service { + t.Helper() + svc, err := lighthouse.From(backpack.New(nil)) + if err != nil { + t.Fatal(err) + } + svc.SetUserLookup(lookup) + return svc +} + +// TestTokenHandler covers RT-01: 401 without a principal or user, 503 with +// an empty secret only after the user check, and a 200 {"token"} body with +// the Laravel JSON headers and no trailing newline, safe under concurrency. +func TestTokenHandler(t *testing.T) { + name := "Ann" + svc := tokenService(t, func(_ context.Context, id uint) (lighthouse.User, bool, error) { + switch id { + case 7: + return lighthouse.User{ID: 7, Name: &name}, true, nil + case 9: + return lighthouse.User{}, false, errors.New("database down") + } + return lighthouse.User{}, false, nil + }) + iss := NewTokenIssuer(tokenTestSecret, time.Hour) + h := TokenHandler(svc, iss) + noSecret := TokenHandler(svc, NewTokenIssuer("", time.Hour)) + call := func(h http.HandlerFunc, p *bouncer.Principal) *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodGet, "/api/realtime/token", nil) + if p != nil { + req = req.WithContext(bouncer.WithUser(req.Context(), p)) + } + rec := httptest.NewRecorder() + h(rec, req) + return rec + } + cases := []struct { + name string + h http.HandlerFunc + p *bouncer.Principal + status int + body string + }{ + {"no_principal", h, nil, 401, `{"error":"Unauthorized"}`}, + {"zero_id", h, &bouncer.Principal{ID: 0}, 401, `{"error":"Unauthorized"}`}, + {"unknown_user", h, &bouncer.Principal{ID: 5}, 401, `{"error":"Unauthorized"}`}, + {"lookup_error", h, &bouncer.Principal{ID: 9}, 401, `{"error":"Unauthorized"}`}, + {"unknown_user_before_secret", noSecret, &bouncer.Principal{ID: 5}, 401, `{"error":"Unauthorized"}`}, + {"empty_secret", noSecret, &bouncer.Principal{ID: 7}, 503, `{"error":"WebSocket not configured"}`}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + rec := call(c.h, c.p) + if rec.Code != c.status || rec.Body.String() != c.body { + t.Fatalf("got %d %s, want %d %s", rec.Code, rec.Body.String(), c.status, c.body) + } + if rec.Header().Get("Content-Type") != "application/json" || rec.Header().Get("Cache-Control") != "no-cache, private" { + t.Fatalf("headers = %v", rec.Header()) + } + }) + } + t.Run("ok", func(t *testing.T) { + rec := call(h, &bouncer.Principal{ID: 7}) + body := rec.Body.String() + if rec.Code != 200 || !strings.HasPrefix(body, `{"token":"`) || !strings.HasSuffix(body, `"}`) || strings.HasSuffix(body, "\n") { + t.Fatalf("got %d %q", rec.Code, body) + } + tok := strings.TrimSuffix(strings.TrimPrefix(body, `{"token":"`), `"}`) + if _, claims := jwtParts(t, tok, tokenTestSecret); !strings.HasPrefix(claims, `{"sub":"7","exp":`) || !strings.HasSuffix(claims, `,"info":{"name":"Ann"}}`) { + t.Fatalf("claims = %s", claims) + } + }) + t.Run("concurrent", func(t *testing.T) { + var wg sync.WaitGroup + for range 16 { + wg.Add(1) + go func() { + defer wg.Done() + if rec := call(h, &bouncer.Principal{ID: 7}); rec.Code != 200 { + t.Errorf("concurrent status %d", rec.Code) + } + }() + } + wg.Wait() + }) +} diff --git a/modules/lighthouse/channel_test.go b/modules/lighthouse/channel_test.go index 48533fc..08256d4 100644 --- a/modules/lighthouse/channel_test.go +++ b/modules/lighthouse/channel_test.go @@ -76,9 +76,22 @@ func TestFormatChannels(t *testing.T) { if got := FormatChannels("", []string{"Collection:5"}); !reflect.DeepEqual(got, []string{"collection:5"}) { t.Fatalf("no namespace = %v", got) } + // PHP treats the namespace "0" as empty; the prefix is applied once. + if got := FormatChannels("0", []string{"Room:1"}); !reflect.DeepEqual(got, []string{"room:1"}) { + t.Fatalf("namespace 0 = %v", got) + } + if got := FormatChannels("acme", []string{"ACME:room:1", "acme:acme:x"}); !reflect.DeepEqual(got, []string{"acme:room:1", "acme:acme:x"}) { + t.Fatalf("prefix once = %v", got) + } + if got := FormatChannels("acme", nil); len(got) != 0 { + t.Fatalf("nil channels = %v", got) + } } func TestClientID(t *testing.T) { + if ClientID(nil) != "" || ClientID(WithClientID(nil, "c-0")) != "c-0" { + t.Fatal("nil ctx handling") + } if ClientID(context.Background()) != "" { t.Fatal("empty ctx has a client id") } @@ -86,35 +99,3 @@ func TestClientID(t *testing.T) { t.Fatalf("ClientID = %q", got) } } - -func TestRegistry(t *testing.T) { - r := NewRegistry() - allow := AuthorizerFunc(func(context.Context, uint, string) Result { return Allowed(nil) }) - if err := r.Register("b", allow); err != nil { - t.Fatal(err) - } - if err := r.Register("a", allow); err != nil { - t.Fatal(err) - } - for _, bad := range []struct { - ns string - a Authorizer - }{{"", allow}, {"x:y", allow}, {"c", nil}, {"a", allow}} { - if err := r.Register(bad.ns, bad.a); err == nil { - t.Errorf("Register(%q) accepted", bad.ns) - } - } - if got := r.Namespaces(); !reflect.DeepEqual(got, []string{"a", "b"}) { - t.Fatalf("Namespaces = %v", got) - } - if _, ok := r.Get("A"); ok { - t.Fatal("lookup is not case-sensitive") - } - if a, ok := r.Get("a"); !ok || !a.Authorize(context.Background(), 1, "a:1").Allowed { - t.Fatal("Get(a) failed") - } - d := Denied("why") - if d.Allowed || d.Reason() != "why" { - t.Fatalf("Denied = %+v", d) - } -} diff --git a/modules/lighthouse/lighthouse_test.go b/modules/lighthouse/lighthouse_test.go new file mode 100644 index 0000000..44bf7a8 --- /dev/null +++ b/modules/lighthouse/lighthouse_test.go @@ -0,0 +1,248 @@ +package lighthouse + +import ( + "bytes" + "context" + "encoding/json" + "log/slog" + "reflect" + "strings" + "testing" + "time" + + "git.golem15.com/golem15/summercms/modules/backpack" + "git.golem15.com/golem15/summercms/modules/compass" + "gorm.io/gorm" +) + +func configApp(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) +} + +// TestFromSelectsDriver covers D-11: realtime.driver picks a registered +// driver (null by default, case-insensitive), an unknown name fails and +// lists the registered drivers, and the realtime.* settings are read once. +func TestFromSelectsDriver(t *testing.T) { + cases := []struct { + name string + driver any + want string + }{ + {"default_null", nil, "null"}, + {"memory", "memory", "memory"}, + {"log_case_insensitive", " LOG ", "log"}, + } + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + kv := map[string]any{} + if c.driver != nil { + kv["realtime.driver"] = c.driver + } + svc, err := From(configApp(t, kv)) + if err != nil { + t.Fatal(err) + } + if got := svc.Driver().Name(); got != c.want { + t.Fatalf("driver = %s, want %s", got, c.want) + } + if svc.Driver().Routes() != nil { + t.Fatalf("%s driver declares routes", c.want) + } + }) + } + t.Run("unknown_lists_registered", func(t *testing.T) { + _, err := From(configApp(t, map[string]any{"realtime.driver": "nope"})) + if err == nil || !strings.Contains(err.Error(), `unknown realtime.driver "nope"`) || !strings.Contains(err.Error(), "acme-nil-driver, log, memory, null") { + t.Fatalf("err = %v", err) + } + }) + t.Run("settings_and_idempotence", func(t *testing.T) { + app := configApp(t, map[string]any{ + "realtime.broadcast_namespace": " acme ", + "realtime.broadcast_queue": "rt", + "realtime.broadcast_timeout": 7, + }) + svc, err := From(app) + if err != nil { + t.Fatal(err) + } + if svc.Namespace() != "acme" || svc.Queue() != "rt" || svc.Timeout() != 7*time.Second || svc.Registry() == nil { + t.Fatalf("settings = %q %q %s", svc.Namespace(), svc.Queue(), svc.Timeout()) + } + again, err := From(app) + if err != nil || again != svc { + t.Fatalf("From twice = %p, %v", again, err) + } + if _, err := From(nil); err == nil { + t.Fatal("From(nil) succeeded") + } + def, err := From(backpack.New(nil)) + if err != nil || def.Queue() != DefaultQueue || def.Timeout() != DefaultTimeout || def.Namespace() != "" { + t.Fatalf("config-less service = %+v, %v", def, err) + } + }) + t.Run("nil_service_accessors", func(t *testing.T) { + var s *Service + if s.Driver() != nil || s.Registry() != nil || s.Namespace() != "" || s.Queue() != DefaultQueue || s.Timeout() != DefaultTimeout || s.Logger() == nil { + t.Fatal("nil service accessors") + } + s.SetUserLookup(nil) + if u, found, err := s.User(context.Background(), 4); err != nil || !found || u.ID != 4 { + t.Fatalf("nil service User = %+v %v %v", u, found, err) + } + }) +} + +// TestDurationSetting covers the seconds-or-duration setting parser. +func TestDurationSetting(t *testing.T) { + app := configApp(t, map[string]any{"a": 5, "b": "250ms", "c": "soon", "d": -3, "e": " "}) + for key, want := range map[string]time.Duration{"a": 5 * time.Second, "b": 250 * time.Millisecond, "c": 0, "d": 0, "e": 0, "missing": 0} { + if got := DurationSetting(app.Config, key); got != want { + t.Errorf("DurationSetting(%s) = %s, want %s", key, got, want) + } + } + if DurationSetting(nil, "a") != 0 { + t.Fatal("nil config") + } +} + +// TestDrivers covers the built-in drivers: the log driver logs channels and +// event but never the payload, the memory driver records copies, and the +// registry panics on misuse. +func TestDrivers(t *testing.T) { + var buf bytes.Buffer + d := &logDriver{log: slog.New(slog.NewTextHandler(&buf, nil))} + ctx := context.Background() + if err := d.Publish(ctx, "room:1", "acme.ping", json.RawMessage(`{"secret":"payload-value"}`)); err != nil { + t.Fatal(err) + } + if err := d.Broadcast(ctx, []string{"room:1", "room:2"}, "acme.ping", json.RawMessage(`{"secret":"payload-value"}`)); err != nil { + t.Fatal(err) + } + out := buf.String() + if !strings.Contains(out, "channel=room:1") || !strings.Contains(out, "event=acme.ping") || strings.Contains(out, "payload-value") { + t.Fatalf("log driver output = %s", out) + } + if (nullDriver{}).Publish(ctx, "", "", nil) != nil || (nullDriver{}).Broadcast(ctx, nil, "", nil) != nil { + t.Fatal("null driver returned an error") + } + + m := NewMemoryDriver() + payload := json.RawMessage(`{"a":1}`) + channels := []string{"x:1", "x:2"} + if err := m.Broadcast(ctx, channels, "e", payload); err != nil { + t.Fatal(err) + } + channels[0] = "mutated" + payload[2] = 'Z' + pubs := m.Publications() + if len(pubs) != 1 || pubs[0].Method != "broadcast" || pubs[0].Channels[0] != "x:1" || string(pubs[0].Payload) != `{"a":1}` || m.Name() != "memory" || m.Routes() != nil { + t.Fatalf("memory driver kept references: %+v", pubs) + } + pubs[0].Channels[0] = "changed" + if m.Publications()[0].Channels[0] != "x:1" { + t.Fatal("Publications returned shared slices") + } + var nilMem *MemoryDriver + if nilMem.Publish(ctx, "a", "b", nil) == nil || nilMem.Publications() != nil { + t.Fatal("nil memory driver") + } + + for name, fn := range map[string]func(){ + "empty_name": func() { RegisterDriver("", func(*backpack.App, *Service) (Driver, error) { return nullDriver{}, nil }) }, + "nil_factory": func() { RegisterDriver("acme-nil", nil) }, + "duplicate": func() { + RegisterDriver("null", func(*backpack.App, *Service) (Driver, error) { return nullDriver{}, nil }) + }, + } { + func() { + defer func() { + if recover() == nil { + t.Errorf("RegisterDriver %s did not panic", name) + } + }() + fn() + }() + } + if _, err := From(configApp(t, map[string]any{"realtime.driver": "acme-nil-driver"})); err == nil || !strings.Contains(err.Error(), "returned nil") { + t.Fatalf("nil driver: %v", err) + } + if _, err := From(configApp(t, map[string]any{"realtime.driver": "acme-broken"})); err == nil || !strings.Contains(err.Error(), "driver acme-broken") { + t.Fatalf("failing factory: %v", err) + } +} + +func init() { + RegisterDriver("acme-nil-driver", func(*backpack.App, *Service) (Driver, error) { return nil, nil }) + RegisterDriver("acme-broken", func(*backpack.App, *Service) (Driver, error) { return nil, context.Canceled }) +} + +// TestBroadcastArgsJSON covers the stored job args: the payload travels as +// a JSON string so JSONB cannot reorder its keys. +func TestBroadcastArgsJSON(t *testing.T) { + a := BroadcastArgs{Channels: []string{"room:1"}, Event: "e", Payload: json.RawMessage(`{"z":1,"a":2}`)} + raw, err := json.Marshal(a) + if err != nil { + t.Fatal(err) + } + if string(raw) != `{"channels":["room:1"],"event":"e","payload":"{\"z\":1,\"a\":2}"}` { + t.Fatalf("stored args = %s", raw) + } + var back BroadcastArgs + if err := json.Unmarshal(raw, &back); err != nil { + t.Fatal(err) + } + if string(back.Payload) != `{"z":1,"a":2}` || back.Event != "e" || back.Kind() != "summer.broadcast" { + t.Fatalf("round trip = %+v", back) + } + if err := json.Unmarshal([]byte(`{"channels":[],"event":"e","payload":""}`), &back); err != nil || back.Payload != nil { + t.Fatalf("empty payload = %s, %v", back.Payload, err) + } + if err := back.UnmarshalJSON([]byte(`not json`)); err == nil { + t.Fatal("malformed args accepted") + } +} + +// TestBind covers the Binding rules. +func TestBind(t *testing.T) { + svc, err := From(backpack.New(nil)) + if err != nil { + t.Fatal(err) + } + channels := func(context.Context, *gorm.DB, *Gadget) ([]string, error) { return nil, nil } + if err := Bind[Gadget](nil, Binding[Gadget]{Channels: channels}); err == nil { + t.Fatal("nil service accepted") + } + if err := Bind[string](svc, Binding[string]{Channels: func(context.Context, *gorm.DB, *string) ([]string, error) { return nil, nil }}); err == nil { + t.Fatal("non-struct type accepted") + } + if err := Bind[Gadget](svc, Binding[Gadget]{}); err == nil { + t.Fatal("nil Channels accepted") + } + if err := Bind[Gadget](svc, Binding[Gadget]{Channels: channels}); err != nil { + t.Fatal(err) + } + if err := Bind[Gadget](svc, Binding[Gadget]{Channels: channels}); err == nil { + t.Fatal("second binding accepted") + } + h := svc.handlerFor(reflect.TypeFor[Gadget]()) + if h == nil || h.alias != "lighthouse.gadget" || h.ttl != DefaultTTL || h.eventName(ActionDeleted) != "deleted.lighthouse.gadget" || !h.allows(nil, ActionUpdated) { + t.Fatalf("binding defaults = %+v", h) + } + if svc.handlerFor(reflect.TypeFor[User]()) != nil { + t.Fatal("a type without a binding or Broadcastable got a handler") + } + if got := defaultAlias(reflect.TypeFor[Widget]()); got != "lighthouse.widget" { + t.Fatalf("defaultAlias = %s", got) + } +} diff --git a/modules/lighthouse/postgres_test.go b/modules/lighthouse/postgres_test.go index 6614081..570959f 100644 --- a/modules/lighthouse/postgres_test.go +++ b/modules/lighthouse/postgres_test.go @@ -22,7 +22,7 @@ var ( lhSQL *sql.DB lhDSN string lhPGErr error - dbSeq atomic.Int64 + dbSeq atomic.Int64 ) func TestMain(m *testing.M) { diff --git a/modules/lighthouse/registry_test.go b/modules/lighthouse/registry_test.go new file mode 100644 index 0000000..6d7c919 --- /dev/null +++ b/modules/lighthouse/registry_test.go @@ -0,0 +1,73 @@ +package lighthouse + +import ( + "context" + "fmt" + "reflect" + "sync" + "testing" +) + +// TestRegistry covers the namespace registry: empty, colon, nil and +// duplicate registrations are refused, lookups are byte-exact, Namespaces +// is sorted, and concurrent lookups are safe (run under -race). +func TestRegistry(t *testing.T) { + r := NewRegistry() + allow := AuthorizerFunc(func(context.Context, uint, string) Result { return Allowed(nil) }) + if err := r.Register("b", allow); err != nil { + t.Fatal(err) + } + if err := r.Register("a", allow); err != nil { + t.Fatal(err) + } + for _, bad := range []struct { + ns string + a Authorizer + }{{"", allow}, {"x:y", allow}, {"c", nil}, {"a", allow}} { + if err := r.Register(bad.ns, bad.a); err == nil { + t.Errorf("Register(%q) accepted", bad.ns) + } + } + if got := r.Namespaces(); !reflect.DeepEqual(got, []string{"a", "b"}) { + t.Fatalf("Namespaces = %v", got) + } + if _, ok := r.Get("A"); ok { + t.Fatal("lookup is not case-sensitive") + } + if a, ok := r.Get("a"); !ok || !a.Authorize(context.Background(), 1, "a:1").Allowed { + t.Fatal("Get(a) failed") + } + d := Denied("why") + if d.Allowed || d.Reason() != "why" { + t.Fatalf("Denied = %+v", d) + } + var nilReg *Registry + if _, ok := nilReg.Get("a"); ok || nilReg.Namespaces() != nil { + t.Fatal("nil registry is not empty") + } + if Allowed(map[string]any{"k": 1}).Reason() != "" { + t.Fatal("an allow carries a reason") + } + + var wg sync.WaitGroup + for i := range 16 { + wg.Add(1) + go func() { + defer wg.Done() + if i%4 == 0 { + _ = r.Register(fmt.Sprintf("ns%d", i), allow) + } + for range 100 { + if _, ok := r.Get("a"); !ok { + t.Error("concurrent Get lost a namespace") + return + } + _ = r.Namespaces() + } + }() + } + wg.Wait() + if n := len(r.Namespaces()); n != 6 { + t.Fatalf("namespaces after concurrent registration = %d, want 6", n) + } +} diff --git a/modules/lighthouse/route_test.go b/modules/lighthouse/route_test.go new file mode 100644 index 0000000..016f1b4 --- /dev/null +++ b/modules/lighthouse/route_test.go @@ -0,0 +1,142 @@ +package lighthouse + +import ( + "context" + "encoding/json" + "net/http" + "strings" + "testing" + + "git.golem15.com/golem15/summercms/modules/pact" +) + +// recRouter records the groups and routes Mount registers. +type recRouter struct { + groups []recGroup + routes []string +} + +type recGroup struct { + raw bool + prefix string + middleware []string + routes []string +} + +func (r *recRouter) group(raw bool, prefix string, mw []string, fn func(pact.Router)) { + inner := &recRouter{} + fn(inner) + r.groups = append(r.groups, recGroup{raw: raw, prefix: prefix, middleware: append([]string(nil), mw...), routes: inner.routes}) +} + +func (r *recRouter) Group(prefix string, mw []string, fn func(pact.Router)) { + r.group(false, prefix, mw, fn) +} +func (r *recRouter) GroupRaw(prefix string, mw []string, fn func(pact.Router)) { + r.group(true, prefix, mw, fn) +} +func (r *recRouter) add(m, p string) { r.routes = append(r.routes, m+" "+p) } +func (r *recRouter) Get(p string, _ http.HandlerFunc, _ ...string) { r.add("GET", p) } +func (r *recRouter) Post(p string, _ http.HandlerFunc, _ ...string) { r.add("POST", p) } +func (r *recRouter) Put(p string, _ http.HandlerFunc, _ ...string) { r.add("PUT", p) } +func (r *recRouter) Patch(p string, _ http.HandlerFunc, _ ...string) { r.add("PATCH", p) } +func (r *recRouter) Delete(p string, _ http.HandlerFunc, _ ...string) { r.add("DELETE", p) } +func (r *recRouter) Where(string, string) {} +func (r *recRouter) WhereIn(string, ...string) {} + +// routeDriver is a driver that declares the given routes. +type routeDriver struct{ routes []Route } + +func (d routeDriver) Name() string { return "acme" } +func (d routeDriver) Routes() []Route { return d.routes } +func (routeDriver) Publish(context.Context, string, string, json.RawMessage) error { + return nil +} +func (routeDriver) Broadcast(context.Context, []string, string, json.RawMessage) error { + return nil +} + +func okHandler(http.ResponseWriter, *http.Request) {} + +// TestMountSurfaces covers D-13 and T-11-19: UserAuth and Public routes +// mount in a group with the surface middleware then the shared middleware, +// ServerToServer routes in a raw group, every method is registered, a +// UserAuth route without a guard is refused before anything is mounted, +// and a nil driver or the null driver mounts nothing. +func TestMountSurfaces(t *testing.T) { + s := Surfaces{UserAuth: []string{"jwt.auth"}, ServerToServer: []string{"acme.s2s"}, Public: []string{"acme.public"}, Middleware: []string{"throttle:ws-api"}} + d := routeDriver{routes: []Route{ + {Name: "token", Method: "get", Path: "/api/rt/token", Surface: UserAuth, Handler: okHandler}, + {Name: "subscribe", Method: http.MethodPost, Path: "/api/rt/subscribe", Surface: ServerToServer, Handler: okHandler}, + {Name: "put", Method: http.MethodPut, Path: "/api/rt/p", Surface: Public, Handler: okHandler}, + {Name: "patch", Method: http.MethodPatch, Path: "/api/rt/p", Surface: Public, Handler: okHandler}, + {Name: "delete", Method: http.MethodDelete, Path: "/api/rt/p", Surface: Public, Handler: okHandler}, + }} + r := &recRouter{} + if err := Mount(r, d, s); err != nil { + t.Fatal(err) + } + want := []recGroup{ + {false, "/", []string{"jwt.auth", "throttle:ws-api"}, []string{"GET /api/rt/token"}}, + {true, "/", []string{"acme.s2s", "throttle:ws-api"}, []string{"POST /api/rt/subscribe"}}, + {false, "/", []string{"acme.public", "throttle:ws-api"}, []string{"PUT /api/rt/p"}}, + {false, "/", []string{"acme.public", "throttle:ws-api"}, []string{"PATCH /api/rt/p"}}, + {false, "/", []string{"acme.public", "throttle:ws-api"}, []string{"DELETE /api/rt/p"}}, + } + if len(r.groups) != len(want) { + t.Fatalf("groups = %+v", r.groups) + } + for i, w := range want { + g := r.groups[i] + if g.raw != w.raw || g.prefix != w.prefix || strings.Join(g.middleware, ",") != strings.Join(w.middleware, ",") || strings.Join(g.routes, ",") != strings.Join(w.routes, ",") { + t.Errorf("group %d = %+v, want %+v", i, g, w) + } + } + if len(s.UserAuth) != 1 || s.UserAuth[0] != "jwt.auth" { + t.Fatal("Mount mutated the caller's surface middleware") + } + + refusals := []struct { + name string + route Route + s Surfaces + want string + }{ + {"user_route_without_guard", Route{Name: "token", Method: "GET", Path: "/t", Surface: UserAuth, Handler: okHandler}, Surfaces{Middleware: []string{"throttle"}}, "Surfaces.UserAuth is empty"}, + {"no_handler", Route{Name: "x", Method: "GET", Path: "/t", Surface: Public}, s, "has no handler"}, + {"relative_path", Route{Name: "x", Method: "GET", Path: "t", Surface: Public, Handler: okHandler}, s, "must be absolute"}, + {"bad_method", Route{Name: "x", Method: "TRACE", Path: "/t", Surface: Public, Handler: okHandler}, s, "unsupported method"}, + {"unknown_surface", Route{Name: "x", Method: "GET", Path: "/t", Surface: Surface(9), Handler: okHandler}, s, "unknown surface Surface(9)"}, + } + for _, c := range refusals { + t.Run(c.name, func(t *testing.T) { + r := &recRouter{} + // A valid route first: nothing may be mounted when a later one + // is invalid. + valid := Route{Name: "ok", Method: "POST", Path: "/ok", Surface: ServerToServer, Handler: okHandler} + err := Mount(r, routeDriver{routes: []Route{valid, c.route}}, c.s) + if err == nil || !strings.Contains(err.Error(), c.want) || !strings.Contains(err.Error(), "driver acme route") { + t.Fatalf("err = %v, want %q", err, c.want) + } + if len(r.groups) != 0 { + t.Fatalf("mounted %d group(s) before refusing", len(r.groups)) + } + }) + } + + if err := Mount(nil, d, s); err == nil { + t.Fatal("nil router accepted") + } + r = &recRouter{} + if err := Mount(r, nil, s); err != nil || len(r.groups) != 0 { + t.Fatalf("nil driver: err %v, groups %d", err, len(r.groups)) + } + if err := Mount(r, nullDriver{}, Surfaces{}); err != nil || len(r.groups) != 0 { + t.Fatalf("null driver: err %v, groups %d", err, len(r.groups)) + } + for surface, name := range map[Surface]string{UserAuth: "UserAuth", ServerToServer: "ServerToServer", Public: "Public", Surface(0): "Surface(0)"} { + if surface.String() != name { + t.Errorf("Surface(%d).String() = %q", int(surface), surface.String()) + } + } +}