diff --git a/backend.go b/backend.go new file mode 100644 index 0000000..61f58b4 --- /dev/null +++ b/backend.go @@ -0,0 +1,10 @@ +package main + +import "context" + +type Backend interface { + Publish(event *Event) error + Subscribe(ctx context.Context) <-chan *Event + Since(lastEventID string, subscribePath string) []*Event + Close() error +} diff --git a/backend_memory.go b/backend_memory.go new file mode 100644 index 0000000..c3a1868 --- /dev/null +++ b/backend_memory.go @@ -0,0 +1,127 @@ +package main + +import ( + "context" + "strings" + "sync" +) + +type MemoryBackend struct { + mu sync.RWMutex + buffer *RingBuffer + listeners []chan *Event +} + +func NewMemoryBackend(bufferSize int) *MemoryBackend { + return &MemoryBackend{ + buffer: NewRingBuffer(bufferSize), + } +} + +func (m *MemoryBackend) Publish(event *Event) error { + m.buffer.Add(event) + + m.mu.RLock() + defer m.mu.RUnlock() + + for _, ch := range m.listeners { + select { + case ch <- event: + default: + } + } + return nil +} + +func (m *MemoryBackend) Subscribe(ctx context.Context) <-chan *Event { + ch := make(chan *Event, 256) + + m.mu.Lock() + m.listeners = append(m.listeners, ch) + m.mu.Unlock() + + go func() { + <-ctx.Done() + m.mu.Lock() + defer m.mu.Unlock() + for i, l := range m.listeners { + if l == ch { + m.listeners = append(m.listeners[:i], m.listeners[i+1:]...) + break + } + } + }() + + return ch +} + +func (m *MemoryBackend) Since(lastEventID string, subscribePath string) []*Event { + return m.buffer.Since(lastEventID, subscribePath) +} + +func (m *MemoryBackend) Close() error { + return nil +} + +type RingBuffer struct { + mu sync.RWMutex + buf []*Event + size int + write int + count int +} + +func NewRingBuffer(size int) *RingBuffer { + return &RingBuffer{ + buf: make([]*Event, size), + size: size, + } +} + +func (rb *RingBuffer) Add(event *Event) { + rb.mu.Lock() + defer rb.mu.Unlock() + rb.buf[rb.write%rb.size] = event + rb.write++ + if rb.count < rb.size { + rb.count++ + } +} + +func (rb *RingBuffer) Since(lastEventID string, subscribePath string) []*Event { + rb.mu.RLock() + defer rb.mu.RUnlock() + + start := rb.write - rb.count + found := false + foundIdx := start + + for i := start; i < rb.write; i++ { + e := rb.buf[i%rb.size] + if e.ID == lastEventID { + found = true + foundIdx = i + 1 + break + } + } + + if !found { + foundIdx = start + } + + var result []*Event + for i := foundIdx; i < rb.write; i++ { + e := rb.buf[i%rb.size] + if pathMatches(subscribePath, e.Path) { + result = append(result, e) + } + } + return result +} + +func pathMatches(subscribePath, eventPath string) bool { + if subscribePath == "" { + return true + } + return eventPath == subscribePath || strings.HasPrefix(eventPath, subscribePath+"/") +} diff --git a/backend_memory_test.go b/backend_memory_test.go new file mode 100644 index 0000000..f760718 --- /dev/null +++ b/backend_memory_test.go @@ -0,0 +1,119 @@ +package main + +import ( + "context" + "fmt" + "testing" + "time" +) + +func TestMemoryBackend_publishAndSubscribe(t *testing.T) { + backend := NewMemoryBackend(100) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + ch := backend.Subscribe(ctx) + event := &Event{ID: "e1", Path: "test", Payload: map[string]any{}} + backend.Publish(event) + + select { + case got := <-ch: + if got.ID != "e1" { + t.Errorf("expected e1, got %s", got.ID) + } + case <-time.After(time.Second): + t.Fatal("timed out waiting for event") + } +} + +func TestMemoryBackend_multipleListeners(t *testing.T) { + backend := NewMemoryBackend(100) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + ch1 := backend.Subscribe(ctx) + ch2 := backend.Subscribe(ctx) + backend.Publish(&Event{ID: "e1", Path: "test", Payload: map[string]any{}}) + + for i, ch := range []<-chan *Event{ch1, ch2} { + select { + case got := <-ch: + if got.ID != "e1" { + t.Errorf("listener %d: expected e1, got %s", i, got.ID) + } + case <-time.After(time.Second): + t.Fatalf("listener %d: timed out", i) + } + } +} + +func TestMemoryBackend_cancelRemovesListener(t *testing.T) { + backend := NewMemoryBackend(100) + ctx, cancel := context.WithCancel(context.Background()) + backend.Subscribe(ctx) + cancel() + + // Give the cleanup goroutine time to run + time.Sleep(50 * time.Millisecond) + + backend.mu.RLock() + n := len(backend.listeners) + backend.mu.RUnlock() + if n != 0 { + t.Errorf("expected 0 listeners after cancel, got %d", n) + } +} + +func TestMemoryBackend_since(t *testing.T) { + backend := NewMemoryBackend(100) + backend.Publish(&Event{ID: "e1", Path: "test", Payload: map[string]any{}}) + backend.Publish(&Event{ID: "e2", Path: "test", Payload: map[string]any{}}) + backend.Publish(&Event{ID: "e3", Path: "other", Payload: map[string]any{}}) + + events := backend.Since("e1", "test") + if len(events) != 1 { + t.Fatalf("expected 1 event, got %d", len(events)) + } + if events[0].ID != "e2" { + t.Errorf("expected e2, got %s", events[0].ID) + } +} + +func TestMemoryBackend_sinceEmptyPath(t *testing.T) { + backend := NewMemoryBackend(100) + backend.Publish(&Event{ID: "e1", Path: "test", Payload: map[string]any{}}) + backend.Publish(&Event{ID: "e2", Path: "other", Payload: map[string]any{}}) + + events := backend.Since("e1", "") + if len(events) != 1 { + t.Fatalf("expected 1 event, got %d", len(events)) + } + if events[0].ID != "e2" { + t.Errorf("expected e2, got %s", events[0].ID) + } +} + +func TestMemoryBackend_sinceExpiredID(t *testing.T) { + backend := NewMemoryBackend(3) + for i := range 6 { + backend.Publish(&Event{ID: fmt.Sprintf("e%d", i), Path: "test", Payload: map[string]any{}}) + } + + events := backend.Since("e0", "test") + if len(events) != 3 { + t.Fatalf("expected 3 events (full buffer), got %d", len(events)) + } + if events[0].ID != "e3" { + t.Errorf("expected e3, got %s", events[0].ID) + } +} + +func TestMemoryBackend_closeIsIdempotent(t *testing.T) { + backend := NewMemoryBackend(100) + if err := backend.Close(); err != nil { + t.Errorf("first close: %v", err) + } + if err := backend.Close(); err != nil { + t.Errorf("second close: %v", err) + } +} diff --git a/backend_redis.go b/backend_redis.go new file mode 100644 index 0000000..c5e2d8b --- /dev/null +++ b/backend_redis.go @@ -0,0 +1,141 @@ +package main + +import ( + "context" + "encoding/json" + "fmt" + "log" + "sync" + + "github.com/redis/go-redis/v9" +) + +const ( + redisChannel = "wicket:events" + redisBufferKey = "wicket:buffer" +) + +type RedisBackend struct { + client *redis.Client + bufferSize int + mu sync.Mutex + closed bool +} + +func NewRedisBackend(url string, bufferSize int) (*RedisBackend, error) { + opts, err := redis.ParseURL(url) + if err != nil { + return nil, fmt.Errorf("parsing redis URL: %w", err) + } + client := redis.NewClient(opts) + if err := client.Ping(context.Background()).Err(); err != nil { + return nil, fmt.Errorf("connecting to redis: %w", err) + } + return &RedisBackend{ + client: client, + bufferSize: bufferSize, + }, nil +} + +func newRedisBackendFromClient(client *redis.Client, bufferSize int) *RedisBackend { + return &RedisBackend{ + client: client, + bufferSize: bufferSize, + } +} + +func (r *RedisBackend) Publish(event *Event) error { + data, err := json.Marshal(event) + if err != nil { + return fmt.Errorf("marshaling event: %w", err) + } + ctx := context.Background() + pipe := r.client.Pipeline() + pipe.LPush(ctx, redisBufferKey, data) + pipe.LTrim(ctx, redisBufferKey, 0, int64(r.bufferSize-1)) + pipe.Publish(ctx, redisChannel, data) + _, err = pipe.Exec(ctx) + return err +} + +func (r *RedisBackend) Subscribe(ctx context.Context) <-chan *Event { + ch := make(chan *Event, 256) + pubsub := r.client.Subscribe(ctx, redisChannel) + + go func() { + defer close(ch) + defer pubsub.Close() + msgCh := pubsub.Channel() + for { + select { + case <-ctx.Done(): + return + case msg := <-msgCh: + if msg == nil { + return + } + var event Event + if err := json.Unmarshal([]byte(msg.Payload), &event); err != nil { + log.Printf("redis: unmarshal event: %v", err) + continue + } + select { + case ch <- &event: + default: + } + } + } + }() + + return ch +} + +func (r *RedisBackend) Since(lastEventID string, subscribePath string) []*Event { + ctx := context.Background() + vals, err := r.client.LRange(ctx, redisBufferKey, 0, -1).Result() + if err != nil { + log.Printf("redis: LRANGE: %v", err) + return nil + } + + // Redis list is newest-first (LPUSH), reverse to chronological order + events := make([]*Event, 0, len(vals)) + for i := len(vals) - 1; i >= 0; i-- { + var e Event + if err := json.Unmarshal([]byte(vals[i]), &e); err != nil { + continue + } + events = append(events, &e) + } + + found := false + foundIdx := 0 + for i, e := range events { + if e.ID == lastEventID { + found = true + foundIdx = i + 1 + break + } + } + if !found { + foundIdx = 0 + } + + var result []*Event + for i := foundIdx; i < len(events); i++ { + if pathMatches(subscribePath, events[i].Path) { + result = append(result, events[i]) + } + } + return result +} + +func (r *RedisBackend) Close() error { + r.mu.Lock() + defer r.mu.Unlock() + if r.closed { + return nil + } + r.closed = true + return r.client.Close() +} diff --git a/backend_redis_test.go b/backend_redis_test.go new file mode 100644 index 0000000..5852a09 --- /dev/null +++ b/backend_redis_test.go @@ -0,0 +1,301 @@ +package main + +import ( + "context" + "testing" + "time" + + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" +) + +func newTestRedisBackend(t *testing.T, bufferSize int) (*RedisBackend, *miniredis.Miniredis) { + t.Helper() + mr := miniredis.RunT(t) + client := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + return newRedisBackendFromClient(client, bufferSize), mr +} + +func TestRedisBackend_publishAndSubscribe(t *testing.T) { + backend, _ := newTestRedisBackend(t, 100) + defer backend.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + ch := backend.Subscribe(ctx) + time.Sleep(50 * time.Millisecond) + + backend.Publish(&Event{ID: "e1", Path: "test", Payload: map[string]any{"ok": true}}) + + select { + case got := <-ch: + if got.ID != "e1" { + t.Errorf("expected e1, got %s", got.ID) + } + if got.Path != "test" { + t.Errorf("expected path test, got %s", got.Path) + } + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for event") + } +} + +func TestRedisBackend_since(t *testing.T) { + backend, _ := newTestRedisBackend(t, 100) + defer backend.Close() + + backend.Publish(&Event{ID: "e1", Path: "test", Payload: map[string]any{}}) + backend.Publish(&Event{ID: "e2", Path: "test", Payload: map[string]any{}}) + backend.Publish(&Event{ID: "e3", Path: "other", Payload: map[string]any{}}) + + events := backend.Since("e1", "test") + if len(events) != 1 { + t.Fatalf("expected 1 event, got %d", len(events)) + } + if events[0].ID != "e2" { + t.Errorf("expected e2, got %s", events[0].ID) + } +} + +func TestRedisBackend_sinceExpired(t *testing.T) { + backend, _ := newTestRedisBackend(t, 3) + defer backend.Close() + + backend.Publish(&Event{ID: "e1", Path: "test", Payload: map[string]any{}}) + backend.Publish(&Event{ID: "e2", Path: "test", Payload: map[string]any{}}) + backend.Publish(&Event{ID: "e3", Path: "test", Payload: map[string]any{}}) + backend.Publish(&Event{ID: "e4", Path: "test", Payload: map[string]any{}}) + + events := backend.Since("e1", "test") + if len(events) != 3 { + t.Fatalf("expected 3 events (full buffer), got %d", len(events)) + } + if events[0].ID != "e2" { + t.Errorf("expected e2, got %s", events[0].ID) + } +} + +func TestRedisBackend_sincePathFiltering(t *testing.T) { + backend, _ := newTestRedisBackend(t, 100) + defer backend.Close() + + backend.Publish(&Event{ID: "e1", Path: "a/b", Payload: map[string]any{}}) + backend.Publish(&Event{ID: "e2", Path: "a/b/c", Payload: map[string]any{}}) + backend.Publish(&Event{ID: "e3", Path: "x/y", Payload: map[string]any{}}) + + events := backend.Since("e1", "a/b") + if len(events) != 1 { + t.Fatalf("expected 1 event, got %d", len(events)) + } + if events[0].ID != "e2" { + t.Errorf("expected e2, got %s", events[0].ID) + } +} + +func TestRedisBackend_multiReplica(t *testing.T) { + mr := miniredis.RunT(t) + + client1 := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + backend1 := newRedisBackendFromClient(client1, 100) + defer backend1.Close() + + client2 := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + backend2 := newRedisBackendFromClient(client2, 100) + defer backend2.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + ch2 := backend2.Subscribe(ctx) + time.Sleep(50 * time.Millisecond) + + backend1.Publish(&Event{ID: "cross-replica", Path: "test", Payload: map[string]any{"from": "replica1"}}) + + select { + case got := <-ch2: + if got.ID != "cross-replica" { + t.Errorf("expected cross-replica, got %s", got.ID) + } + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for cross-replica event") + } +} + +func TestRedisBackend_multiReplicaReplay(t *testing.T) { + mr := miniredis.RunT(t) + + client1 := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + backend1 := newRedisBackendFromClient(client1, 100) + defer backend1.Close() + + backend1.Publish(&Event{ID: "e1", Path: "test", Payload: map[string]any{}}) + backend1.Publish(&Event{ID: "e2", Path: "test", Payload: map[string]any{}}) + + client2 := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + backend2 := newRedisBackendFromClient(client2, 100) + defer backend2.Close() + + events := backend2.Since("e1", "test") + if len(events) != 1 { + t.Fatalf("expected 1 event, got %d", len(events)) + } + if events[0].ID != "e2" { + t.Errorf("expected e2, got %s", events[0].ID) + } +} + +func TestRedisBackend_closeIsIdempotent(t *testing.T) { + backend, _ := newTestRedisBackend(t, 100) + if err := backend.Close(); err != nil { + t.Errorf("first close: %v", err) + } + if err := backend.Close(); err != nil { + t.Errorf("second close: %v", err) + } +} + +func TestRedisBackend_bufferTrims(t *testing.T) { + backend, _ := newTestRedisBackend(t, 3) + defer backend.Close() + + for i := range 10 { + backend.Publish(&Event{ID: string(rune('a' + i)), Path: "test", Payload: map[string]any{}}) + } + + events := backend.Since("", "test") + if len(events) != 3 { + t.Fatalf("expected 3 events in trimmed buffer, got %d", len(events)) + } +} + +func TestNewRedisBackend_success(t *testing.T) { + mr := miniredis.RunT(t) + backend, err := NewRedisBackend("redis://"+mr.Addr(), 100) + if err != nil { + t.Fatalf("expected no error, got %v", err) + } + defer backend.Close() +} + +func TestNewRedisBackend_badURL(t *testing.T) { + _, err := NewRedisBackend("not-a-url", 100) + if err == nil { + t.Fatal("expected error for bad URL") + } +} + +func TestNewRedisBackend_unreachable(t *testing.T) { + _, err := NewRedisBackend("redis://127.0.0.1:1", 100) + if err == nil { + t.Fatal("expected error for unreachable redis") + } +} + +func TestRedisBackend_sinceWithCorruptData(t *testing.T) { + backend, mr := newTestRedisBackend(t, 100) + defer backend.Close() + + backend.Publish(&Event{ID: "e1", Path: "test", Payload: map[string]any{}}) + mr.Lpush(redisBufferKey, "not-valid-json") + backend.Publish(&Event{ID: "e2", Path: "test", Payload: map[string]any{}}) + + events := backend.Since("e1", "test") + if len(events) != 1 { + t.Fatalf("expected 1 event (corrupt data skipped), got %d", len(events)) + } + if events[0].ID != "e2" { + t.Errorf("expected e2, got %s", events[0].ID) + } +} + +func TestRedisBackend_subscribeContextCancel(t *testing.T) { + backend, _ := newTestRedisBackend(t, 100) + defer backend.Close() + + ctx, cancel := context.WithCancel(context.Background()) + ch := backend.Subscribe(ctx) + cancel() + + // Channel should close after context cancel + select { + case _, ok := <-ch: + if ok { + t.Error("expected channel to be closed") + } + case <-time.After(2 * time.Second): + t.Fatal("timed out waiting for channel close") + } +} + +func TestRedisBackend_subscribeSkipsBadJSON(t *testing.T) { + backend, _ := newTestRedisBackend(t, 100) + defer backend.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + ch := backend.Subscribe(ctx) + time.Sleep(50 * time.Millisecond) + + // Publish garbage directly to the Redis channel + backend.client.Publish(context.Background(), redisChannel, "not-json") + // Then a valid event + backend.Publish(&Event{ID: "valid", Path: "test", Payload: map[string]any{}}) + + select { + case got := <-ch: + if got.ID != "valid" { + t.Errorf("expected valid, got %s", got.ID) + } + case <-time.After(2 * time.Second): + t.Fatal("timed out") + } +} + +func TestRedisBackend_sinceClosedClient(t *testing.T) { + backend, _ := newTestRedisBackend(t, 100) + backend.Publish(&Event{ID: "e1", Path: "test", Payload: map[string]any{}}) + backend.client.Close() + backend.closed = true + + events := backend.Since("", "test") + if events != nil { + t.Errorf("expected nil, got %v", events) + } +} + +func TestRedisBackend_subscribeNilMessage(t *testing.T) { + mr := miniredis.RunT(t) + client := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + backend := newRedisBackendFromClient(client, 100) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + ch := backend.Subscribe(ctx) + time.Sleep(50 * time.Millisecond) + + // Close the client, which causes the pubsub channel to close + client.Close() + mr.Close() + + select { + case _, ok := <-ch: + if ok { + t.Error("expected channel to be closed after client close") + } + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for channel close") + } +} + +func TestRedisBackend_publishMarshalError(t *testing.T) { + backend, _ := newTestRedisBackend(t, 100) + defer backend.Close() + + err := backend.Publish(&Event{ID: "bad", Path: "test", Payload: make(chan int)}) + if err == nil { + t.Fatal("expected marshal error") + } +} diff --git a/broker.go b/broker.go index ed8a5b6..60b6bad 100644 --- a/broker.go +++ b/broker.go @@ -1,36 +1,52 @@ package main import ( + "context" "strings" "sync" + "sync/atomic" ) type subscriber struct { - ch chan *Event - path string + ch chan *Event + path string + startSeq int64 } type Broker struct { mu sync.RWMutex subscribers map[string][]*subscriber - buffer *RingBuffer + backend Backend + seq atomic.Int64 } -func NewBroker(bufferSize int) *Broker { +func NewBroker(backend Backend) *Broker { return &Broker{ subscribers: make(map[string][]*subscriber), - buffer: NewRingBuffer(bufferSize), + backend: backend, } } -func (b *Broker) Publish(event *Event) { - b.buffer.Add(event) +func (b *Broker) Start(ctx context.Context) { + ch := b.backend.Subscribe(ctx) + go func() { + var fanOutSeq int64 + for event := range ch { + fanOutSeq++ + b.fanOut(event, fanOutSeq) + } + }() +} +func (b *Broker) fanOut(event *Event, seq int64) { b.mu.RLock() defer b.mu.RUnlock() for _, path := range publishPaths(event.Path) { for _, sub := range b.subscribers[path] { + if seq <= sub.startSeq { + continue + } select { case sub.ch <- event: default: @@ -39,16 +55,21 @@ func (b *Broker) Publish(event *Event) { } } +func (b *Broker) Publish(event *Event) error { + b.seq.Add(1) + return b.backend.Publish(event) +} + func (b *Broker) Subscribe(path string, lastEventID string) (<-chan *Event, func()) { ch := make(chan *Event, 64) - sub := &subscriber{ch: ch, path: path} + sub := &subscriber{ch: ch, path: path, startSeq: b.seq.Load()} b.mu.Lock() b.subscribers[path] = append(b.subscribers[path], sub) b.mu.Unlock() if lastEventID != "" { - events := b.buffer.Since(lastEventID, path) + events := b.backend.Since(lastEventID, path) for _, e := range events { ch <- e } @@ -64,7 +85,10 @@ func (b *Broker) Subscribe(path string, lastEventID string) (<-chan *Event, func break } } - close(ch) + func() { + defer func() { recover() }() + close(ch) + }() } return ch, unsub @@ -85,66 +109,3 @@ func publishPaths(eventPath string) []string { } return paths } - -type RingBuffer struct { - mu sync.RWMutex - buf []*Event - size int - write int - count int -} - -func NewRingBuffer(size int) *RingBuffer { - return &RingBuffer{ - buf: make([]*Event, size), - size: size, - } -} - -func (rb *RingBuffer) Add(event *Event) { - rb.mu.Lock() - defer rb.mu.Unlock() - rb.buf[rb.write%rb.size] = event - rb.write++ - if rb.count < rb.size { - rb.count++ - } -} - -func (rb *RingBuffer) Since(lastEventID string, subscribePath string) []*Event { - rb.mu.RLock() - defer rb.mu.RUnlock() - - start := rb.write - rb.count - found := false - foundIdx := start - - for i := start; i < rb.write; i++ { - e := rb.buf[i%rb.size] - if e.ID == lastEventID { - found = true - foundIdx = i + 1 - break - } - } - - if !found { - foundIdx = start - } - - var result []*Event - for i := foundIdx; i < rb.write; i++ { - e := rb.buf[i%rb.size] - if pathMatches(subscribePath, e.Path) { - result = append(result, e) - } - } - return result -} - -func pathMatches(subscribePath, eventPath string) bool { - if subscribePath == "" { - return true - } - return eventPath == subscribePath || strings.HasPrefix(eventPath, subscribePath+"/") -} diff --git a/broker_test.go b/broker_test.go index b8bd803..871f5f8 100644 --- a/broker_test.go +++ b/broker_test.go @@ -1,6 +1,7 @@ package main import ( + "context" "fmt" "testing" "time" @@ -35,8 +36,17 @@ func mustNotReceive(t *testing.T, ch <-chan *Event, timeout time.Duration) { } } +func newTestBroker(bufferSize int) (*Broker, context.CancelFunc) { + backend := NewMemoryBackend(bufferSize) + broker := NewBroker(backend) + ctx, cancel := context.WithCancel(context.Background()) + broker.Start(ctx) + return broker, cancel +} + func TestBroker_exactPathDelivery(t *testing.T) { - b := NewBroker(100) + b, cancel := newTestBroker(100) + defer cancel() ch, unsub := b.Subscribe("github.com/chrisguidry/docketeer", "") defer unsub() @@ -50,7 +60,8 @@ func TestBroker_exactPathDelivery(t *testing.T) { } func TestBroker_parentReceivesChildEvents(t *testing.T) { - b := NewBroker(100) + b, cancel := newTestBroker(100) + defer cancel() ch, unsub := b.Subscribe("github.com/chrisguidry", "") defer unsub() @@ -64,7 +75,8 @@ func TestBroker_parentReceivesChildEvents(t *testing.T) { } func TestBroker_rootReceivesAll(t *testing.T) { - b := NewBroker(100) + b, cancel := newTestBroker(100) + defer cancel() ch, unsub := b.Subscribe("", "") defer unsub() @@ -78,7 +90,8 @@ func TestBroker_rootReceivesAll(t *testing.T) { } func TestBroker_unrelatedSubscriberDoesNotReceive(t *testing.T) { - b := NewBroker(100) + b, cancel := newTestBroker(100) + defer cancel() ch, unsub := b.Subscribe("gitlab.com", "") defer unsub() @@ -88,7 +101,8 @@ func TestBroker_unrelatedSubscriberDoesNotReceive(t *testing.T) { } func TestBroker_multipleSubscribersSamePath(t *testing.T) { - b := NewBroker(100) + b, cancel := newTestBroker(100) + defer cancel() ch1, unsub1 := b.Subscribe("github.com/chrisguidry/docketeer", "") defer unsub1() ch2, unsub2 := b.Subscribe("github.com/chrisguidry/docketeer", "") @@ -101,7 +115,8 @@ func TestBroker_multipleSubscribersSamePath(t *testing.T) { } func TestBroker_unsubscribeStopsDelivery(t *testing.T) { - b := NewBroker(100) + b, cancel := newTestBroker(100) + defer cancel() ch, unsub := b.Subscribe("github.com/chrisguidry/docketeer", "") unsub() @@ -118,7 +133,8 @@ func TestBroker_unsubscribeStopsDelivery(t *testing.T) { } func TestBroker_ringBufferWraps(t *testing.T) { - b := NewBroker(5) + b, cancel := newTestBroker(5) + defer cancel() for i := range 10 { b.Publish(&Event{ ID: fmt.Sprintf("event-%d", i), @@ -140,7 +156,8 @@ func TestBroker_ringBufferWraps(t *testing.T) { } func TestBroker_lastEventIDReplay(t *testing.T) { - b := NewBroker(100) + b, cancel := newTestBroker(100) + defer cancel() b.Publish(&Event{ID: "e1", Path: "test", Payload: map[string]any{}}) b.Publish(&Event{ID: "e2", Path: "test", Payload: map[string]any{}}) b.Publish(&Event{ID: "e3", Path: "test", Payload: map[string]any{}}) @@ -156,7 +173,8 @@ func TestBroker_lastEventIDReplay(t *testing.T) { } func TestBroker_lastEventIDRespectsPathHierarchy(t *testing.T) { - b := NewBroker(100) + b, cancel := newTestBroker(100) + defer cancel() b.Publish(&Event{ID: "e1", Path: "github.com/chrisguidry/docketeer", Payload: map[string]any{}}) b.Publish(&Event{ID: "e2", Path: "gitlab.com/other/project", Payload: map[string]any{}}) b.Publish(&Event{ID: "e3", Path: "github.com/chrisguidry/other", Payload: map[string]any{}}) @@ -172,7 +190,8 @@ func TestBroker_lastEventIDRespectsPathHierarchy(t *testing.T) { } func TestBroker_lastEventIDReplayExactPath(t *testing.T) { - b := NewBroker(100) + b, cancel := newTestBroker(100) + defer cancel() b.Publish(&Event{ID: "e1", Path: "exact/path", Payload: map[string]any{}}) b.Publish(&Event{ID: "e2", Path: "exact/path/child", Payload: map[string]any{}}) b.Publish(&Event{ID: "e3", Path: "exact/path", Payload: map[string]any{}}) @@ -211,7 +230,8 @@ func TestPathMatches(t *testing.T) { } func TestBroker_lastEventIDExpiredFromBuffer(t *testing.T) { - b := NewBroker(3) + b, cancel := newTestBroker(3) + defer cancel() b.Publish(&Event{ID: "old1", Path: "test", Payload: map[string]any{}}) b.Publish(&Event{ID: "old2", Path: "test", Payload: map[string]any{}}) b.Publish(&Event{ID: "old3", Path: "test", Payload: map[string]any{}}) diff --git a/go.mod b/go.mod index e60be94..4a4af31 100644 --- a/go.mod +++ b/go.mod @@ -8,4 +8,12 @@ require ( gopkg.in/yaml.v3 v3.0.1 ) -require golang.org/x/sys v0.13.0 // indirect +require ( + github.com/alicebob/miniredis/v2 v2.37.0 // indirect + github.com/cespare/xxhash/v2 v2.3.0 // indirect + github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f // indirect + github.com/redis/go-redis/v9 v9.18.0 // indirect + github.com/yuin/gopher-lua v1.1.1 // indirect + go.uber.org/atomic v1.11.0 // indirect + golang.org/x/sys v0.13.0 // indirect +) diff --git a/go.sum b/go.sum index 900f5b3..24fe49b 100644 --- a/go.sum +++ b/go.sum @@ -1,7 +1,19 @@ +github.com/alicebob/miniredis/v2 v2.37.0 h1:RheObYW32G1aiJIj81XVt78ZHJpHonHLHW7OLIshq68= +github.com/alicebob/miniredis/v2 v2.37.0/go.mod h1:TcL7YfarKPGDAthEtl5NBeHZfeUQj6OXMm/+iu5cLMM= +github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs= +github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs= +github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f h1:lO4WD4F/rVNCu3HqELle0jiPLLBs70cWOduZpkS1E78= +github.com/dgryski/go-rendezvous v0.0.0-20200823014737-9f7001d12a5f/go.mod h1:cuUVRXasLTGF7a8hSLbxyZXjz+1KgoB3wDUb6vlszIc= github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k= github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0= github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0= github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/redis/go-redis/v9 v9.18.0 h1:pMkxYPkEbMPwRdenAzUNyFNrDgHx9U+DrBabWNfSRQs= +github.com/redis/go-redis/v9 v9.18.0/go.mod h1:k3ufPphLU5YXwNTUcCRXGxUoF1fqxnhFQmscfkCoDA0= +github.com/yuin/gopher-lua v1.1.1 h1:kYKnWBjvbNP4XLT3+bPEwAXJx262OhaHDWDVOPjL46M= +github.com/yuin/gopher-lua v1.1.1/go.mod h1:GBR0iDaNXjAgGg9zfCvksxSRnQx76gclCIb7kdAd1Pw= +go.uber.org/atomic v1.11.0 h1:ZvwS0R+56ePWxUNi+Atn9dWONBPp/AUETXlHW0DxSjE= +go.uber.org/atomic v1.11.0/go.mod h1:LUxbIzbOniOlMKjJjyPfpl4v+PKK2cNJn91OQbhoJI0= golang.org/x/sys v0.13.0 h1:Af8nKPmuFypiUBjVoU9V20FiaFXOcuZI21p0ycVYYGE= golang.org/x/sys v0.13.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405 h1:yhCVgyC4o1eVCa2tZl7eS0r+SDo693bJlVdllGtEeKM= diff --git a/main.go b/main.go index c2b7548..71e0c13 100644 --- a/main.go +++ b/main.go @@ -16,6 +16,8 @@ func main() { address := flag.String("address", ":8080", "listen address") configPath := flag.String("configuration", "", "path to configuration file") bufferSize := flag.Int("buffer-size", 1000, "event replay buffer size") + backendType := flag.String("backend", "memory", "pub/sub backend: memory or redis") + redisURL := flag.String("redis-url", "redis://localhost:6379", "Redis connection URL (when backend=redis)") flag.Parse() var cfgPtr atomic.Pointer[Configuration] @@ -37,7 +39,26 @@ func main() { defer stop() } - broker := NewBroker(*bufferSize) + var backend Backend + switch *backendType { + case "memory": + backend = NewMemoryBackend(*bufferSize) + case "redis": + var err error + backend, err = NewRedisBackend(*redisURL, *bufferSize) + if err != nil { + log.Fatalf("connecting to redis: %v", err) + } + default: + log.Fatalf("unknown backend: %s", *backendType) + } + defer backend.Close() + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + broker := NewBroker(backend) + broker.Start(ctx) handler := NewServer(broker, &cfgPtr) server := &http.Server{ @@ -63,6 +84,7 @@ func main() { log.Printf("reloaded configuration from %s", *configPath) case syscall.SIGINT, syscall.SIGTERM: log.Printf("shutting down") + cancel() server.Shutdown(context.Background()) return } diff --git a/payload_test.go b/payload_test.go index 099cca9..ef69339 100644 --- a/payload_test.go +++ b/payload_test.go @@ -40,7 +40,8 @@ func postAndReceive(t *testing.T, ts *httptest.Server, broker *Broker, contentTy } func TestServer_postFormDataStoredAsText(t *testing.T) { - ts, broker := newTestServer(nil) + ts, broker, cancel := newTestServer(nil) + defer cancel() defer ts.Close() event := postAndReceive(t, ts, broker, "application/x-www-form-urlencoded", "foo=bar&baz=qux") @@ -55,7 +56,8 @@ func TestServer_postFormDataStoredAsText(t *testing.T) { } func TestServer_postPlainTextStoredAsText(t *testing.T) { - ts, broker := newTestServer(nil) + ts, broker, cancel := newTestServer(nil) + defer cancel() defer ts.Close() event := postAndReceive(t, ts, broker, "text/plain", "hello world") @@ -70,7 +72,8 @@ func TestServer_postPlainTextStoredAsText(t *testing.T) { } func TestServer_postBinaryBase64Encoded(t *testing.T) { - ts, broker := newTestServer(nil) + ts, broker, cancel := newTestServer(nil) + defer cancel() defer ts.Close() event := postAndReceive(t, ts, broker, "application/octet-stream", "\x00\x01\x02\x03") @@ -86,7 +89,8 @@ func TestServer_postBinaryBase64Encoded(t *testing.T) { } func TestServer_postIncludesMethod(t *testing.T) { - ts, broker := newTestServer(nil) + ts, broker, cancel := newTestServer(nil) + defer cancel() defer ts.Close() event := postAndReceive(t, ts, broker, "application/json", `{"ok":true}`) @@ -97,7 +101,8 @@ func TestServer_postIncludesMethod(t *testing.T) { } func TestServer_postMalformedContentTypeFallsBackToBase64(t *testing.T) { - ts, broker := newTestServer(nil) + ts, broker, cancel := newTestServer(nil) + defer cancel() defer ts.Close() ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) @@ -161,7 +166,8 @@ func TestIsTextContent(t *testing.T) { } func TestServer_postJSONViaSSERoundTrip(t *testing.T) { - ts, _ := newTestServer(nil) + ts, _, cancel := newTestServer(nil) + defer cancel() defer ts.Close() ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) diff --git a/server.go b/server.go index ae6f03b..0906d5d 100644 --- a/server.go +++ b/server.go @@ -102,7 +102,11 @@ func (s *Server) handlePost(w http.ResponseWriter, r *http.Request) { Payload: payload, } - s.broker.Publish(event) + if err := s.broker.Publish(event); err != nil { + log.Printf("publish failed for %s: %v", path, err) + http.Error(w, "publish failed", http.StatusInternalServerError) + return + } log.Printf("published %s %s (%s) id=%s", r.Method, path, r.Header.Get("Content-Type"), event.ID) w.WriteHeader(http.StatusAccepted) } diff --git a/server_test.go b/server_test.go index 852dda0..dc6d134 100644 --- a/server_test.go +++ b/server_test.go @@ -1,9 +1,11 @@ package main import ( + "context" "crypto/hmac" "crypto/sha256" "encoding/hex" + "fmt" "net/http" "net/http/httptest" "strings" @@ -12,18 +14,22 @@ import ( "time" ) -func newTestServer(cfg *Configuration) (*httptest.Server, *Broker) { - broker := NewBroker(100) +func newTestServer(cfg *Configuration) (*httptest.Server, *Broker, context.CancelFunc) { + backend := NewMemoryBackend(100) + broker := NewBroker(backend) + ctx, cancel := context.WithCancel(context.Background()) + broker.Start(ctx) var cfgPtr atomic.Pointer[Configuration] if cfg != nil { cfgPtr.Store(cfg) } handler := NewServer(broker, &cfgPtr) - return httptest.NewServer(handler), broker + return httptest.NewServer(handler), broker, cancel } func TestServer_healthEndpoint(t *testing.T) { - ts, _ := newTestServer(nil) + ts, _, cancel := newTestServer(nil) + defer cancel() defer ts.Close() resp, err := http.Get(ts.URL + "/_health") @@ -42,7 +48,8 @@ func TestServer_healthEndpoint(t *testing.T) { } func TestServer_postPublishesEvent(t *testing.T) { - ts, _ := newTestServer(nil) + ts, _, cancel := newTestServer(nil) + defer cancel() defer ts.Close() resp, err := http.Post(ts.URL+"/test/topic", "application/json", strings.NewReader(`{"hello":"world"}`)) @@ -56,7 +63,8 @@ func TestServer_postPublishesEvent(t *testing.T) { } func TestServer_postTrailingSlashNormalized(t *testing.T) { - ts, broker := newTestServer(nil) + ts, broker, cancel := newTestServer(nil) + defer cancel() defer ts.Close() ch, unsub := broker.Subscribe("test/topic", "") @@ -84,7 +92,8 @@ func TestServer_postWithValidHMAC(t *testing.T) { }, }, } - ts, _ := newTestServer(cfg) + ts, _, cancel := newTestServer(cfg) + defer cancel() defer ts.Close() body := `{"action":"push"}` @@ -115,7 +124,8 @@ func TestServer_postWithInvalidSignature(t *testing.T) { }, }, } - ts, _ := newTestServer(cfg) + ts, _, cancel := newTestServer(cfg) + defer cancel() defer ts.Close() req, _ := http.NewRequest("POST", ts.URL+"/secure/path", strings.NewReader(`{"bad":"data"}`)) @@ -141,7 +151,8 @@ func TestServer_postToUnconfiguredPath(t *testing.T) { }, }, } - ts, _ := newTestServer(cfg) + ts, _, cancel := newTestServer(cfg) + defer cancel() defer ts.Close() resp, err := http.Post(ts.URL+"/open/path", "application/json", strings.NewReader(`{"ok":true}`)) @@ -155,7 +166,8 @@ func TestServer_postToUnconfiguredPath(t *testing.T) { } func TestServer_getWithoutSSEAccept(t *testing.T) { - ts, _ := newTestServer(nil) + ts, _, cancel := newTestServer(nil) + defer cancel() defer ts.Close() resp, err := http.Get(ts.URL + "/test/topic") @@ -169,7 +181,8 @@ func TestServer_getWithoutSSEAccept(t *testing.T) { } func TestServer_corsHeaders(t *testing.T) { - ts, _ := newTestServer(nil) + ts, _, cancel := newTestServer(nil) + defer cancel() defer ts.Close() resp, err := http.Post(ts.URL+"/test", "application/json", strings.NewReader(`{}`)) @@ -184,7 +197,8 @@ func TestServer_corsHeaders(t *testing.T) { } func TestServer_optionsPreflight(t *testing.T) { - ts, _ := newTestServer(nil) + ts, _, cancel := newTestServer(nil) + defer cancel() defer ts.Close() req, _ := http.NewRequest("OPTIONS", ts.URL+"/test", nil) @@ -202,7 +216,8 @@ func TestServer_optionsPreflight(t *testing.T) { } func TestServer_methodNotAllowed(t *testing.T) { - ts, _ := newTestServer(nil) + ts, _, cancel := newTestServer(nil) + defer cancel() defer ts.Close() req, _ := http.NewRequest("DELETE", ts.URL+"/test", nil) @@ -226,7 +241,8 @@ func TestServer_postWithBadVerifierConfig(t *testing.T) { }, }, } - ts, _ := newTestServer(cfg) + ts, _, cancel := newTestServer(cfg) + defer cancel() defer ts.Close() req, _ := http.NewRequest("POST", ts.URL+"/bad/path", strings.NewReader(`{}`)) @@ -242,7 +258,8 @@ func TestServer_postWithBadVerifierConfig(t *testing.T) { } func TestServer_postInvalidJSON(t *testing.T) { - ts, _ := newTestServer(nil) + ts, _, cancel := newTestServer(nil) + defer cancel() defer ts.Close() resp, err := http.Post(ts.URL+"/test", "application/json", strings.NewReader(`{not json`)) @@ -266,7 +283,11 @@ func (w *bareResponseWriter) Write(b []byte) (int, error) { return w.body.Write( func (w *bareResponseWriter) WriteHeader(code int) { w.code = code } func TestServer_sseWithoutFlusher(t *testing.T) { - broker := NewBroker(100) + backend := NewMemoryBackend(100) + broker := NewBroker(backend) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + broker.Start(ctx) var cfgPtr atomic.Pointer[Configuration] handler := NewServer(broker, &cfgPtr) @@ -291,7 +312,8 @@ func TestServer_postInheritsVerificationFromParent(t *testing.T) { }, }, } - ts, _ := newTestServer(cfg) + ts, _, cancel := newTestServer(cfg) + defer cancel() defer ts.Close() body := `{"action":"push"}` @@ -324,6 +346,33 @@ func TestServer_postInheritsVerificationFromParent(t *testing.T) { } } +type failingBackend struct{ MemoryBackend } + +func (f *failingBackend) Publish(*Event) error { + return fmt.Errorf("backend unavailable") +} + +func TestServer_postPublishError(t *testing.T) { + backend := &failingBackend{MemoryBackend: *NewMemoryBackend(100)} + broker := NewBroker(backend) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + broker.Start(ctx) + var cfgPtr atomic.Pointer[Configuration] + handler := NewServer(broker, &cfgPtr) + ts := httptest.NewServer(handler) + defer ts.Close() + + resp, err := http.Post(ts.URL+"/test", "application/json", strings.NewReader(`{}`)) + if err != nil { + t.Fatalf("POST failed: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusInternalServerError { + t.Errorf("expected 500, got %d", resp.StatusCode) + } +} + func TestServer_postWithMissingSignatureOnSecuredPath(t *testing.T) { cfg := &Configuration{ Paths: map[string]PathConfiguration{ @@ -334,7 +383,8 @@ func TestServer_postWithMissingSignatureOnSecuredPath(t *testing.T) { }, }, } - ts, _ := newTestServer(cfg) + ts, _, cancel := newTestServer(cfg) + defer cancel() defer ts.Close() resp, err := http.Post(ts.URL+"/secure/path", "application/json", strings.NewReader(`{}`)) diff --git a/sse_test.go b/sse_test.go index 6e5055b..8c4b319 100644 --- a/sse_test.go +++ b/sse_test.go @@ -46,11 +46,12 @@ func sseSubscribe(ctx context.Context, url string, headers map[string]string) <- } func TestServer_postAndSSEReceive(t *testing.T) { - ts, _ := newTestServer(nil) + ts, _, cancel := newTestServer(nil) + defer cancel() defer ts.Close() - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() + ctx, cancelSSE := context.WithCancel(context.Background()) + defer cancelSSE() events := sseSubscribe(ctx, ts.URL+"/test/topic", nil) time.Sleep(50 * time.Millisecond) @@ -73,7 +74,8 @@ func TestServer_sseWithValidBearerToken(t *testing.T) { "private/topic": {SubscribeSecret: "my-token"}, }, } - ts, _ := newTestServer(cfg) + ts, _, cancel := newTestServer(cfg) + defer cancel() defer ts.Close() req, _ := http.NewRequest("GET", ts.URL+"/private/topic", nil) @@ -97,7 +99,8 @@ func TestServer_sseWithWrongToken(t *testing.T) { "private/topic": {SubscribeSecret: "my-token"}, }, } - ts, _ := newTestServer(cfg) + ts, _, cancel := newTestServer(cfg) + defer cancel() defer ts.Close() req, _ := http.NewRequest("GET", ts.URL+"/private/topic", nil) @@ -114,7 +117,8 @@ func TestServer_sseWithWrongToken(t *testing.T) { } func TestServer_sseToOpenPath(t *testing.T) { - ts, _ := newTestServer(nil) + ts, _, cancel := newTestServer(nil) + defer cancel() defer ts.Close() req, _ := http.NewRequest("GET", ts.URL+"/open/topic", nil) @@ -136,11 +140,12 @@ func TestServer_sseToOpenPath(t *testing.T) { } func TestServer_prefixSubscription(t *testing.T) { - ts, _ := newTestServer(nil) + ts, _, cancel := newTestServer(nil) + defer cancel() defer ts.Close() - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() + ctx, cancelSSE := context.WithCancel(context.Background()) + defer cancelSSE() events := sseSubscribe(ctx, ts.URL+"/github.com/chrisguidry", nil) time.Sleep(50 * time.Millisecond) @@ -158,11 +163,12 @@ func TestServer_prefixSubscription(t *testing.T) { } func TestServer_prefixSubscriptionWithTrailingSlash(t *testing.T) { - ts, _ := newTestServer(nil) + ts, _, cancel := newTestServer(nil) + defer cancel() defer ts.Close() - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() + ctx, cancelSSE := context.WithCancel(context.Background()) + defer cancelSSE() events := sseSubscribe(ctx, ts.URL+"/test/", nil) time.Sleep(50 * time.Millisecond) @@ -180,7 +186,8 @@ func TestServer_prefixSubscriptionWithTrailingSlash(t *testing.T) { } func TestServer_lastEventIDReplay(t *testing.T) { - ts, broker := newTestServer(nil) + ts, broker, cancel := newTestServer(nil) + defer cancel() defer ts.Close() broker.Publish(&Event{ @@ -194,8 +201,8 @@ func TestServer_lastEventIDReplay(t *testing.T) { Payload: map[string]any{"n": 2}, }) - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() + ctx, cancelSSE := context.WithCancel(context.Background()) + defer cancelSSE() events := sseSubscribe(ctx, ts.URL+"/test/topic", map[string]string{ "Last-Event-ID": "replay-1", @@ -212,11 +219,12 @@ func TestServer_lastEventIDReplay(t *testing.T) { } func TestServer_filterQueryParam(t *testing.T) { - ts, _ := newTestServer(nil) + ts, _, cancel := newTestServer(nil) + defer cancel() defer ts.Close() - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() + ctx, cancelSSE := context.WithCancel(context.Background()) + defer cancelSSE() events := sseSubscribe(ctx, ts.URL+"/test/topic?filter=payload.ref:refs/heads/main", nil) time.Sleep(50 * time.Millisecond) @@ -242,7 +250,8 @@ func TestServer_sseWithMissingAuth(t *testing.T) { "private/topic": {SubscribeSecret: "my-token"}, }, } - ts, _ := newTestServer(cfg) + ts, _, cancel := newTestServer(cfg) + defer cancel() defer ts.Close() req, _ := http.NewRequest("GET", ts.URL+"/private/topic", nil) @@ -258,16 +267,20 @@ func TestServer_sseWithMissingAuth(t *testing.T) { } func TestServer_sseChannelClosed(t *testing.T) { - broker := NewBroker(100) + backend := NewMemoryBackend(100) + broker := NewBroker(backend) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + broker.Start(ctx) var cfgPtr atomic.Pointer[Configuration] handler := NewServer(broker, &cfgPtr) ts := httptest.NewServer(handler) defer ts.Close() - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() + sseCtx, sseCancel := context.WithCancel(context.Background()) + defer sseCancel() - events := sseSubscribe(ctx, ts.URL+"/test/topic", nil) + events := sseSubscribe(sseCtx, ts.URL+"/test/topic", nil) time.Sleep(50 * time.Millisecond) broker.mu.Lock() @@ -292,11 +305,12 @@ func TestServer_sseChannelClosed(t *testing.T) { } func TestServer_sseSkipsUnmarshalableEvent(t *testing.T) { - ts, broker := newTestServer(nil) + ts, broker, cancel := newTestServer(nil) + defer cancel() defer ts.Close() - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() + ctx, cancelSSE := context.WithCancel(context.Background()) + defer cancelSSE() events := sseSubscribe(ctx, ts.URL+"/test/topic", nil) time.Sleep(50 * time.Millisecond) @@ -323,11 +337,12 @@ func TestServer_sseSkipsUnmarshalableEvent(t *testing.T) { } func TestServer_textPayloadStoredAsString(t *testing.T) { - ts, _ := newTestServer(nil) + ts, _, cancel := newTestServer(nil) + defer cancel() defer ts.Close() - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() + ctx, cancelSSE := context.WithCancel(context.Background()) + defer cancelSSE() done := make(chan string, 1) go func() {