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%)
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
260
modules/lighthouse/centrifugo/client_test.go
Normal file
260
modules/lighthouse/centrifugo/client_test.go
Normal file
@@ -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<z>"}`)); 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<z>"},"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")
|
||||
}
|
||||
}
|
||||
222
modules/lighthouse/centrifugo/proxy_test.go
Normal file
222
modules/lighthouse/centrifugo/proxy_test.go
Normal file
@@ -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()
|
||||
})
|
||||
}
|
||||
210
modules/lighthouse/centrifugo/token_test.go
Normal file
210
modules/lighthouse/centrifugo/token_test.go
Normal file
@@ -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()
|
||||
})
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
248
modules/lighthouse/lighthouse_test.go
Normal file
248
modules/lighthouse/lighthouse_test.go
Normal file
@@ -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)
|
||||
}
|
||||
}
|
||||
73
modules/lighthouse/registry_test.go
Normal file
73
modules/lighthouse/registry_test.go
Normal file
@@ -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)
|
||||
}
|
||||
}
|
||||
142
modules/lighthouse/route_test.go
Normal file
142
modules/lighthouse/route_test.go
Normal file
@@ -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())
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user