From eec06f5b5fbd1a0c1cb6377f63eb88d9e2649741 Mon Sep 17 00:00:00 2001 From: Algis Dumbris Date: Sun, 15 Mar 2026 09:18:36 +0200 Subject: [PATCH] =?UTF-8?q?feat:=20web=20UI=20notifications=20=E2=80=94=20?= =?UTF-8?q?unread=20badges,=20new=20message=20line,=20auto=20mark-as-read?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Backend: - GET /api/notifications/unread returns channel + DM unread counts - POST /api/notifications/mark-read updates inbox_state for channels/DMs - SSE broadcaster wired into message send for real-time push - last_read_message_id added to channel/DM message responses Frontend: - Notification store tracks unread counts per channel/DM - SSE listener for new_message and unread_update events - Red circular badges in sidebar (Slack-style) - "New messages" separator line in channel/DM views - Auto mark-as-read after 2 seconds of viewing Co-Authored-By: Claude Opus 4.6 --- internal/api/broadcaster.go | 109 ++++++ internal/api/channels_handler.go | 19 +- internal/api/messages_handler.go | 24 ++ internal/api/notifications_handler.go | 251 +++++++++++++ internal/api/notifications_handler_test.go | 373 ++++++++++++++++++++ internal/api/router.go | 11 + internal/channels/service.go | 1 + internal/channels/store.go | 3 +- internal/messaging/service.go | 30 ++ internal/messaging/store.go | 183 ++++++++++ internal/messaging/types.go | 7 + internal/web/dist/index.html | 18 +- web/src/lib/api/client.ts | 8 + web/src/lib/api/sse.ts | 2 +- web/src/lib/components/Sidebar.svelte | 22 +- web/src/lib/stores/notifications.ts | 108 ++++++ web/src/routes/+layout.svelte | 23 ++ web/src/routes/channels/[name]/+page.svelte | 43 ++- web/src/routes/dm/[name]/+page.svelte | 43 ++- 19 files changed, 1255 insertions(+), 23 deletions(-) create mode 100644 internal/api/broadcaster.go create mode 100644 internal/api/notifications_handler.go create mode 100644 internal/api/notifications_handler_test.go create mode 100644 web/src/lib/stores/notifications.ts diff --git a/internal/api/broadcaster.go b/internal/api/broadcaster.go new file mode 100644 index 0000000..754e0b9 --- /dev/null +++ b/internal/api/broadcaster.go @@ -0,0 +1,109 @@ +package api + +import ( + "context" + "log/slog" + + "github.com/synapbus/synapbus/internal/agents" + "github.com/synapbus/synapbus/internal/channels" +) + +// NewMessageEvent is broadcast when a new message is sent. +type NewMessageEvent struct { + Channel string `json:"channel,omitempty"` // set for channel messages + FromAgent string `json:"from_agent,omitempty"` // set for DMs + ToAgent string `json:"to_agent,omitempty"` // set for DMs + MessageID int64 `json:"message_id"` +} + +// UnreadUpdateEvent is broadcast when unread counts change (e.g. mark-read). +type UnreadUpdateEvent struct { + Channel string `json:"channel,omitempty"` + Agent string `json:"agent,omitempty"` + UnreadCount int `json:"unread_count"` +} + +// EventBroadcaster broadcasts real-time events to connected SSE clients. +type EventBroadcaster interface { + BroadcastNewMessage(ctx context.Context, ownerID int64, event NewMessageEvent) + BroadcastUnreadUpdate(ctx context.Context, ownerID int64, event UnreadUpdateEvent) +} + +// SSEBroadcaster implements EventBroadcaster using the SSEHub. +type SSEBroadcaster struct { + hub *SSEHub + agentService *agents.AgentService + channelService *channels.Service + logger *slog.Logger +} + +// NewSSEBroadcaster creates a broadcaster that sends events via SSE. +func NewSSEBroadcaster(hub *SSEHub, agentService *agents.AgentService, channelService *channels.Service) *SSEBroadcaster { + return &SSEBroadcaster{ + hub: hub, + agentService: agentService, + channelService: channelService, + logger: slog.Default().With("component", "api.broadcaster"), + } +} + +// 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{ + Type: "new_message", + Data: event, + }) +} + +// BroadcastUnreadUpdate sends an unread_update event to the given owner. +func (b *SSEBroadcaster) BroadcastUnreadUpdate(_ context.Context, ownerID int64, event UnreadUpdateEvent) { + b.hub.Broadcast(ownerID, SSEEvent{ + Type: "unread_update", + Data: event, + }) +} + +// BroadcastDM broadcasts a new_message event for a direct message. +// It resolves the recipient agent's owner and sends the event to them. +func (b *SSEBroadcaster) BroadcastDM(ctx context.Context, msg NewMessageEvent) { + if msg.ToAgent == "" { + return + } + + agent, err := b.agentService.GetAgent(ctx, msg.ToAgent) + if err != nil { + b.logger.Debug("could not resolve recipient owner for SSE broadcast", + "to_agent", msg.ToAgent, "error", err) + return + } + + b.BroadcastNewMessage(ctx, agent.OwnerID, msg) +} + +// BroadcastChannelMessage broadcasts a new_message event for a channel message. +// It resolves all channel members' owners and sends the event to each unique owner. +func (b *SSEBroadcaster) BroadcastChannelMessage(ctx context.Context, channelID int64, msg NewMessageEvent) { + if b.channelService == nil { + return + } + + members, err := b.channelService.GetMembers(ctx, channelID) + if err != nil { + b.logger.Debug("could not get channel members for SSE broadcast", + "channel_id", channelID, "error", err) + return + } + + // Collect unique owner IDs to avoid duplicate broadcasts. + seen := make(map[int64]bool) + for _, m := range members { + agent, err := b.agentService.GetAgent(ctx, m.AgentName) + if err != nil { + continue + } + if !seen[agent.OwnerID] { + seen[agent.OwnerID] = true + b.BroadcastNewMessage(ctx, agent.OwnerID, msg) + } + } +} diff --git a/internal/api/channels_handler.go b/internal/api/channels_handler.go index 1ccc58b..8bd5c10 100644 --- a/internal/api/channels_handler.go +++ b/internal/api/channels_handler.go @@ -201,7 +201,7 @@ func (h *ChannelsHandler) JoinChannel(w http.ResponseWriter, r *http.Request) { // ChannelMessages handles GET /api/channels/{name}/messages. func (h *ChannelsHandler) ChannelMessages(w http.ResponseWriter, r *http.Request) { - _, ok := OwnerIDFromContext(r.Context()) + ownerID, ok := OwnerIDFromContext(r.Context()) if !ok { writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required")) return @@ -228,9 +228,22 @@ func (h *ChannelsHandler) ChannelMessages(w http.ResponseWriter, r *http.Request return } + // Compute last_read_message_id across owned agents + var lastReadMessageID int64 + ownedAgents, err := h.agentService.ListAgents(r.Context(), ownerID) + if err == nil { + for _, agent := range ownedAgents { + lr, err := h.msgService.GetLastReadForChannel(r.Context(), agent.Name, ch.ID) + if err == nil && lr > lastReadMessageID { + lastReadMessageID = lr + } + } + } + writeJSON(w, http.StatusOK, map[string]any{ - "messages": paginated.Messages, - "total": paginated.Total, + "messages": paginated.Messages, + "total": paginated.Total, + "last_read_message_id": lastReadMessageID, }) } diff --git a/internal/api/messages_handler.go b/internal/api/messages_handler.go index 40db0be..3f04a34 100644 --- a/internal/api/messages_handler.go +++ b/internal/api/messages_handler.go @@ -18,9 +18,15 @@ import ( type MessagesHandler struct { msgService *messaging.MessagingService agentService *agents.AgentService + broadcaster *SSEBroadcaster logger *slog.Logger } +// SetBroadcaster sets the event broadcaster for real-time SSE notifications. +func (h *MessagesHandler) SetBroadcaster(b *SSEBroadcaster) { + h.broadcaster = b +} + // NewMessagesHandler creates a new messages handler. func NewMessagesHandler(msgService *messaging.MessagingService, agentService *agents.AgentService) *MessagesHandler { return &MessagesHandler{ @@ -309,6 +315,24 @@ func (h *MessagesHandler) SendMessage(w http.ResponseWriter, r *http.Request) { return } + // Broadcast real-time event to connected SSE clients + if h.broadcaster != nil { + event := NewMessageEvent{ + MessageID: msg.ID, + FromAgent: msg.FromAgent, + ToAgent: msg.ToAgent, + } + if msg.ChannelID != nil && h.broadcaster.channelService != nil { + ch, chErr := h.broadcaster.channelService.GetChannel(r.Context(), *msg.ChannelID) + if chErr == nil { + event.Channel = ch.Name + } + h.broadcaster.BroadcastChannelMessage(r.Context(), *msg.ChannelID, event) + } else { + h.broadcaster.BroadcastDM(r.Context(), event) + } + } + writeJSON(w, http.StatusCreated, msg) } diff --git a/internal/api/notifications_handler.go b/internal/api/notifications_handler.go new file mode 100644 index 0000000..c4a0102 --- /dev/null +++ b/internal/api/notifications_handler.go @@ -0,0 +1,251 @@ +package api + +import ( + "encoding/json" + "log/slog" + "net/http" + + "github.com/synapbus/synapbus/internal/agents" + "github.com/synapbus/synapbus/internal/channels" + "github.com/synapbus/synapbus/internal/messaging" +) + +// NotificationsHandler handles REST API requests for notification badges. +type NotificationsHandler struct { + msgService *messaging.MessagingService + agentService *agents.AgentService + channelService *channels.Service + logger *slog.Logger +} + +// NewNotificationsHandler creates a new notifications handler. +func NewNotificationsHandler(msgService *messaging.MessagingService, agentService *agents.AgentService, channelService *channels.Service) *NotificationsHandler { + return &NotificationsHandler{ + msgService: msgService, + agentService: agentService, + channelService: channelService, + logger: slog.Default().With("component", "api.notifications"), + } +} + +// channelUnread is the JSON shape for a channel's unread info. +type channelUnread struct { + Name string `json:"name"` + UnreadCount int `json:"unread_count"` + LastMessageID int64 `json:"last_message_id"` +} + +// dmUnread is the JSON shape for a DM peer's unread info. +type dmUnread struct { + Agent string `json:"agent"` + UnreadCount int `json:"unread_count"` + LastMessageID int64 `json:"last_message_id"` +} + +// UnreadCounts handles GET /api/notifications/unread. +func (h *NotificationsHandler) UnreadCounts(w http.ResponseWriter, r *http.Request) { + ownerID, ok := OwnerIDFromContext(r.Context()) + if !ok { + writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required")) + return + } + + ownedAgents, err := h.agentService.ListAgents(r.Context(), ownerID) + if err != nil { + h.logger.Error("list agents failed", "error", err) + writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to list agents")) + return + } + + if len(ownedAgents) == 0 { + writeJSON(w, http.StatusOK, map[string]any{ + "channels": []channelUnread{}, + "dms": []dmUnread{}, + "total_unread": 0, + }) + return + } + + // Aggregate channel summaries across all owned agents + channelMap := make(map[string]*channelUnread) + for _, agent := range ownedAgents { + if h.channelService == nil { + break + } + summaries, err := h.channelService.GetChannelSummaries(r.Context(), agent.Name) + if err != nil { + h.logger.Error("get channel summaries failed", "agent", agent.Name, "error", err) + continue + } + for _, cs := range summaries { + existing, ok := channelMap[cs.Name] + if !ok { + channelMap[cs.Name] = &channelUnread{ + Name: cs.Name, + UnreadCount: cs.UnreadCount, + LastMessageID: cs.LastMessageID, + } + } else { + // Take the max unread count (different agents may see different counts) + if cs.UnreadCount > existing.UnreadCount { + existing.UnreadCount = cs.UnreadCount + } + if cs.LastMessageID > existing.LastMessageID { + existing.LastMessageID = cs.LastMessageID + } + } + } + } + + channelsList := make([]channelUnread, 0, len(channelMap)) + for _, cu := range channelMap { + channelsList = append(channelsList, *cu) + } + + // Aggregate DM unread counts across all owned agents + dmMap := make(map[string]*dmUnread) + for _, agent := range ownedAgents { + counts, err := h.msgService.GetDMUnreadCounts(r.Context(), agent.Name) + if err != nil { + h.logger.Error("get dm unread counts failed", "agent", agent.Name, "error", err) + continue + } + for _, dc := range counts { + existing, ok := dmMap[dc.Agent] + if !ok { + dmMap[dc.Agent] = &dmUnread{ + Agent: dc.Agent, + UnreadCount: dc.UnreadCount, + LastMessageID: dc.LastMessageID, + } + } else { + existing.UnreadCount += dc.UnreadCount + if dc.LastMessageID > existing.LastMessageID { + existing.LastMessageID = dc.LastMessageID + } + } + } + } + + dmsList := make([]dmUnread, 0, len(dmMap)) + for _, du := range dmMap { + dmsList = append(dmsList, *du) + } + + totalUnread := 0 + for _, cu := range channelsList { + totalUnread += cu.UnreadCount + } + for _, du := range dmsList { + totalUnread += du.UnreadCount + } + + writeJSON(w, http.StatusOK, map[string]any{ + "channels": channelsList, + "dms": dmsList, + "total_unread": totalUnread, + }) +} + +// MarkRead handles POST /api/notifications/mark-read. +func (h *NotificationsHandler) MarkRead(w http.ResponseWriter, r *http.Request) { + ownerID, ok := OwnerIDFromContext(r.Context()) + if !ok { + writeJSON(w, http.StatusUnauthorized, errorBody("unauthorized", "Authentication required")) + return + } + + var req struct { + Type string `json:"type"` + Target string `json:"target"` + LastMessageID int64 `json:"last_message_id"` + } + + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + writeJSON(w, http.StatusBadRequest, errorBody("invalid_request", "Invalid JSON body")) + return + } + + if req.Type == "" || req.Target == "" || req.LastMessageID <= 0 { + writeJSON(w, http.StatusBadRequest, errorBody("validation_error", "type, target, and last_message_id are required")) + return + } + + if req.Type != "channel" && req.Type != "dm" { + writeJSON(w, http.StatusBadRequest, errorBody("validation_error", "type must be 'channel' or 'dm'")) + return + } + + ownedAgents, err := h.agentService.ListAgents(r.Context(), ownerID) + if err != nil { + h.logger.Error("list agents failed", "error", err) + writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to list agents")) + return + } + + if len(ownedAgents) == 0 { + writeJSON(w, http.StatusBadRequest, errorBody("no_agents", "No agents registered")) + return + } + + agentNames := make([]string, len(ownedAgents)) + for i, a := range ownedAgents { + agentNames[i] = a.Name + } + + if req.Type == "channel" { + if h.channelService == nil { + writeJSON(w, http.StatusBadRequest, errorBody("not_available", "Channel service not available")) + return + } + + ch, err := h.channelService.GetChannelByName(r.Context(), req.Target) + if err != nil { + writeJSON(w, http.StatusNotFound, errorBody("not_found", "Channel not found")) + return + } + + // Get conversation IDs for messages in this channel up to the given message ID + convIDs, err := h.msgService.GetConversationIDsForChannel(r.Context(), ch.ID, req.LastMessageID) + if err != nil { + h.logger.Error("get conversation ids failed", "error", err) + writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to get conversations")) + return + } + + // Update inbox state for all owned agents on all relevant conversations + for _, agentName := range agentNames { + for _, convID := range convIDs { + if err := h.msgService.UpdateInboxState(r.Context(), agentName, convID, req.LastMessageID); err != nil { + h.logger.Error("update inbox state failed", + "agent", agentName, + "conversation_id", convID, + "error", err, + ) + } + } + } + } else { + // DM mark-read + convIDs, err := h.msgService.GetConversationIDsForDM(r.Context(), agentNames, req.Target, req.LastMessageID) + if err != nil { + h.logger.Error("get dm conversation ids failed", "error", err) + writeJSON(w, http.StatusInternalServerError, errorBody("server_error", "Failed to get conversations")) + return + } + + for _, agentName := range agentNames { + for _, convID := range convIDs { + if err := h.msgService.UpdateInboxState(r.Context(), agentName, convID, req.LastMessageID); err != nil { + h.logger.Error("update inbox state failed", + "agent", agentName, + "conversation_id", convID, + "error", err, + ) + } + } + } + } + + writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) +} diff --git a/internal/api/notifications_handler_test.go b/internal/api/notifications_handler_test.go new file mode 100644 index 0000000..998ca69 --- /dev/null +++ b/internal/api/notifications_handler_test.go @@ -0,0 +1,373 @@ +package api + +import ( + "bytes" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "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" +) + +func setupNotificationsRouter(t *testing.T) (chi.Router, *messaging.MessagingService, *agents.AgentService, *channels.Service) { + t.Helper() + db := newTestDBFull(t) + + msgStore := messaging.NewSQLiteMessageStore(db) + msgService := messaging.NewMessagingService(msgStore, nil) + + agentStore := agents.NewSQLiteAgentStore(db) + agentService := agents.NewAgentService(agentStore, nil) + + channelStore := channels.NewSQLiteChannelStore(db) + channelService := channels.NewService(channelStore, msgService, nil) + + // Seed agents + seedTestAgent(t, db, "human-agent", 1) + seedTestAgent(t, db, "bot-alice", 2) + seedTestAgent(t, db, "bot-bob", 2) + + handler := NewNotificationsHandler(msgService, agentService, channelService) + messagesHandler := NewMessagesHandler(msgService, agentService) + channelsHandler := NewChannelsHandler(channelService, agentService, msgService) + + router := chi.NewRouter() + router.Group(func(r chi.Router) { + r.Use(OwnerAuthMiddleware) + r.Get("/api/notifications/unread", handler.UnreadCounts) + r.Post("/api/notifications/mark-read", handler.MarkRead) + r.Get("/api/channels/{name}/messages", channelsHandler.ChannelMessages) + r.Get("/api/agents/{name}/messages", messagesHandler.DMMessages) + }) + + return router, msgService, agentService, channelService +} + +func TestUnreadCounts_Unauthenticated(t *testing.T) { + router, _, _, _ := setupNotificationsRouter(t) + + req := httptest.NewRequest("GET", "/api/notifications/unread", nil) + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusUnauthorized { + t.Errorf("status = %d, want %d", rr.Code, http.StatusUnauthorized) + } +} + +func TestUnreadCounts_NoAgents(t *testing.T) { + router, _, _, _ := setupNotificationsRouter(t) + + // Owner 99 has no agents + req := httptest.NewRequest("GET", "/api/notifications/unread", nil) + req.Header.Set("X-Owner-ID", "99") + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want %d, body: %s", rr.Code, http.StatusOK, rr.Body.String()) + } + + var resp struct { + Channels []channelUnread `json:"channels"` + DMs []dmUnread `json:"dms"` + TotalUnread int `json:"total_unread"` + } + json.Unmarshal(rr.Body.Bytes(), &resp) + + if len(resp.Channels) != 0 { + t.Errorf("channels = %d, want 0", len(resp.Channels)) + } + if len(resp.DMs) != 0 { + t.Errorf("dms = %d, want 0", len(resp.DMs)) + } + if resp.TotalUnread != 0 { + t.Errorf("total_unread = %d, want 0", resp.TotalUnread) + } +} + +func TestUnreadCounts_WithDMs(t *testing.T) { + router, msgService, _, _ := setupNotificationsRouter(t) + + ctx := t.Context() + + // bot-alice sends DMs to human-agent (owner 1) + _, err := msgService.SendMessage(ctx, "bot-alice", "human-agent", "Hello from Alice", messaging.SendOptions{Subject: "dm"}) + if err != nil { + t.Fatalf("send message: %v", err) + } + _, err = msgService.SendMessage(ctx, "bot-alice", "human-agent", "Second message from Alice", messaging.SendOptions{Subject: "dm"}) + if err != nil { + t.Fatalf("send message: %v", err) + } + + // bot-bob sends a DM to human-agent + _, err = msgService.SendMessage(ctx, "bot-bob", "human-agent", "Hello from Bob", messaging.SendOptions{Subject: "dm"}) + if err != nil { + t.Fatalf("send message: %v", err) + } + + req := httptest.NewRequest("GET", "/api/notifications/unread", nil) + req.Header.Set("X-Owner-ID", "1") + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("status = %d, want %d, body: %s", rr.Code, http.StatusOK, rr.Body.String()) + } + + var resp struct { + Channels []channelUnread `json:"channels"` + DMs []dmUnread `json:"dms"` + TotalUnread int `json:"total_unread"` + } + json.Unmarshal(rr.Body.Bytes(), &resp) + + if len(resp.DMs) != 2 { + t.Fatalf("dms = %d, want 2; body: %s", len(resp.DMs), rr.Body.String()) + } + + // Find alice and bob counts + dmsByAgent := make(map[string]dmUnread) + for _, dm := range resp.DMs { + dmsByAgent[dm.Agent] = dm + } + + if alice, ok := dmsByAgent["bot-alice"]; !ok { + t.Error("expected DM from bot-alice") + } else if alice.UnreadCount != 2 { + t.Errorf("bot-alice unread = %d, want 2", alice.UnreadCount) + } + + if bob, ok := dmsByAgent["bot-bob"]; !ok { + t.Error("expected DM from bot-bob") + } else if bob.UnreadCount != 1 { + t.Errorf("bot-bob unread = %d, want 1", bob.UnreadCount) + } + + if resp.TotalUnread != 3 { + t.Errorf("total_unread = %d, want 3", resp.TotalUnread) + } +} + +func TestMarkRead_DM(t *testing.T) { + router, msgService, _, _ := setupNotificationsRouter(t) + + ctx := t.Context() + + // bot-alice sends 3 DMs to human-agent + msg1, _ := msgService.SendMessage(ctx, "bot-alice", "human-agent", "Message 1", messaging.SendOptions{Subject: "dm"}) + _, _ = msgService.SendMessage(ctx, "bot-alice", "human-agent", "Message 2", messaging.SendOptions{Subject: "dm"}) + msg3, _ := msgService.SendMessage(ctx, "bot-alice", "human-agent", "Message 3", messaging.SendOptions{Subject: "dm"}) + + // Mark read up to msg1 + body, _ := json.Marshal(map[string]any{ + "type": "dm", + "target": "bot-alice", + "last_message_id": msg1.ID, + }) + req := httptest.NewRequest("POST", "/api/notifications/mark-read", bytes.NewReader(body)) + req.Header.Set("X-Owner-ID", "1") + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("mark-read status = %d, want %d, body: %s", rr.Code, http.StatusOK, rr.Body.String()) + } + + // Check unread — should have 2 unread still + req = httptest.NewRequest("GET", "/api/notifications/unread", nil) + req.Header.Set("X-Owner-ID", "1") + rr = httptest.NewRecorder() + router.ServeHTTP(rr, req) + + var resp struct { + DMs []dmUnread `json:"dms"` + TotalUnread int `json:"total_unread"` + } + json.Unmarshal(rr.Body.Bytes(), &resp) + + if len(resp.DMs) != 1 { + t.Fatalf("dms = %d, want 1; body: %s", len(resp.DMs), rr.Body.String()) + } + if resp.DMs[0].UnreadCount != 2 { + t.Errorf("unread after mark-read = %d, want 2", resp.DMs[0].UnreadCount) + } + + // Mark all read up to msg3 + body, _ = json.Marshal(map[string]any{ + "type": "dm", + "target": "bot-alice", + "last_message_id": msg3.ID, + }) + req = httptest.NewRequest("POST", "/api/notifications/mark-read", bytes.NewReader(body)) + req.Header.Set("X-Owner-ID", "1") + rr = httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("mark-read status = %d, want %d", rr.Code, http.StatusOK) + } + + // Check unread — should be 0 + req = httptest.NewRequest("GET", "/api/notifications/unread", nil) + req.Header.Set("X-Owner-ID", "1") + rr = httptest.NewRecorder() + router.ServeHTTP(rr, req) + + json.Unmarshal(rr.Body.Bytes(), &resp) + if resp.TotalUnread != 0 { + t.Errorf("total_unread after marking all read = %d, want 0; body: %s", resp.TotalUnread, rr.Body.String()) + } +} + +func TestMarkRead_Validation(t *testing.T) { + router, _, _, _ := setupNotificationsRouter(t) + + tests := []struct { + name string + body map[string]any + want int + }{ + { + name: "missing type", + body: map[string]any{"target": "foo", "last_message_id": 1}, + want: http.StatusBadRequest, + }, + { + name: "missing target", + body: map[string]any{"type": "dm", "last_message_id": 1}, + want: http.StatusBadRequest, + }, + { + name: "missing last_message_id", + body: map[string]any{"type": "dm", "target": "foo"}, + want: http.StatusBadRequest, + }, + { + name: "invalid type", + body: map[string]any{"type": "invalid", "target": "foo", "last_message_id": 1}, + want: http.StatusBadRequest, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + body, _ := json.Marshal(tt.body) + req := httptest.NewRequest("POST", "/api/notifications/mark-read", bytes.NewReader(body)) + req.Header.Set("X-Owner-ID", "1") + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != tt.want { + t.Errorf("status = %d, want %d, body: %s", rr.Code, tt.want, rr.Body.String()) + } + }) + } +} + +func TestDMMessages_IncludesLastRead(t *testing.T) { + router, msgService, _, _ := setupNotificationsRouter(t) + + ctx := t.Context() + + // Send DMs from bot-alice to human-agent + msg1, _ := msgService.SendMessage(ctx, "bot-alice", "human-agent", "Hello", messaging.SendOptions{Subject: "dm"}) + _, _ = msgService.SendMessage(ctx, "bot-alice", "human-agent", "World", messaging.SendOptions{Subject: "dm"}) + + // Mark read up to msg1 + body, _ := json.Marshal(map[string]any{ + "type": "dm", + "target": "bot-alice", + "last_message_id": msg1.ID, + }) + req := httptest.NewRequest("POST", "/api/notifications/mark-read", bytes.NewReader(body)) + req.Header.Set("X-Owner-ID", "1") + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("mark-read status = %d", rr.Code) + } + + // GET DM messages should include last_read_message_id + req = httptest.NewRequest("GET", "/api/agents/bot-alice/messages", nil) + req.Header.Set("X-Owner-ID", "1") + rr = httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("dm messages status = %d, body: %s", rr.Code, rr.Body.String()) + } + + var resp map[string]any + json.Unmarshal(rr.Body.Bytes(), &resp) + + lastRead, ok := resp["last_read_message_id"] + if !ok { + t.Fatal("response missing last_read_message_id") + } + if int64(lastRead.(float64)) != msg1.ID { + t.Errorf("last_read_message_id = %v, want %d", lastRead, msg1.ID) + } +} + +func TestChannelMessages_IncludesLastRead(t *testing.T) { + router, _, _, channelService := setupNotificationsRouter(t) + + ctx := t.Context() + + // Create a channel and have human-agent join + ch, err := channelService.CreateChannel(ctx, channels.CreateChannelRequest{ + Name: "test-channel", + CreatedBy: "human-agent", + }) + if err != nil { + t.Fatalf("create channel: %v", err) + } + + // Broadcast messages + msgs, err := channelService.BroadcastMessage(ctx, ch.ID, "human-agent", "Hello channel", 5, "") + if err != nil { + t.Fatalf("broadcast: %v", err) + } + + // Mark read up to the first channel message + body, _ := json.Marshal(map[string]any{ + "type": "channel", + "target": "test-channel", + "last_message_id": msgs[0].ID, + }) + req := httptest.NewRequest("POST", "/api/notifications/mark-read", bytes.NewReader(body)) + req.Header.Set("X-Owner-ID", "1") + rr := httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("mark-read status = %d, body: %s", rr.Code, rr.Body.String()) + } + + // GET channel messages should include last_read_message_id + req = httptest.NewRequest("GET", "/api/channels/test-channel/messages", nil) + req.Header.Set("X-Owner-ID", "1") + rr = httptest.NewRecorder() + router.ServeHTTP(rr, req) + + if rr.Code != http.StatusOK { + t.Fatalf("channel messages status = %d, body: %s", rr.Code, rr.Body.String()) + } + + var resp map[string]any + json.Unmarshal(rr.Body.Bytes(), &resp) + + _, ok := resp["last_read_message_id"] + if !ok { + t.Fatal("response missing last_read_message_id") + } +} diff --git a/internal/api/router.go b/internal/api/router.go index 6c10671..d574d13 100644 --- a/internal/api/router.go +++ b/internal/api/router.go @@ -85,6 +85,13 @@ func NewRouterWithConfig(cfg RouterConfig) chi.Router { if cfg.MsgService != nil && cfg.AgentService != nil { messagesHandler := NewMessagesHandler(cfg.MsgService, cfg.AgentService) agentsHandler := NewAgentsHandler(cfg.AgentService, cfg.TraceStore, cfg.ChannelService) + notificationsHandler := NewNotificationsHandler(cfg.MsgService, cfg.AgentService, cfg.ChannelService) + + // Wire up SSE broadcaster for real-time events + if cfg.SSEHub != nil { + broadcaster := NewSSEBroadcaster(cfg.SSEHub, cfg.AgentService, cfg.ChannelService) + messagesHandler.SetBroadcaster(broadcaster) + } r.Group(func(r chi.Router) { r.Use(authMiddleware) @@ -109,6 +116,10 @@ func NewRouterWithConfig(cfg RouterConfig) chi.Router { r.Delete("/api/agents/{name}", agentsHandler.DeleteAgent) r.Post("/api/agents/{name}/revoke-key", agentsHandler.RevokeKey) r.Get("/api/agents/{name}/messages", messagesHandler.DMMessages) + + // Notifications + r.Get("/api/notifications/unread", notificationsHandler.UnreadCounts) + r.Post("/api/notifications/mark-read", notificationsHandler.MarkRead) }) // API Keys diff --git a/internal/channels/service.go b/internal/channels/service.go index 21c12d9..967bb75 100644 --- a/internal/channels/service.go +++ b/internal/channels/service.go @@ -16,6 +16,7 @@ type ChannelSummary struct { ID int64 `json:"id"` Name string `json:"name"` UnreadCount int `json:"unread"` + LastMessageID int64 `json:"last_message_id"` LastMessageAt *time.Time `json:"last_message_at"` } diff --git a/internal/channels/store.go b/internal/channels/store.go index 3d841c6..85a0f84 100644 --- a/internal/channels/store.go +++ b/internal/channels/store.go @@ -381,6 +381,7 @@ func (s *SQLiteChannelStore) GetChannelSummaries(ctx context.Context, agentName WHERE ist.agent_name = ? AND ist.conversation_id = m.conversation_id), 0) AND m.from_agent != ? ) AS unread_count, + COALESCE((SELECT MAX(m3.id) FROM messages m3 WHERE m3.channel_id = c.id), 0) AS last_message_id, (SELECT MAX(m2.created_at) FROM messages m2 WHERE m2.channel_id = c.id) AS last_message_at FROM channels c JOIN channel_members cm ON cm.channel_id = c.id AND cm.agent_name = ? @@ -396,7 +397,7 @@ func (s *SQLiteChannelStore) GetChannelSummaries(ctx context.Context, agentName for rows.Next() { var cs ChannelSummary var lastMsg sql.NullString - if err := rows.Scan(&cs.ID, &cs.Name, &cs.UnreadCount, &lastMsg); err != nil { + if err := rows.Scan(&cs.ID, &cs.Name, &cs.UnreadCount, &cs.LastMessageID, &lastMsg); err != nil { return nil, fmt.Errorf("scan channel summary: %w", err) } if lastMsg.Valid { diff --git a/internal/messaging/service.go b/internal/messaging/service.go index 01ae391..3de44d6 100644 --- a/internal/messaging/service.go +++ b/internal/messaging/service.go @@ -459,6 +459,36 @@ func (s *MessagingService) GetDMMessages(ctx context.Context, ownedAgents []stri return messages, nil } +// GetDMUnreadCounts returns unread DM counts grouped by peer agent. +func (s *MessagingService) GetDMUnreadCounts(ctx context.Context, agentName string) ([]DMUnreadCount, error) { + return s.store.GetDMUnreadCounts(ctx, agentName) +} + +// GetLastReadForChannel returns the last_read_message_id for an agent in a channel. +func (s *MessagingService) GetLastReadForChannel(ctx context.Context, agentName string, channelID int64) (int64, error) { + return s.store.GetLastReadForChannel(ctx, agentName, channelID) +} + +// GetLastReadForDM returns the last_read_message_id for owned agents in a DM with a peer. +func (s *MessagingService) GetLastReadForDM(ctx context.Context, agentNames []string, peerAgent string) (int64, error) { + return s.store.GetLastReadForDM(ctx, agentNames, peerAgent) +} + +// UpdateInboxState updates the read position for an agent in a conversation. +func (s *MessagingService) UpdateInboxState(ctx context.Context, agentName string, conversationID int64, lastReadMsgID int64) error { + return s.store.UpdateInboxState(ctx, agentName, conversationID, lastReadMsgID) +} + +// GetConversationIDsForChannel returns conversation IDs in a channel with messages up to lastMessageID. +func (s *MessagingService) GetConversationIDsForChannel(ctx context.Context, channelID int64, lastMessageID int64) ([]int64, error) { + return s.store.GetConversationIDsForChannel(ctx, channelID, lastMessageID) +} + +// GetConversationIDsForDM returns conversation IDs for DMs between owned agents and a peer. +func (s *MessagingService) GetConversationIDsForDM(ctx context.Context, agentNames []string, peerAgent string, lastMessageID int64) ([]int64, error) { + return s.store.GetConversationIDsForDM(ctx, agentNames, peerAgent, lastMessageID) +} + // GetConversation returns a conversation and its messages. func (s *MessagingService) GetConversation(ctx context.Context, id int64) (*Conversation, []*Message, error) { conv, err := s.store.GetConversation(ctx, id) diff --git a/internal/messaging/store.go b/internal/messaging/store.go index 8957dbb..b575a1b 100644 --- a/internal/messaging/store.go +++ b/internal/messaging/store.go @@ -34,6 +34,11 @@ type MessageStore interface { GetPendingDMs(ctx context.Context, agentName string, limit int) ([]*Message, error) GetRecentMentions(ctx context.Context, agentName string, limit int) ([]*Message, error) GetSystemNotifications(ctx context.Context, agentName string, limit int) ([]*Message, error) + GetDMUnreadCounts(ctx context.Context, agentName string) ([]DMUnreadCount, error) + GetLastReadForChannel(ctx context.Context, agentName string, channelID int64) (int64, error) + GetLastReadForDM(ctx context.Context, agentNames []string, peerAgent string) (int64, error) + GetConversationIDsForChannel(ctx context.Context, channelID int64, lastMessageID int64) ([]int64, error) + GetConversationIDsForDM(ctx context.Context, agentNames []string, peerAgent string, lastMessageID int64) ([]int64, error) } // SQLiteMessageStore implements MessageStore using SQLite. @@ -773,6 +778,184 @@ func scanMessageFromRows(rows *sql.Rows) (*Message, error) { return &msg, nil } +// GetDMUnreadCounts returns unread DM counts grouped by peer agent. +// For each unique from_agent that has sent DMs to agentName, it computes +// how many messages have id > last_read_message_id (from inbox_state). +func (s *SQLiteMessageStore) GetDMUnreadCounts(ctx context.Context, agentName string) ([]DMUnreadCount, error) { + rows, err := s.db.QueryContext(ctx, + `SELECT + m.from_agent, + COUNT(CASE WHEN m.id > COALESCE( + (SELECT ist.last_read_message_id FROM inbox_state ist + WHERE ist.agent_name = ? AND ist.conversation_id = m.conversation_id), 0) + THEN 1 END) AS unread_count, + MAX(m.id) AS last_message_id + FROM messages m + WHERE m.to_agent = ? + AND m.channel_id IS NULL + AND m.from_agent != 'system' + GROUP BY m.from_agent + ORDER BY m.from_agent`, + agentName, agentName, + ) + if err != nil { + return nil, fmt.Errorf("get dm unread counts: %w", err) + } + defer rows.Close() + + var results []DMUnreadCount + for rows.Next() { + var d DMUnreadCount + if err := rows.Scan(&d.Agent, &d.UnreadCount, &d.LastMessageID); err != nil { + return nil, fmt.Errorf("scan dm unread count: %w", err) + } + results = append(results, d) + } + if results == nil { + results = []DMUnreadCount{} + } + return results, rows.Err() +} + +// GetLastReadForChannel returns the effective last_read_message_id for an agent +// across all conversations in a given channel. +func (s *SQLiteMessageStore) GetLastReadForChannel(ctx context.Context, agentName string, channelID int64) (int64, error) { + var lastRead sql.NullInt64 + err := s.db.QueryRowContext(ctx, + `SELECT MIN(COALESCE(ist.last_read_message_id, 0)) + FROM (SELECT DISTINCT conversation_id FROM messages WHERE channel_id = ?) conv + LEFT JOIN inbox_state ist ON ist.conversation_id = conv.conversation_id AND ist.agent_name = ?`, + channelID, agentName, + ).Scan(&lastRead) + if err != nil { + return 0, fmt.Errorf("get last read for channel: %w", err) + } + if lastRead.Valid { + return lastRead.Int64, nil + } + return 0, nil +} + +// GetLastReadForDM returns the effective last_read_message_id for an owned agent +// in DM conversations with a specific peer agent. +func (s *SQLiteMessageStore) GetLastReadForDM(ctx context.Context, agentNames []string, peerAgent string) (int64, error) { + if len(agentNames) == 0 { + return 0, nil + } + + placeholders := make([]string, len(agentNames)) + args := make([]any, 0, len(agentNames)*2+1) + for i, a := range agentNames { + placeholders[i] = "?" + args = append(args, a) + } + inClause := strings.Join(placeholders, ",") + + // Find conversations between owned agents and peer agent (DMs only) + // and get the max last_read_message_id + query := fmt.Sprintf( + `SELECT COALESCE(MAX(ist.last_read_message_id), 0) + FROM inbox_state ist + WHERE ist.agent_name IN (%s) + AND ist.conversation_id IN ( + SELECT DISTINCT m.conversation_id FROM messages m + WHERE m.channel_id IS NULL + AND ((m.from_agent IN (%s) AND m.to_agent = ?) + OR (m.from_agent = ? AND m.to_agent IN (%s))) + )`, + inClause, inClause, inClause, + ) + for _, a := range agentNames { + args = append(args, a) + } + args = append(args, peerAgent, peerAgent) + for _, a := range agentNames { + args = append(args, a) + } + + var lastRead int64 + err := s.db.QueryRowContext(ctx, query, args...).Scan(&lastRead) + if err != nil { + return 0, fmt.Errorf("get last read for dm: %w", err) + } + return lastRead, nil +} + +// GetConversationIDsForChannel returns conversation IDs in a channel with messages up to lastMessageID. +func (s *SQLiteMessageStore) GetConversationIDsForChannel(ctx context.Context, channelID int64, lastMessageID int64) ([]int64, error) { + rows, err := s.db.QueryContext(ctx, + `SELECT DISTINCT conversation_id FROM messages + WHERE channel_id = ? AND id <= ?`, + channelID, lastMessageID, + ) + if err != nil { + return nil, fmt.Errorf("get conversation ids for channel: %w", err) + } + defer rows.Close() + + var ids []int64 + for rows.Next() { + var id int64 + if err := rows.Scan(&id); err != nil { + return nil, fmt.Errorf("scan conversation id: %w", err) + } + ids = append(ids, id) + } + return ids, rows.Err() +} + +// GetConversationIDsForDM returns conversation IDs for DMs between owned agents and a peer agent, +// with messages up to lastMessageID. +func (s *SQLiteMessageStore) GetConversationIDsForDM(ctx context.Context, agentNames []string, peerAgent string, lastMessageID int64) ([]int64, error) { + if len(agentNames) == 0 { + return nil, nil + } + + placeholders := make([]string, len(agentNames)) + args := make([]any, 0, len(agentNames)*2+2) + for i, a := range agentNames { + placeholders[i] = "?" + args = append(args, a) + } + inClause := strings.Join(placeholders, ",") + + query := fmt.Sprintf( + `SELECT DISTINCT conversation_id FROM messages + WHERE channel_id IS NULL + AND id <= ? + AND ((from_agent IN (%s) AND to_agent = ?) + OR (from_agent = ? AND to_agent IN (%s)))`, + inClause, inClause, + ) + + // Reorder args: agentNames for first IN, lastMessageID, agentNames for second IN... + finalArgs := make([]any, 0, len(agentNames)*2+3) + finalArgs = append(finalArgs, lastMessageID) + for _, a := range agentNames { + finalArgs = append(finalArgs, a) + } + finalArgs = append(finalArgs, peerAgent, peerAgent) + for _, a := range agentNames { + finalArgs = append(finalArgs, a) + } + + rows, err := s.db.QueryContext(ctx, query, finalArgs...) + if err != nil { + return nil, fmt.Errorf("get conversation ids for dm: %w", err) + } + defer rows.Close() + + var ids []int64 + for rows.Next() { + var id int64 + if err := rows.Scan(&id); err != nil { + return nil, fmt.Errorf("scan conversation id: %w", err) + } + ids = append(ids, id) + } + return ids, rows.Err() +} + // scanMessage scans a single message from sql.Row. func scanMessage(row *sql.Row) (*Message, error) { var msg Message diff --git a/internal/messaging/types.go b/internal/messaging/types.go index 2b0ebb8..f7b6461 100644 --- a/internal/messaging/types.go +++ b/internal/messaging/types.go @@ -72,3 +72,10 @@ type InboxState struct { LastReadMessageID int64 `json:"last_read_message_id"` UpdatedAt time.Time `json:"updated_at"` } + +// DMUnreadCount holds the unread DM count for a specific peer agent. +type DMUnreadCount struct { + Agent string `json:"agent"` + UnreadCount int `json:"unread_count"` + LastMessageID int64 `json:"last_message_id"` +} diff --git a/internal/web/dist/index.html b/internal/web/dist/index.html index b183fda..868494e 100644 --- a/internal/web/dist/index.html +++ b/internal/web/dist/index.html @@ -8,29 +8,29 @@ - - - + + + - - - + + +
diff --git a/web/src/routes/channels/[name]/+page.svelte b/web/src/routes/channels/[name]/+page.svelte index 4028d70..fbf8c26 100644 --- a/web/src/routes/channels/[name]/+page.svelte +++ b/web/src/routes/channels/[name]/+page.svelte @@ -1,7 +1,9 @@