From 4f2ae26d58eb1f53df22a0f8f906b14bd74529ae Mon Sep 17 00:00:00 2001 From: Angel Beltran Date: Tue, 20 Jan 2026 21:11:02 -0700 Subject: [PATCH] filter ss events by account --- internal/server/api/sse/router.go | 2 +- internal/server/middleware/events.go | 9 ++++- internal/server/sse/sse.go | 51 ++++++++++++++++++---------- 3 files changed, 42 insertions(+), 20 deletions(-) diff --git a/internal/server/api/sse/router.go b/internal/server/api/sse/router.go index 7e45528..7d08ed6 100644 --- a/internal/server/api/sse/router.go +++ b/internal/server/api/sse/router.go @@ -61,7 +61,7 @@ func (r *sseRouter) serveEvents(c *gin.Context) (response.Response, error) { hdr.Set("Cache-Control", "no-cache") w.Flush() - err := r.sse.Listen(ctx, func(ctx context.Context, e *sse.Event) error { + err := r.sse.Listen(ctx, acctID, func(ctx context.Context, e *sse.Event) error { log.Debugf("sending event of type %s", e.Type) e.Write(w) diff --git a/internal/server/middleware/events.go b/internal/server/middleware/events.go index cd288c3..5940742 100644 --- a/internal/server/middleware/events.go +++ b/internal/server/middleware/events.go @@ -2,6 +2,7 @@ package middleware import ( "context" + "fmt" "path" "ruben/inventory2/internal/logging" "ruben/inventory2/internal/server/sse" @@ -40,6 +41,7 @@ func (p *UpdateNotificationPublisher) Group(pathPattern string) *UpdateNotificat return &p2 } +// Publish must be applied AFTER a middleware puts in the Identity into the gin.Context. func (p *UpdateNotificationPublisher) Publish(pathPattern string) gin.HandlerFunc { trimBasePathSegs := getPathSegments(p.trimBasePath) trimmedBasePathPattern := path.Join(p.basePathPattern, pathPattern) @@ -85,12 +87,17 @@ func (p *UpdateNotificationPublisher) Publish(pathPattern string) gin.HandlerFun parentEvent = e } + acctID := GetIdentity(c).Account.AccountID for _, e := range events { go func() { ctx, cancel := context.WithTimeout(context.Background(), time.Minute) defer cancel() - if err := p.sse.Send(ctx, e, nil); err != nil { + if err := p.sse.Send(ctx, sse.Event{ + AccountID: acctID, + Type: e, + Data: []byte(fmt.Sprintf(`{"eventType": %q}`, e)), + }); err != nil { p.log.Errorf("failed to send sse event to listener: %v", err) } }() diff --git a/internal/server/sse/sse.go b/internal/server/sse/sse.go index f0efd56..6e0ac5e 100644 --- a/internal/server/sse/sse.go +++ b/internal/server/sse/sse.go @@ -4,6 +4,7 @@ import ( "bytes" "context" "fmt" + "maps" "net/http" "sync" ) @@ -11,18 +12,18 @@ import ( type ( Queue struct { in chan Event - out map[int]chan Event + out map[int64]map[int]chan Event ctx context.Context cancel context.CancelFunc - lock sync.Mutex - prevID int + lock sync.Mutex } Event struct { - Type string - Data []byte + AccountID int64 + Type string + Data []byte } ) @@ -30,7 +31,7 @@ func NewQueue() *Queue { ctx, cancel := context.WithCancel(context.Background()) return &Queue{ in: make(chan Event), - out: make(map[int]chan Event), + out: make(map[int64]map[int]chan Event), ctx: ctx, cancel: cancel, @@ -46,9 +47,13 @@ func (q *Queue) Start(ctx context.Context) error { // we're done piping events return nil case e := <-q.in: - // share event will all subscribers + if e.AccountID == 0 { + continue + } + + // share event with listeners on the account q.lock.Lock() - for _, out := range q.out { + for _, out := range q.out[e.AccountID] { out <- e } q.lock.Unlock() @@ -56,21 +61,34 @@ func (q *Queue) Start(ctx context.Context) error { } } -func (q *Queue) Listen(ctx context.Context, fn func(context.Context, *Event) error) error { +func (q *Queue) Listen(ctx context.Context, acctID int64, fn func(context.Context, *Event) error) error { // create new out pipe and append it to the queue out := make(chan Event, 1) q.lock.Lock() - q.prevID += 1 - id := q.prevID - q.out[id] = out + outs := q.out[acctID] + if outs == nil { + outs = make(map[int]chan Event) + q.out[acctID] = outs + } + + var maxID int + for n := range maps.Keys(outs) { + maxID = max(maxID, n) + } + id := maxID + 1 + + outs[id] = out q.lock.Unlock() // delete the pipe when done listening defer func() { q.lock.Lock() - delete(q.out, id) + delete(q.out[acctID], id) + if len(q.out[acctID]) == 0 { + delete(q.out, acctID) + } q.lock.Unlock() }() @@ -91,7 +109,7 @@ func (q *Queue) Listen(ctx context.Context, fn func(context.Context, *Event) err } } -func (q *Queue) Send(ctx context.Context, eventType string, data []byte) error { +func (q *Queue) Send(ctx context.Context, e Event) error { select { case <-q.ctx.Done(): // the queue has closed @@ -100,10 +118,7 @@ func (q *Queue) Send(ctx context.Context, eventType string, data []byte) error { // sender ran out of time return fmt.Errorf("provided context canceled: %w", ctx.Err()) // push the event onto the queue - case q.in <- Event{ - Type: eventType, - Data: data, - }: + case q.in <- e: } return nil