From 2e8a49ca5ae13605ca94d7e9c31487461e1c072b Mon Sep 17 00:00:00 2001 From: QiuSW <105186638@qq.com> Date: Mon, 5 Oct 2026 16:12:39 +0800 Subject: [PATCH] feat(#1): add agent-key authenticated SSE stream /api/agent-events Per-agent SSE (metadata-only new_message events, id=message_id, 30s heartbeat, Last-Event-ID replay capped at 200 with resync_required, max 5 connections per agent, slow clients dropped). Mounted behind the agent API-key middleware. /api/events is unchanged. Co-Authored-By: Claude Sonnet 5.5 --- cmd/synapbus/main.go | 8 + internal/api/agent_events.go | 270 +++++++++++++++++ internal/api/agent_events_test.go | 489 ++++++++++++++++++++++++++++++ internal/api/broadcaster.go | 53 ++++ internal/api/sse_handler.go | 10 +- internal/messaging/event_meta.go | 71 +++++ internal/messaging/store.go | 1 + 7 files changed, 901 insertions(+), 1 deletion(-) create mode 100644 internal/api/agent_events.go create mode 100644 internal/api/agent_events_test.go create mode 100644 internal/messaging/event_meta.go diff --git a/cmd/synapbus/main.go b/cmd/synapbus/main.go index a95f16c..5518a82 100644 --- a/cmd/synapbus/main.go +++ b/cmd/synapbus/main.go @@ -841,6 +841,14 @@ func runServe(cmd *cobra.Command, args []string) error { // Register broadcaster as a message listener so SSE events fire // for messages sent via MCP (agents) as well as the REST API. msgService.AddMessageListener(sseBroadcaster) + sseBroadcaster.SetMessageService(msgService) + + // Per-agent SSE stream (agent API key auth, not session auth). + sseHub.SetAgentEventBacklog(msgService) + r.Group(func(r chi.Router) { + r.Use(agents.RequiredAuthMiddlewareWithOAuth(agentService, apiKeyService, oauthProvider)) + r.Get("/api/agent-events", sseHub.HandleAgentEvents) + }) // Initialize push notification service pushStore := push.NewSQLiteStore(db.DB) diff --git a/internal/api/agent_events.go b/internal/api/agent_events.go new file mode 100644 index 0000000..159ac62 --- /dev/null +++ b/internal/api/agent_events.go @@ -0,0 +1,270 @@ +package api + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "strconv" + "time" + + "github.com/synapbus/synapbus/internal/agents" + "github.com/synapbus/synapbus/internal/messaging" +) + +const ( + // maxAgentConnections is the maximum number of simultaneous + // /api/agent-events connections per agent. + maxAgentConnections = 5 + // maxAgentBacklog is the maximum number of missed events replayed on + // reconnect; beyond this a single resync_required event is sent instead. + maxAgentBacklog = 200 + // agentSubBuffer is the per-connection live event buffer. A connection + // whose buffer is full is dropped rather than blocking the broadcaster. + agentSubBuffer = 64 + // agentWriteTimeout bounds a single write to a slow client. + agentWriteTimeout = 10 * time.Second +) + +// AgentEventBacklog loads body-free metadata for messages visible to an agent, +// used to replay events after a reconnect with Last-Event-ID. +type AgentEventBacklog interface { + ListEventMetaAfter(ctx context.Context, agentName string, afterID int64, limit int) ([]*messaging.EventMeta, error) +} + +// AgentMessageEvent is the payload of a new_message event on the agent stream. +// It carries metadata only, never the message body. +type AgentMessageEvent struct { + MessageID int64 `json:"message_id"` + Channel string `json:"channel,omitempty"` + FromAgent string `json:"from_agent,omitempty"` + ToAgent string `json:"to_agent,omitempty"` + Subject string `json:"subject,omitempty"` +} + +// agentSub is one connected /api/agent-events client. +type agentSub struct { + ch chan AgentMessageEvent +} + +// SetAgentEventBacklog configures the source used to replay missed events. +func (h *SSEHub) SetAgentEventBacklog(b AgentEventBacklog) { + h.mu.Lock() + defer h.mu.Unlock() + h.agentBacklog = b +} + +// BroadcastAgentMessage delivers a new_message event to every connection of the +// named agents. Connections that cannot keep up are dropped. +func (h *SSEHub) BroadcastAgentMessage(agentNames []string, ev AgentMessageEvent) { + type slowSub struct { + name string + sub *agentSub + } + var slow []slowSub + + h.mu.RLock() + for _, name := range agentNames { + for sub := range h.agentSubs[name] { + select { + case sub.ch <- ev: + default: + slow = append(slow, slowSub{name, sub}) + } + } + } + h.mu.RUnlock() + + for _, s := range slow { + h.logger.Warn("dropping slow agent SSE client", "agent", s.name) + h.removeAgentSub(s.name, s.sub) + } +} + +// addAgentSub registers a connection for the agent. It returns nil when the +// agent already has maxAgentConnections open connections. +func (h *SSEHub) addAgentSub(agentName string) *agentSub { + h.mu.Lock() + defer h.mu.Unlock() + + set := h.agentSubs[agentName] + if len(set) >= maxAgentConnections { + return nil + } + if set == nil { + set = make(map[*agentSub]struct{}) + h.agentSubs[agentName] = set + } + sub := &agentSub{ch: make(chan AgentMessageEvent, agentSubBuffer)} + set[sub] = struct{}{} + h.logger.Info("agent SSE client connected", "agent", agentName) + return sub +} + +func (h *SSEHub) removeAgentSub(agentName string, sub *agentSub) { + h.mu.Lock() + defer h.mu.Unlock() + + set := h.agentSubs[agentName] + if _, ok := set[sub]; ok { + delete(set, sub) + close(sub.ch) + h.logger.Info("agent SSE client disconnected", "agent", agentName) + } + if len(set) == 0 { + delete(h.agentSubs, agentName) + } +} + +// closeAgentSubsLocked disconnects all agent subscriptions. Caller holds h.mu. +func (h *SSEHub) closeAgentSubsLocked() { + for name, set := range h.agentSubs { + for sub := range set { + close(sub.ch) + } + delete(h.agentSubs, name) + } +} + +// HandleAgentEvents handles GET /api/agent-events. It must be mounted behind +// the agent API-key middleware; the authenticated agent is read from the +// request context. Only metadata for messages visible to that agent is sent. +func (h *SSEHub) HandleAgentEvents(w http.ResponseWriter, r *http.Request) { + agent, ok := agents.AgentFromContext(r.Context()) + if !ok || agent == nil { + http.Error(w, `{"error":"unauthorized","message":"An agent API key is required"}`, http.StatusUnauthorized) + return + } + + flusher, ok := w.(http.Flusher) + if !ok { + http.Error(w, "Streaming not supported", http.StatusInternalServerError) + return + } + + var lastID int64 + hasLast := false + if v := r.Header.Get("Last-Event-ID"); v != "" { + id, err := strconv.ParseInt(v, 10, 64) + if err != nil || id < 0 { + http.Error(w, `{"error":"bad_request","message":"Last-Event-ID must be a message id"}`, http.StatusBadRequest) + return + } + lastID, hasLast = id, true + } + + // Subscribe before replaying so no message falls between the replay query + // and the live stream; duplicates are removed by id below. + sub := h.addAgentSub(agent.Name) + if sub == nil { + http.Error(w, fmt.Sprintf(`{"error":"too_many_connections","message":"Agent %q already has the maximum of %d event stream connections"}`, agent.Name, maxAgentConnections), http.StatusTooManyRequests) + return + } + defer h.removeAgentSub(agent.Name, sub) + + w.Header().Set("Content-Type", "text/event-stream") + w.Header().Set("Cache-Control", "no-cache") + w.Header().Set("Connection", "keep-alive") + w.Header().Set("X-Accel-Buffering", "no") + + rc := http.NewResponseController(w) + write := func(id int64, eventType string, data any) bool { + _ = rc.SetWriteDeadline(time.Now().Add(agentWriteTimeout)) + return writeSSEEvent(w, flusher, id, eventType, data) + } + + if !write(0, "connected", map[string]any{ + "agent": agent.Name, + "timestamp": time.Now().Format(time.RFC3339), + }) { + return + } + + // sentUpTo is the highest message id already delivered; live events at or + // below it are duplicates of the replay. + sentUpTo := lastID + if hasLast { + h.mu.RLock() + backlog := h.agentBacklog + h.mu.RUnlock() + + if backlog != nil { + metas, err := backlog.ListEventMetaAfter(r.Context(), agent.Name, lastID, maxAgentBacklog+1) + if err != nil || len(metas) > maxAgentBacklog { + if err != nil { + h.logger.Warn("agent SSE backlog failed", "agent", agent.Name, "error", err) + } + if !write(0, "resync_required", map[string]any{ + "after_id": lastID, + "limit": maxAgentBacklog, + }) { + return + } + } else { + for _, m := range metas { + ev := AgentMessageEvent{MessageID: m.MessageID, Subject: m.Subject} + if m.Channel != "" { + ev.Channel = m.Channel + } else { + ev.FromAgent, ev.ToAgent = m.FromAgent, m.ToAgent + } + if !write(m.MessageID, "new_message", ev) { + return + } + sentUpTo = m.MessageID + } + } + } + } + + interval := h.heartbeat + if interval <= 0 { + interval = 30 * time.Second + } + ticker := time.NewTicker(interval) + defer ticker.Stop() + + ctx := r.Context() + for { + select { + case <-ctx.Done(): + return + case ev, ok := <-sub.ch: + if !ok { + return + } + if ev.MessageID <= sentUpTo { + continue + } + if !write(ev.MessageID, "new_message", ev) { + return + } + sentUpTo = ev.MessageID + case <-ticker.C: + if !write(0, "heartbeat", map[string]any{ + "timestamp": time.Now().Format(time.RFC3339), + }) { + return + } + } + } +} + +// writeSSEEvent writes one SSE frame. A positive id is emitted as the "id:" +// line. It reports whether the write succeeded. +func writeSSEEvent(w http.ResponseWriter, flusher http.Flusher, id int64, eventType string, data any) bool { + jsonData, err := json.Marshal(data) + if err != nil { + return false + } + if id > 0 { + if _, err := fmt.Fprintf(w, "id: %d\n", id); err != nil { + return false + } + } + if _, err := fmt.Fprintf(w, "event: %s\ndata: %s\n\n", eventType, jsonData); err != nil { + return false + } + flusher.Flush() + return true +} diff --git a/internal/api/agent_events_test.go b/internal/api/agent_events_test.go new file mode 100644 index 0000000..ca71685 --- /dev/null +++ b/internal/api/agent_events_test.go @@ -0,0 +1,489 @@ +package api + +import ( + "bufio" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "strconv" + "strings" + "testing" + "time" + + "github.com/go-chi/chi/v5" + _ "modernc.org/sqlite" + + "github.com/synapbus/synapbus/internal/agents" + "github.com/synapbus/synapbus/internal/channels" + "github.com/synapbus/synapbus/internal/messaging" +) + +// agentEventsEnv wires real services over a temporary in-memory SQLite +// database behind the real agent API-key middleware. +type agentEventsEnv struct { + t *testing.T + srv *httptest.Server + hub *SSEHub + msgSvc *messaging.MessagingService + chanSvc *channels.Service + keys map[string]string // agent name -> API key + channels map[string]int64 // channel name -> id +} + +func newAgentEventsEnv(t *testing.T, agentNames ...string) *agentEventsEnv { + t.Helper() + db := newTestDBFull(t) + + agentSvc := agents.NewAgentService(agents.NewSQLiteAgentStore(db), nil) + msgSvc := messaging.NewMessagingService(messaging.NewSQLiteMessageStore(db), nil) + chanSvc := channels.NewService(channels.NewSQLiteChannelStore(db), msgSvc, nil) + + hub := NewSSEHub() + hub.SetAgentEventBacklog(msgSvc) + b := NewSSEBroadcaster(hub, agentSvc, chanSvc) + b.SetMessageService(msgSvc) + msgSvc.AddMessageListener(b) + + env := &agentEventsEnv{ + t: t, hub: hub, msgSvc: msgSvc, chanSvc: chanSvc, + keys: map[string]string{}, channels: map[string]int64{}, + } + ctx := context.Background() + for _, name := range agentNames { + _, key, err := agentSvc.Register(ctx, name, name, "ai", nil, 1) + if err != nil { + t.Fatalf("register %s: %v", name, err) + } + env.keys[name] = key + } + + r := chi.NewRouter() + r.Group(func(r chi.Router) { + r.Use(agents.RequiredAuthMiddlewareWithOAuth(agentSvc, nil, nil)) + r.Get("/api/agent-events", hub.HandleAgentEvents) + }) + env.srv = httptest.NewServer(r) + t.Cleanup(env.srv.Close) + t.Cleanup(hub.Close) // runs before srv.Close so open streams end + return env +} + +func (e *agentEventsEnv) makeChannel(name string, members ...string) { + e.t.Helper() + ctx := context.Background() + ch, err := e.chanSvc.CreateChannel(ctx, channels.CreateChannelRequest{Name: name, CreatedBy: members[0]}) + if err != nil { + e.t.Fatalf("create channel: %v", err) + } + e.channels[name] = ch.ID + for _, m := range members { + if err := e.chanSvc.JoinChannel(ctx, ch.ID, m); err != nil { + e.t.Fatalf("join %s: %v", m, err) + } + } +} + +func (e *agentEventsEnv) dm(from, to, body, subject string) int64 { + e.t.Helper() + m, err := e.msgSvc.SendMessage(context.Background(), from, to, body, messaging.SendOptions{Subject: subject}) + if err != nil { + e.t.Fatalf("send dm: %v", err) + } + return m.ID +} + +func (e *agentEventsEnv) post(from, channel, body string) int64 { + e.t.Helper() + id := e.channels[channel] + m, err := e.msgSvc.SendMessage(context.Background(), from, "", body, messaging.SendOptions{ChannelID: &id}) + if err != nil { + e.t.Fatalf("post: %v", err) + } + return m.ID +} + +// sseFrame is one parsed SSE frame. +type sseFrame struct { + ID string + Event string + Data string +} + +type sseClient struct { + t *testing.T + frames chan sseFrame + cancel context.CancelFunc + status int + body string +} + +// connect opens the stream. For non-200 statuses it records status/body only. +func (e *agentEventsEnv) connect(key, lastEventID string) *sseClient { + e.t.Helper() + ctx, cancel := context.WithCancel(context.Background()) + req, _ := http.NewRequestWithContext(ctx, "GET", e.srv.URL+"/api/agent-events", nil) + if key != "" { + req.Header.Set("Authorization", "Bearer "+key) + } + if lastEventID != "" { + req.Header.Set("Last-Event-ID", lastEventID) + } + resp, err := http.DefaultClient.Do(req) + if err != nil { + cancel() + e.t.Fatalf("connect: %v", err) + } + c := &sseClient{t: e.t, frames: make(chan sseFrame, 512), cancel: cancel, status: resp.StatusCode} + e.t.Cleanup(cancel) + if resp.StatusCode != http.StatusOK { + var sb strings.Builder + sc := bufio.NewScanner(resp.Body) + for sc.Scan() { + sb.WriteString(sc.Text()) + } + resp.Body.Close() + c.body = sb.String() + return c + } + go func() { + defer resp.Body.Close() + defer close(c.frames) + sc := bufio.NewScanner(resp.Body) + var f sseFrame + for sc.Scan() { + line := sc.Text() + switch { + case line == "": + if f.Event != "" { + c.frames <- f + } + f = sseFrame{} + case strings.HasPrefix(line, "id: "): + f.ID = line[4:] + case strings.HasPrefix(line, "event: "): + f.Event = line[7:] + case strings.HasPrefix(line, "data: "): + f.Data = line[6:] + } + } + }() + return c +} + +func (c *sseClient) next() sseFrame { + c.t.Helper() + select { + case f, ok := <-c.frames: + if !ok { + c.t.Fatal("stream closed while waiting for event") + } + return f + case <-time.After(5 * time.Second): + c.t.Fatal("timed out waiting for SSE event") + } + return sseFrame{} +} + +// nextMessage skips nothing: it expects the next frame to be new_message. +func (c *sseClient) nextMessage() (sseFrame, map[string]any) { + c.t.Helper() + f := c.next() + if f.Event != "new_message" { + c.t.Fatalf("event = %q (data %s), want new_message", f.Event, f.Data) + } + var m map[string]any + if err := json.Unmarshal([]byte(f.Data), &m); err != nil { + c.t.Fatalf("bad json %q: %v", f.Data, err) + } + return f, m +} + +func (c *sseClient) expectConnected(agent string) { + c.t.Helper() + f := c.next() + if f.Event != "connected" || !strings.Contains(f.Data, `"agent":"`+agent+`"`) { + c.t.Fatalf("first frame = %+v, want connected for %s", f, agent) + } +} + +func TestAgentEvents_Auth(t *testing.T) { + env := newAgentEventsEnv(t, "alice") + tests := []struct { + name string + key string + }{ + {"missing key", ""}, + {"invalid key", "not-a-real-key"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + c := env.connect(tt.key, "") + if c.status != http.StatusUnauthorized { + t.Fatalf("status = %d, want 401", c.status) + } + }) + } + + t.Run("valid key gets connected event", func(t *testing.T) { + c := env.connect(env.keys["alice"], "") + if c.status != http.StatusOK { + t.Fatalf("status = %d, want 200", c.status) + } + c.expectConnected("alice") + }) +} + +func TestAgentEvents_HandlerWithoutAgentContext(t *testing.T) { + hub := NewSSEHub() + rr := httptest.NewRecorder() + hub.HandleAgentEvents(rr, httptest.NewRequest("GET", "/api/agent-events", nil)) + if rr.Code != http.StatusUnauthorized { + t.Fatalf("status = %d, want 401", rr.Code) + } +} + +func TestAgentEvents_Delivery(t *testing.T) { + const secret = "TOP-SECRET-BODY-TEXT" + + tests := []struct { + name string + // send performs the action; returns the message id to expect for "bob". + send func(env *agentEventsEnv) int64 + want map[string]any // expected fields (besides message_id) + absent []string // fields that must not be present + viewers []string + }{ + { + name: "dm to bob", + send: func(env *agentEventsEnv) int64 { return env.dm("alice", "bob", secret, "hello subject") }, + want: map[string]any{"from_agent": "alice", "to_agent": "bob", "subject": "hello subject"}, + absent: []string{"channel", "body"}, + }, + { + name: "channel message to member", + send: func(env *agentEventsEnv) int64 { + env.makeChannel("room", "alice", "bob") + return env.post("alice", "room", secret) + }, + want: map[string]any{"channel": "room"}, + absent: []string{"to_agent", "from_agent", "body"}, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + env := newAgentEventsEnv(t, "alice", "bob") + bob := env.connect(env.keys["bob"], "") + bob.expectConnected("bob") + alice := env.connect(env.keys["alice"], "") + alice.expectConnected("alice") + + id := tt.send(env) + + f, m := bob.nextMessage() + if f.ID != strconv.FormatInt(id, 10) { + t.Errorf("id line = %q, want %d", f.ID, id) + } + if int64(m["message_id"].(float64)) != id { + t.Errorf("message_id = %v, want %d", m["message_id"], id) + } + for k, v := range tt.want { + if m[k] != v { + t.Errorf("%s = %v, want %v", k, m[k], v) + } + } + for _, k := range tt.absent { + if _, ok := m[k]; ok { + t.Errorf("field %q must not be present: %s", k, f.Data) + } + } + if strings.Contains(f.Data, secret) { + t.Errorf("event leaks body: %s", f.Data) + } + + // Sender must not get its own message: send a canary to alice + // and make sure it is the next thing she sees. + canary := env.dm("bob", "alice", "canary", "") + f2, _ := alice.nextMessage() + if f2.ID != strconv.FormatInt(canary, 10) { + t.Errorf("alice's next event id = %s, want canary %d (own message leaked?)", f2.ID, canary) + } + }) + } +} + +func TestAgentEvents_Isolation(t *testing.T) { + env := newAgentEventsEnv(t, "alice", "bob", "carol") + env.makeChannel("private-ish", "bob", "carol") // alice is NOT a member + env.makeChannel("shared", "alice", "bob") + + alice := env.connect(env.keys["alice"], "") + alice.expectConnected("alice") + + // None of these may reach alice. + env.dm("bob", "carol", "b->c", "") // DM between others + env.post("bob", "private-ish", "not for you") // channel alice has not joined + env.dm("alice", "bob", "alice->bob", "") // her own message + + // Canary she is allowed to see; it must be the very next event. + canary := env.post("bob", "shared", "visible") + f, m := alice.nextMessage() + if f.ID != strconv.FormatInt(canary, 10) || m["channel"] != "shared" { + t.Fatalf("alice got %+v, want only the canary in 'shared'", f) + } +} + +func TestAgentEvents_Resume(t *testing.T) { + t.Run("replays visible newer messages once, in order", func(t *testing.T) { + env := newAgentEventsEnv(t, "alice", "bob", "carol") + env.makeChannel("room", "alice", "bob") + env.makeChannel("other", "bob", "carol") + + first := env.dm("bob", "alice", "old", "") // client already has this + m2 := env.dm("bob", "alice", "missed dm", "s") + env.dm("bob", "carol", "not alice's", "") + env.post("bob", "other", "not a member") + m5 := env.post("bob", "room", "missed channel") + env.dm("alice", "bob", "own", "") // own message excluded + + c := env.connect(env.keys["alice"], strconv.FormatInt(first, 10)) + c.expectConnected("alice") + + f, m := c.nextMessage() + if f.ID != strconv.FormatInt(m2, 10) || m["subject"] != "s" { + t.Fatalf("first replay = %+v", f) + } + f, m = c.nextMessage() + if f.ID != strconv.FormatInt(m5, 10) || m["channel"] != "room" { + t.Fatalf("second replay = %+v", f) + } + + // Live message after replay arrives exactly once. + live := env.dm("bob", "alice", "live", "") + f, _ = c.nextMessage() + if f.ID != strconv.FormatInt(live, 10) { + t.Fatalf("live event id = %s, want %d (duplicate or missing?)", f.ID, live) + } + }) + + t.Run("exactly 200 are replayed", func(t *testing.T) { + env := newAgentEventsEnv(t, "alice", "bob") + var ids []int64 + for i := 0; i < maxAgentBacklog; i++ { + ids = append(ids, env.dm("bob", "alice", "x", "")) + } + c := env.connect(env.keys["alice"], "0") + c.expectConnected("alice") + for i, id := range ids { + f, _ := c.nextMessage() + if f.ID != strconv.FormatInt(id, 10) { + t.Fatalf("replay #%d id = %s, want %d", i, f.ID, id) + } + } + }) + + t.Run("more than 200 sends resync_required only", func(t *testing.T) { + env := newAgentEventsEnv(t, "alice", "bob") + for i := 0; i < maxAgentBacklog+1; i++ { + env.dm("bob", "alice", "x", "") + } + c := env.connect(env.keys["alice"], "0") + c.expectConnected("alice") + f := c.next() + if f.Event != "resync_required" { + t.Fatalf("event = %q, want resync_required", f.Event) + } + // No replay; the next event is the live one. + live := env.dm("bob", "alice", "live", "") + lf, _ := c.nextMessage() + if lf.ID != strconv.FormatInt(live, 10) { + t.Fatalf("after resync got id %s, want live %d", lf.ID, live) + } + }) + + t.Run("invalid Last-Event-ID is rejected", func(t *testing.T) { + env := newAgentEventsEnv(t, "alice") + c := env.connect(env.keys["alice"], "abc") + if c.status != http.StatusBadRequest { + t.Fatalf("status = %d, want 400", c.status) + } + }) +} + +func TestAgentEvents_ConnectionLimit(t *testing.T) { + env := newAgentEventsEnv(t, "alice", "bob") + + var conns []*sseClient + for i := 0; i < maxAgentConnections; i++ { + c := env.connect(env.keys["alice"], "") + if c.status != http.StatusOK { + t.Fatalf("connection %d status = %d, want 200", i+1, c.status) + } + c.expectConnected("alice") + conns = append(conns, c) + } + + over := env.connect(env.keys["alice"], "") + if over.status != http.StatusTooManyRequests { + t.Fatalf("6th connection status = %d, want 429", over.status) + } + if !strings.Contains(over.body, "too_many_connections") { + t.Errorf("429 body = %q, want explicit error", over.body) + } + + // The limit is per agent. + other := env.connect(env.keys["bob"], "") + if other.status != http.StatusOK { + t.Fatalf("other agent status = %d, want 200", other.status) + } + + // Closing a connection frees a slot. + conns[0].cancel() + deadline := time.Now().Add(5 * time.Second) + for { + c := env.connect(env.keys["alice"], "") + if c.status == http.StatusOK { + break + } + if time.Now().After(deadline) { + t.Fatal("slot was not released after disconnect") + } + time.Sleep(20 * time.Millisecond) + } +} + +func TestAgentEvents_SlowClientDroppedWithoutBlocking(t *testing.T) { + hub := NewSSEHub() + sub := hub.addAgentSub("alice") + if sub == nil { + t.Fatal("addAgentSub returned nil") + } + done := make(chan struct{}) + go func() { + defer close(done) + for i := 1; i <= agentSubBuffer+10; i++ { // nobody reads sub.ch + hub.BroadcastAgentMessage([]string{"alice"}, AgentMessageEvent{MessageID: int64(i)}) + } + }() + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("broadcast blocked on a slow client") + } + hub.mu.RLock() + n := len(hub.agentSubs["alice"]) + hub.mu.RUnlock() + if n != 0 { + t.Fatalf("slow subscription still registered (%d)", n) + } +} + +func TestAgentEvents_Heartbeat(t *testing.T) { + env := newAgentEventsEnv(t, "alice") + env.hub.heartbeat = 50 * time.Millisecond + c := env.connect(env.keys["alice"], "") + c.expectConnected("alice") + if f := c.next(); f.Event != "heartbeat" { + t.Fatalf("event = %q, want heartbeat", f.Event) + } +} diff --git a/internal/api/broadcaster.go b/internal/api/broadcaster.go index 683015f..b897351 100644 --- a/internal/api/broadcaster.go +++ b/internal/api/broadcaster.go @@ -35,6 +35,7 @@ type SSEBroadcaster struct { hub *SSEHub agentService *agents.AgentService channelService *channels.Service + msgService *messaging.MessagingService // optional: resolves conversation subjects logger *slog.Logger } @@ -48,6 +49,12 @@ func NewSSEBroadcaster(hub *SSEHub, agentService *agents.AgentService, channelSe } } +// SetMessageService sets the messaging service used to resolve conversation +// subjects for agent events. Optional; without it events carry no subject. +func (b *SSEBroadcaster) SetMessageService(svc *messaging.MessagingService) { + b.msgService = svc +} + // BroadcastNewMessage sends a new_message event to the given owner. func (b *SSEBroadcaster) BroadcastNewMessage(_ context.Context, ownerID int64, event NewMessageEvent) { b.hub.Broadcast(ownerID, SSEEvent{ @@ -128,4 +135,50 @@ func (b *SSEBroadcaster) OnMessageSent(ctx context.Context, msg *messaging.Messa } else { b.BroadcastDM(ctx, event) } + + b.broadcastToAgents(ctx, msg) +} + +// broadcastToAgents pushes a body-free new_message event to the per-agent +// stream (GET /api/agent-events): the DM recipient, or the members of the +// channel at send time. The sender is never notified of its own message. +func (b *SSEBroadcaster) broadcastToAgents(ctx context.Context, msg *messaging.Message) { + var recipients []string + if msg.ChannelID != nil { + if b.channelService == nil { + return + } + members, err := b.channelService.GetMembers(ctx, *msg.ChannelID) + if err != nil { + b.logger.Debug("could not get channel members for agent SSE broadcast", + "channel_id", *msg.ChannelID, "error", err) + return + } + for _, m := range members { + if m.AgentName != msg.FromAgent { + recipients = append(recipients, m.AgentName) + } + } + } else if msg.ToAgent != "" && msg.ToAgent != msg.FromAgent { + recipients = []string{msg.ToAgent} + } + if len(recipients) == 0 { + return + } + + ev := AgentMessageEvent{ + MessageID: msg.ID, + FromAgent: msg.FromAgent, + ToAgent: msg.ToAgent, + } + if msg.ChannelID != nil { + ev.FromAgent = "" + if ch, err := b.channelService.GetChannel(ctx, *msg.ChannelID); err == nil { + ev.Channel = ch.Name + } + } + if b.msgService != nil { + ev.Subject = b.msgService.GetConversationSubject(ctx, msg.ConversationID) + } + b.hub.BroadcastAgentMessage(recipients, ev) } diff --git a/internal/api/sse_handler.go b/internal/api/sse_handler.go index 58a4d9b..91bdce2 100644 --- a/internal/api/sse_handler.go +++ b/internal/api/sse_handler.go @@ -19,6 +19,11 @@ type SSEEvent struct { type SSEHub struct { mu sync.RWMutex clients map[int64]map[chan SSEEvent]struct{} // ownerID -> set of channels + + // Per-agent subscriptions for GET /api/agent-events (see agent_events.go). + agentSubs map[string]map[*agentSub]struct{} + agentBacklog AgentEventBacklog + heartbeat time.Duration nextID int64 logger *slog.Logger } @@ -26,7 +31,9 @@ type SSEHub struct { // NewSSEHub creates a new SSE hub. func NewSSEHub() *SSEHub { return &SSEHub{ - clients: make(map[int64]map[chan SSEEvent]struct{}), + clients: make(map[int64]map[chan SSEEvent]struct{}), + agentSubs: make(map[string]map[*agentSub]struct{}), + heartbeat: 30 * time.Second, logger: slog.Default().With("component", "api.sse"), } } @@ -76,6 +83,7 @@ func (h *SSEHub) Close() { } delete(h.clients, ownerID) } + h.closeAgentSubsLocked() h.logger.Info("all SSE clients disconnected") } diff --git a/internal/messaging/event_meta.go b/internal/messaging/event_meta.go new file mode 100644 index 0000000..98e1cf4 --- /dev/null +++ b/internal/messaging/event_meta.go @@ -0,0 +1,71 @@ +package messaging + +import ( + "context" + "fmt" +) + +// EventMeta is the metadata-only view of a message used by event feeds +// (e.g. the agent SSE stream). It deliberately carries no message body. +type EventMeta struct { + MessageID int64 + FromAgent string + ToAgent string // set for DMs + Channel string // set for channel messages + Subject string +} + +// ListEventMetaAfter returns metadata for messages with id > afterID that are +// visible to agentName, ordered by id ascending, at most limit rows. +// +// Visible means: DMs addressed to the agent, plus channel messages in channels +// the agent is a member of (same scope as SearchMessages). Messages sent by the +// agent itself are excluded. +// +// Note: channel messages are stored with to_agent NULL (see InsertMessage), so +// the channel branch matches both NULL and ''. +func (s *SQLiteMessageStore) ListEventMetaAfter(ctx context.Context, agentName string, afterID int64, limit int) ([]*EventMeta, error) { + if limit <= 0 { + limit = 50 + } + rows, err := s.db.QueryContext(ctx, + `SELECT m.id, m.from_agent, COALESCE(m.to_agent, ''), COALESCE(ch.name, ''), COALESCE(cv.subject, '') + FROM messages m + LEFT JOIN channels ch ON ch.id = m.channel_id + LEFT JOIN conversations cv ON cv.id = m.conversation_id + WHERE m.id > ? AND m.from_agent <> ? + AND (m.to_agent = ? OR (m.channel_id IS NOT NULL AND (m.to_agent IS NULL OR m.to_agent = '') AND EXISTS (SELECT 1 FROM channel_members cm WHERE cm.channel_id = m.channel_id AND cm.agent_name = ?))) + ORDER BY m.id ASC + LIMIT ?`, + afterID, agentName, agentName, agentName, limit, + ) + if err != nil { + return nil, fmt.Errorf("query event meta: %w", err) + } + defer rows.Close() + + var out []*EventMeta + for rows.Next() { + var e EventMeta + if err := rows.Scan(&e.MessageID, &e.FromAgent, &e.ToAgent, &e.Channel, &e.Subject); err != nil { + return nil, fmt.Errorf("scan event meta: %w", err) + } + out = append(out, &e) + } + return out, rows.Err() +} + +// ListEventMetaAfter returns body-free metadata for messages visible to +// agentName with id > afterID, ascending by id, at most limit rows. +func (s *MessagingService) ListEventMetaAfter(ctx context.Context, agentName string, afterID int64, limit int) ([]*EventMeta, error) { + return s.store.ListEventMetaAfter(ctx, agentName, afterID, limit) +} + +// GetConversationSubject returns the subject of a conversation, or "" if unknown. +func (s *MessagingService) GetConversationSubject(ctx context.Context, conversationID int64) string { + conv, err := s.store.GetConversation(ctx, conversationID) + if err != nil || conv == nil { + return "" + } + return conv.Subject +} diff --git a/internal/messaging/store.go b/internal/messaging/store.go index 2b525d0..a7e4dbf 100644 --- a/internal/messaging/store.go +++ b/internal/messaging/store.go @@ -41,6 +41,7 @@ type MessageStore interface { GetConversationIDsForChannel(ctx context.Context, channelID int64, lastMessageID int64) ([]int64, error) GetConversationIDsForDM(ctx context.Context, agentNames []string, peerAgent string, lastMessageID int64) ([]int64, error) GetReplyCounts(ctx context.Context, messageIDs []int64) (map[int64]int, error) + ListEventMetaAfter(ctx context.Context, agentName string, afterID int64, limit int) ([]*EventMeta, error) } // SQLiteMessageStore implements MessageStore using SQLite.