From 8a5d5e1f59429def82ee3291b7f2a398cdc83455 Mon Sep 17 00:00:00 2001 From: Algis Dumbris Date: Mon, 11 May 2026 15:09:38 +0300 Subject: [PATCH] =?UTF-8?q?feat(020):=20US1=20=E2=80=94=20proactive=20inje?= =?UTF-8?q?ction=20on=20MCP=20tool=20responses?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Wraps eligible MCP tool handlers (my_status, send_message, search, execute; get_replies excluded as pure metadata) with a middleware that appends relevant_context to the JSON response. Retrieval reuses the existing search.Service hybrid pipeline; owner scoping filters out memories from other owners' agents (SC-008). Pin overlay is a marked TODO for US3. Components: - internal/search/injection.go (+ test): BuildContextPacket with token budget greedy fill, score floor, truncation flag, CoreMemoryProvider interface stubbed for US2. - internal/mcp/injection_wrap.go (+ test): WrapInjection middleware, registered via SetInjection on the existing handler. - internal/mcp/injection_e2e_test.go: adversarial cross-owner test asserts H1 cannot see H2's memories on any wrapped tool. - internal/messaging/memory_injections.go (+ test): 24h audit ring, hourly cleanup tick wired into stalemate worker. Discovery during impl: claim_messages/read_inbox/read_channel live as actions inside the execute bridge, not as registered top-level MCP tools. They inherit injection through the execute wrapper. This commit also bundles pre-existing working-tree changes for the 027 "remove approval noise" cleanup (migration 027, design doc, removal of reminder/escalate logic from stalemate worker, related trims in goals_tools.go and tools_hybrid.go). The two changes touch the same files (stalemate.go, tools_hybrid.go) and bundling them keeps history readable. Co-Authored-By: Claude Opus 4.7 (1M context) --- cmd/synapbus/main.go | 19 +- ...-internal-only-disable-approvals-design.md | 102 +++ internal/mcp/goals_tools.go | 90 +-- internal/mcp/injection_e2e_test.go | 197 +++++ internal/mcp/injection_wrap.go | 251 ++++++ internal/mcp/injection_wrap_test.go | 218 +++++ internal/mcp/server.go | 5 +- internal/mcp/tools_hybrid.go | 122 ++- internal/messaging/memory_injections.go | 147 ++++ internal/messaging/memory_injections_test.go | 166 ++++ internal/messaging/stalemate.go | 757 ++---------------- internal/messaging/stalemate_test.go | 642 +-------------- internal/search/injection.go | 329 ++++++++ internal/search/injection_test.go | 258 ++++++ .../schema/027_remove_approval_noise.sql | 61 ++ 15 files changed, 1934 insertions(+), 1430 deletions(-) create mode 100644 docs/superpowers/specs/2026-05-10-internal-only-disable-approvals-design.md create mode 100644 internal/mcp/injection_e2e_test.go create mode 100644 internal/mcp/injection_wrap.go create mode 100644 internal/mcp/injection_wrap_test.go create mode 100644 internal/messaging/memory_injections.go create mode 100644 internal/messaging/memory_injections_test.go create mode 100644 internal/search/injection.go create mode 100644 internal/search/injection_test.go create mode 100644 internal/storage/schema/027_remove_approval_noise.sql diff --git a/cmd/synapbus/main.go b/cmd/synapbus/main.go index f336673..bcd47bf 100644 --- a/cmd/synapbus/main.go +++ b/cmd/synapbus/main.go @@ -599,7 +599,7 @@ func runServe(cmd *cobra.Command, args []string) error { db.DB, ) mcpSrv.WireGoalsTools(goalsToolReg) - slog.Info("spec-018 MCP tools wired (create_goal, propose_task_tree, propose_agent, claim_task, request_resource, list_resources)") + slog.Info("spec-018 MCP tools wired (create_goal, propose_task_tree, claim_task, request_resource, list_resources, complete_goal)") // Set up SQL query executor for agents (uses read pool if available) queryDB := db.QueryDB() @@ -645,12 +645,10 @@ func runServe(cmd *cobra.Command, args []string) error { slog.Info("stalemate worker disabled by SYNAPBUS_DISABLE_STALEMATE_WORKER=1") } else { stalemateConfig := messaging.ParseStalemateConfig() - stalemateWorker = messaging.NewStalemateWorker(db.DB, msgService, &channelLookupAdapter{channelService: channelService}, stalemateConfig) + stalemateWorker = messaging.NewStalemateWorker(db.DB, msgService, stalemateConfig) stalemateWorker.Start() slog.Info("stalemate worker started", "processing_timeout", stalemateConfig.ProcessingTimeout.String(), - "reminder_after", stalemateConfig.ReminderAfter.String(), - "escalate_after", stalemateConfig.EscalateAfter.String(), "interval", stalemateConfig.Interval.String(), ) } @@ -1158,19 +1156,6 @@ func ensureDefaultMCPClient(ctx context.Context, db *sql.DB, bcryptCost int) { ) } -// channelLookupAdapter adapts channels.Service to messaging.ChannelLookup. -type channelLookupAdapter struct { - channelService *channels.Service -} - -func (a *channelLookupAdapter) GetChannelIDByName(ctx context.Context, name string) (int64, error) { - ch, err := a.channelService.GetChannelByName(ctx, name) - if err != nil { - return 0, err - } - return ch.ID, nil -} - // trustAdjusterAdapter adapts trust.Service to reactions.TrustAdjuster. type trustAdjusterAdapter struct { svc *trust.Service diff --git a/docs/superpowers/specs/2026-05-10-internal-only-disable-approvals-design.md b/docs/superpowers/specs/2026-05-10-internal-only-disable-approvals-design.md new file mode 100644 index 0000000..d3c6d40 --- /dev/null +++ b/docs/superpowers/specs/2026-05-10-internal-only-disable-approvals-design.md @@ -0,0 +1,102 @@ +# Internal-only mode: remove approvals & escalations + +**Date:** 2026-05-10 +**Status:** Design +**Owner:** Algis + +## Problem + +SynapBus today assumes a human is in the loop: a stalemate worker DMs reminders after 4h, escalates to `#approvals` after 48h, and the dynamic-agent-spawning flow (spec 018) gates new agents and task trees on human approval. In practice the user is the only operator, the approval queue stalls, and the volume of reminder/escalation messages drowns out signal. The user is moving to a single daily summary (separate `#summary-daily` channel + summarizer agent already in progress) and treats SynapBus as an internal-only comms + data store — nothing publishes externally. + +The goal is to remove the human-in-the-loop surfaces so the message stream stops generating noise the user will never read. + +## Scope + +### Removed + +1. **Stalemate reminders** — `StalemateWorker.sendPendingReminders` and supporting helpers (`reminderExists`, the 4h ReminderAfter knob). +2. **Stalemate escalations** — `StalemateWorker.escalatePendingMessages` and `checkWorkflowStalemates`, plus the 48h EscalateAfter knob and `#approvals` lookup path. +3. **`propose_agent` MCP tool** (spec 018). It writes a `pending` row to `agent_proposals` for human approval via `#approvals` and there is no automated consumer of that table. Removing the tool leaves agent creation to the admin CLI, which matches the internal-only stance. + + **Note:** `propose_task_tree` is intentionally KEPT despite its name — it is not an approval gate. It directly inserts tasks in `approved` status and auto-transitions the goal to `active`. Removing it would break the spec-018 goal/task flow. +4. *(Reactions service intentionally untouched — it's a generic workflow primitive that also drives trust adjustments. Once no upstream feature creates approval-bearing messages, the `approve` / `reject` reaction paths become dormant on their own.)* + +### Kept + +- **`StalemateWorker.ProcessingTimeout`** (24h auto-fail of claimed-but-abandoned messages). Protects the inbox from crashed agents; not human-facing. +- **The `#approvals` channel row** in the `channels` table. Cheaper to leave than to migrate; user can drop via admin CLI later. +- **Webhook / K8s runner approval gates** (spec 003). User confirmed these are out of scope. +- **Trust system** (spec 011). No approval surface, just delegation. + +### One-shot DB cleanup + +New migration `internal/storage/schema/027_remove_approval_noise.sql`: + +```sql +-- Drop reminder and escalation system DMs. +DELETE FROM messages +WHERE subject LIKE 'stalemate-reminder:%' + OR subject LIKE 'stalemate-escalation:%'; + +-- Drop everything in the #approvals channel. +DELETE FROM messages +WHERE channel_id = (SELECT id FROM channels WHERE name = 'approvals'); + +-- Drop pending agent proposals (table itself stays for reversibility). +DELETE FROM agent_proposals; +``` + +`VACUUM` cannot run inside a migration transaction, so reclaiming disk is a separate `synapbus admin vacuum` command (or a manual `kubectl exec ... sqlite3 ... 'VACUUM;'`). Out of scope for this change unless trivial to wire up. + +## Architecture impact + +``` +Before: + agent → MCP propose_agent → agent_proposals row → human reacts in #approvals + → spawn or reject + message claimed → StalemateWorker (every 15m) → 4h reminder DM + → 48h escalation to #approvals + → 24h auto-fail (KEEP) + +After: + agent → MCP create_agent (existing direct path) → agent registered + message claimed → StalemateWorker (every 15m) → 24h auto-fail +``` + +Net code deletion. No new components, no new config surface, no new dependencies. + +## Components touched + +| File | Change | +|------|--------| +| `internal/messaging/stalemate.go` | Delete `sendPendingReminders`, `escalatePendingMessages`, `checkWorkflowStalemates`, `reminderExists`, `escalationExists`. Trim `StalemateConfig` to `ProcessingTimeout` + `Interval`. Remove `ReminderAfter` / `EscalateAfter` env vars. | +| `internal/messaging/stalemate_test.go` | Delete tests for removed methods; keep ProcessingTimeout tests. | +| `internal/messaging/options.go` | Remove channelLookup wiring if it's only used by escalation. | +| `internal/messaging/service.go` | Remove escalation hooks if any. | +| `internal/mcp/goals_tools.go` (spec 018) | Delete `propose_agent` tool registration (`proposeAgentTool`) and its `handleProposeAgent` handler. Keep `propose_task_tree` and the rest of the registrar. | +| `internal/storage/schema/027_remove_approval_noise.sql` | New migration. | +| `cmd/synapbus/admin.go`, `cmd/synapbus/main.go` | Remove any escalation-related flags. | +| `CLAUDE.md` (project + user) | Update SynapBus protocol section to drop "#approvals" + "stalemate auto-fails after 24h" mention of escalation. Keep claim-process-done loop. | +| User's `~/.claude/CLAUDE.md` | Same — drop approval-channel references and the auto-report trigger for "Need approval → #approvals". | + +## Testing + +- Existing `stalemate_test.go` cases for `ProcessingTimeout` continue to pass. +- New test: confirm `StalemateWorker.tick()` no longer queries pending messages for reminder/escalation candidates (no rows touched, no DMs sent). +- New test: confirm `propose_agent` MCP tool returns "tool not found" / is unregistered. +- Migration test: apply `027_remove_approval_noise.sql` to a fixture DB containing stalemate DMs + an `#approvals` message + an `agent_proposals` row; assert all three are gone, other messages untouched. +- No UI testing required — Web UI just stops showing approval-channel content because the channel is empty. + +## Risks & mitigations + +- **An external agent calls `propose_agent` after deletion.** MCP returns an unknown-tool error; agent's runbook should tolerate this. Acceptable because the user controls all agents. +- **Hidden consumer of escalation messages.** Search confirms reminders/escalations are only produced by `StalemateWorker` and consumed by humans. Low risk. +- **Migration deletes too much.** The `LIKE 'stalemate-%'` pattern is narrow and the `#approvals` channel is internal-only; nothing user-authored lives there. Take a `data/synapbus.db` backup before applying in prod (kubic). + +## Out of scope + +- Webhook/K8s runner human gates (spec 003). +- Removing the `#approvals` channel row. +- Adding `SYNAPBUS_APPROVALS_DISABLED` env flag — code deletion is reversible via git revert. +- Daily summarizer agent + `#summary-daily` channel — already in progress in a separate effort. +- Reclaiming disk via `VACUUM` — separate admin command if needed. diff --git a/internal/mcp/goals_tools.go b/internal/mcp/goals_tools.go index d456e93..035ff0e 100644 --- a/internal/mcp/goals_tools.go +++ b/internal/mcp/goals_tools.go @@ -52,18 +52,21 @@ func NewGoalsToolRegistrar( } } -// RegisterAllOnServer attaches create_goal, propose_task_tree, -// propose_agent, claim_task, request_resource, list_resources, and -// complete_goal to the MCP server. +// RegisterAllOnServer attaches create_goal, propose_task_tree, claim_task, +// request_resource, list_resources, and complete_goal to the MCP server. +// +// propose_agent was removed when SynapBus moved to internal-only mode: it +// wrote a pending row to agent_proposals for human approval via #approvals, +// and that approval surface no longer exists. Agents are created via the +// admin CLI directly. See migration 027_remove_approval_noise.sql. func (r *GoalsToolRegistrar) RegisterAllOnServer(s *server.MCPServer) { s.AddTool(r.createGoalTool(), r.handleCreateGoal) s.AddTool(r.proposeTaskTreeTool(), r.handleProposeTaskTree) - s.AddTool(r.proposeAgentTool(), r.handleProposeAgent) s.AddTool(r.claimTaskTool(), r.handleClaimTask) s.AddTool(r.requestResourceTool(), r.handleRequestResource) s.AddTool(r.listResourcesTool(), r.handleListResources) s.AddTool(r.completeGoalTool(), r.handleCompleteGoal) - r.logger.Info("spec-018 MCP tools registered", "count", 7) + r.logger.Info("spec-018 MCP tools registered", "count", 6) } // --- Tool Definitions --- @@ -87,18 +90,6 @@ func (r *GoalsToolRegistrar) proposeTaskTreeTool() mcplib.Tool { ) } -func (r *GoalsToolRegistrar) proposeAgentTool() mcplib.Tool { - return mcplib.NewTool("propose_agent", - mcplib.WithDescription("Propose creating a new specialist agent. Writes an agent_proposals row so a human can approve it via the #approvals channel. Returns the proposal id."), - mcplib.WithString("name", mcplib.Description("Desired agent name"), mcplib.Required()), - mcplib.WithString("display_name", mcplib.Description("Human-readable display name")), - mcplib.WithString("system_prompt", mcplib.Description("System prompt for the spawned agent"), mcplib.Required()), - mcplib.WithString("tool_scope", mcplib.Description("Comma-separated scope (e.g. 'messages:read,messages:send')")), - mcplib.WithNumber("parent_task_id", mcplib.Description("The task this agent will work on")), - mcplib.WithString("autonomy_tier", mcplib.Description("supervised | assisted | autonomous (default assisted)")), - ) -} - func (r *GoalsToolRegistrar) claimTaskTool() mcplib.Tool { return mcplib.NewTool("claim_task", mcplib.WithDescription("Atomically claim an approved task. Returns the claimed task or an error if another agent got it first."), @@ -239,71 +230,6 @@ func (r *GoalsToolRegistrar) handleProposeTaskTree(ctx context.Context, req mcpl }) } -func (r *GoalsToolRegistrar) handleProposeAgent(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) { - agentName, ok := extractAgentName(ctx) - if !ok { - return mcplib.NewToolResultError("authentication required"), nil - } - if r.db == nil { - return mcplib.NewToolResultError("db not configured"), nil - } - name := req.GetString("name", "") - if name == "" { - return mcplib.NewToolResultError("name is required"), nil - } - systemPrompt := req.GetString("system_prompt", "") - if systemPrompt == "" { - return mcplib.NewToolResultError("system_prompt is required"), nil - } - toolScope := req.GetString("tool_scope", "") - if toolScope == "" { - toolScope = "[]" - } else if !strings.HasPrefix(toolScope, "[") { - // Accept comma-separated convenience form. - parts := strings.Split(toolScope, ",") - for i := range parts { - parts[i] = `"` + strings.TrimSpace(parts[i]) + `"` - } - toolScope = "[" + strings.Join(parts, ",") + "]" - } - tier := req.GetString("autonomy_tier", "assisted") - parentTaskID := int64(req.GetInt("parent_task_id", 0)) - if parentTaskID <= 0 { - return mcplib.NewToolResultError("parent_task_id is required (proposals must attach to a task)"), nil - } - - agent, err := r.agents.GetAgent(ctx, agentName) - if err != nil { - return mcplib.NewToolResultError(fmt.Sprintf("resolve caller: %s", err)), nil - } - - // Resolve goal_id from the parent task. - var goalID int64 - if err := r.db.QueryRowContext(ctx, - `SELECT goal_id FROM goal_tasks WHERE id=?`, parentTaskID).Scan(&goalID); err != nil { - return mcplib.NewToolResultError(fmt.Sprintf("resolve goal from task %d: %s", parentTaskID, err)), nil - } - - res, err := r.db.ExecContext(ctx, ` - INSERT INTO agent_proposals ( - proposer_agent_id, goal_id, parent_task_id, proposed_name, - proposed_model, proposed_system_prompt, proposed_tool_scope_json, - proposed_autonomy_tier, status - ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, 'pending')`, - agent.ID, goalID, parentTaskID, name, - "gemini-2.5-flash", systemPrompt, toolScope, tier) - if err != nil { - return mcplib.NewToolResultError(fmt.Sprintf("insert proposal: %s", err)), nil - } - id, _ := res.LastInsertId() - - return resultJSON(map[string]any{ - "proposal_id": id, - "status": "pending", - "name": name, - }) -} - func (r *GoalsToolRegistrar) handleClaimTask(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) { agentName, ok := extractAgentName(ctx) if !ok { diff --git a/internal/mcp/injection_e2e_test.go b/internal/mcp/injection_e2e_test.go new file mode 100644 index 0000000..ba29570 --- /dev/null +++ b/internal/mcp/injection_e2e_test.go @@ -0,0 +1,197 @@ +package mcp + +import ( + "context" + "database/sql" + "encoding/json" + "testing" + + mcplib "github.com/mark3labs/mcp-go/mcp" + _ "modernc.org/sqlite" + + "github.com/synapbus/synapbus/internal/agents" + "github.com/synapbus/synapbus/internal/messaging" + "github.com/synapbus/synapbus/internal/search" + "github.com/synapbus/synapbus/internal/trace" +) + +// TestInjection_CrossOwner_NoLeak is the SC-008 adversarial test: two +// owners (H1, H2) each have an agent + memories in #open-brain. When +// each agent invokes the same wrapped tool with the same query, their +// `relevant_context.memories` are disjoint along owner boundaries. +func TestInjection_CrossOwner_NoLeak(t *testing.T) { + db := newTestDB(t) + ctx := context.Background() + + // Two human owners. + if _, err := db.Exec( + `INSERT OR IGNORE INTO users (id, username, password_hash, display_name) + VALUES (1, 'h1', 'hash', 'H1'), (2, 'h2', 'hash', 'H2')`, + ); err != nil { + t.Fatalf("seed users: %v", err) + } + // One agent per owner. + if _, err := db.Exec( + `INSERT OR IGNORE INTO agents (name, display_name, type, owner_id, api_key_hash, status) + VALUES ('a-h1', 'A1', 'ai', 1, 'k1', 'active'), + ('a-h2', 'A2', 'ai', 2, 'k2', 'active')`, + ); err != nil { + t.Fatalf("seed agents: %v", err) + } + // Both join an open-brain channel — broadly readable. + if _, err := db.Exec( + `INSERT OR IGNORE INTO channels (id, name, description, type, created_by) + VALUES (1, 'open-brain', 'shared', 'standard', 'system')`, + ); err != nil { + t.Fatalf("seed channel: %v", err) + } + if _, err := db.Exec( + `INSERT OR IGNORE INTO channel_members (channel_id, agent_name) + VALUES (1, 'a-h1'), (1, 'a-h2')`, + ); err != nil { + t.Fatalf("seed members: %v", err) + } + // One memory per owner, both about the same topic. + seedMemory(t, db, 1, "a-h1", "Kuzu graph DB is in H1's research notes") + seedMemory(t, db, 1, "a-h2", "Kuzu graph DB also appears in H2's separate research") + + // Stand up a real search.Service (FTS-only, no embeddings). + tracer := trace.NewTracer(db) + t.Cleanup(func() { tracer.Close() }) + msgStore := messaging.NewSQLiteMessageStore(db) + msgService := messaging.NewMessagingService(msgStore, tracer) + searchSvc := search.NewService(db, nil, nil, msgService) + + // Configure wrap: low score floor so the test isn't flaky on FTS. + cfg := WrapConfig{ + Cfg: messaging.MemoryConfig{ + InjectionEnabled: true, + InjectionBudgetTokens: 500, + InjectionMaxItems: 5, + InjectionMinScore: 0.0, + }, + SearchSvc: searchSvc, + QuerySource: func(_ context.Context, _ string, _ map[string]any, _ map[string]any) string { + return "Kuzu" + }, + } + wrapped := WrapInjection(stubHandler(map[string]any{"ok": true}), "search_messages", cfg) + + // Call as H1. + h1Ctx := agents.ContextWithAgent(ctx, &agents.Agent{Name: "a-h1", OwnerID: 1}) + h1Res, err := wrapped(h1Ctx, mcplib.CallToolRequest{}) + if err != nil { + t.Fatalf("h1 wrapped: %v", err) + } + h1Body := unmarshalText(t, h1Res) + + // Call as H2. + h2Ctx := agents.ContextWithAgent(ctx, &agents.Agent{Name: "a-h2", OwnerID: 2}) + h2Res, err := wrapped(h2Ctx, mcplib.CallToolRequest{}) + if err != nil { + t.Fatalf("h2 wrapped: %v", err) + } + h2Body := unmarshalText(t, h2Res) + + h1Memories := extractMemoryFromAgents(h1Body) + h2Memories := extractMemoryFromAgents(h2Body) + + for _, fromAgent := range h1Memories { + if fromAgent != "a-h1" { + t.Errorf("H1 saw memory from %q (cross-owner leak)", fromAgent) + } + } + for _, fromAgent := range h2Memories { + if fromAgent != "a-h2" { + t.Errorf("H2 saw memory from %q (cross-owner leak)", fromAgent) + } + } + + // Disjoint sets: no h1 memory id may appear in h2's response. + h1IDs := extractMemoryIDs(h1Body) + h2IDs := extractMemoryIDs(h2Body) + for id := range h1IDs { + if _, dup := h2IDs[id]; dup { + t.Errorf("memory id %d leaked across owners", id) + } + } +} + +func seedMemory(t *testing.T, db *sql.DB, channelID int64, fromAgent, body string) { + t.Helper() + convRes, err := db.Exec( + `INSERT INTO conversations (created_by, channel_id) VALUES (?, ?)`, + fromAgent, channelID, + ) + if err != nil { + t.Fatalf("seed conversation: %v", err) + } + convID, _ := convRes.LastInsertId() + if _, err := db.Exec( + `INSERT INTO messages (conversation_id, from_agent, channel_id, body, priority, status, metadata) + VALUES (?, ?, ?, ?, 5, 'pending', '{}')`, + convID, fromAgent, channelID, body, + ); err != nil { + t.Fatalf("seed message: %v", err) + } +} + +func unmarshalText(t *testing.T, res *mcplib.CallToolResult) map[string]any { + t.Helper() + if res == nil || len(res.Content) != 1 { + t.Fatalf("unexpected result: %+v", res) + } + tc, ok := res.Content[0].(mcplib.TextContent) + if !ok { + t.Fatalf("not text content: %T", res.Content[0]) + } + var m map[string]any + if err := json.Unmarshal([]byte(tc.Text), &m); err != nil { + t.Fatalf("unmarshal: %v: %s", err, tc.Text) + } + return m +} + +func extractMemoryFromAgents(body map[string]any) []string { + rc, ok := body["relevant_context"].(map[string]any) + if !ok { + return nil + } + mems, ok := rc["memories"].([]any) + if !ok { + return nil + } + out := make([]string, 0, len(mems)) + for _, raw := range mems { + m, ok := raw.(map[string]any) + if !ok { + continue + } + if fa, ok := m["from_agent"].(string); ok { + out = append(out, fa) + } + } + return out +} + +func extractMemoryIDs(body map[string]any) map[int64]struct{} { + rc, ok := body["relevant_context"].(map[string]any) + if !ok { + return map[int64]struct{}{} + } + mems, ok := rc["memories"].([]any) + if !ok { + return map[int64]struct{}{} + } + out := map[int64]struct{}{} + for _, raw := range mems { + m, ok := raw.(map[string]any) + if !ok { + continue + } + if id, ok := m["id"].(float64); ok { + out[int64(id)] = struct{}{} + } + } + return out +} diff --git a/internal/mcp/injection_wrap.go b/internal/mcp/injection_wrap.go new file mode 100644 index 0000000..f5dd298 --- /dev/null +++ b/internal/mcp/injection_wrap.go @@ -0,0 +1,251 @@ +// Proactive-memory injection middleware — wraps MCP tool handlers so +// that successful JSON responses gain a `relevant_context` field per +// `specs/020-proactive-memory-dream-worker/contracts/mcp-injection.md`. +package mcp + +import ( + "context" + "encoding/json" + "log/slog" + + mcplib "github.com/mark3labs/mcp-go/mcp" + "github.com/mark3labs/mcp-go/server" + + "github.com/synapbus/synapbus/internal/agents" + "github.com/synapbus/synapbus/internal/messaging" + "github.com/synapbus/synapbus/internal/search" +) + +// ToolHandler is the mcp-go tool handler signature. Re-exported as an +// alias so the wrapper signature reads cleanly at registration sites. +type ToolHandler = server.ToolHandlerFunc + +// QuerySource derives the retrieval query for one tool invocation. It +// is given the inner handler's parsed args (best-effort: nil when args +// don't fit map[string]any) and the inner handler's parsed JSON result +// (nil on error or non-JSON). It must be cheap; called on every wrapped +// tool call. +type QuerySource func(ctx context.Context, toolName string, args map[string]any, result map[string]any) string + +// WrapConfig parameterizes WrapInjection. +type WrapConfig struct { + // Cfg is the messaging.MemoryConfig snapshot taken at server + // startup. When Cfg.InjectionEnabled is false, WrapInjection + // returns the handler unchanged. + Cfg messaging.MemoryConfig + + // SearchSvc drives retrieval. Required when Cfg.InjectionEnabled. + SearchSvc *search.Service + + // Injections is the 24h audit ring. May be nil — Record errors are + // logged and the wrapper continues. + Injections *messaging.MemoryInjections + + // QuerySource derives the retrieval query for this tool. Required. + QuerySource QuerySource + + // IncludeCore is true for session-start tools (e.g. my_status). + // Only those get the per-(owner, agent) core-memory blob injected. + IncludeCore bool + + // CoreProvider is consulted when IncludeCore=true. May be nil + // (US2 not yet wired) — wrapper still functions, just skips core. + CoreProvider search.CoreMemoryProvider + + // Logger is used for non-fatal failures. Defaults to slog.Default. + Logger *slog.Logger +} + +// WrapInjection returns a ToolHandler that wraps `inner` with the +// proactive-memory injection middleware described in +// `contracts/mcp-injection.md`. +// +// When Cfg.InjectionEnabled is false, the original handler is returned +// unchanged so the response payload exactly matches the pre-feature +// shape (FR-012, SC-009). +// +// Otherwise, the wrapper: +// 1. Runs the inner handler. +// 2. If the result is an error or not a single TextContent of JSON +// object shape, returns the result unchanged. +// 3. Builds a ContextPacket via search.BuildContextPacket using the +// query derived from cfg.QuerySource. +// 4. If the packet is non-empty (>=1 memory or core memory set), merges +// `relevant_context: ` into the JSON body and re-marshals. +// 5. Records the injection to the 24h audit ring asynchronously. +func WrapInjection(inner ToolHandler, toolName string, cfg WrapConfig) ToolHandler { + if !cfg.Cfg.InjectionEnabled { + return inner + } + logger := cfg.Logger + if logger == nil { + logger = slog.Default().With("component", "mcp-injection") + } + return func(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) { + res, err := inner(ctx, req) + if err != nil { + return res, err + } + if res == nil || res.IsError { + return res, nil + } + + // Locate the JSON text content. Non-JSON or multi-content + // payloads pass through unchanged. + idx, text, ok := singleJSONText(res) + if !ok { + return res, nil + } + var body map[string]any + if err := json.Unmarshal([]byte(text), &body); err != nil { + return res, nil + } + + agent, ok := callerAgent(ctx) + if !ok || agent == nil { + // No identity → no owner scope → no injection. + return res, nil + } + + // Derive the retrieval query. nil args is fine; nil result is + // fine — the source decides what to do. + argsMap, _ := req.Params.Arguments.(map[string]any) + query := "" + if cfg.QuerySource != nil { + query = cfg.QuerySource(ctx, toolName, argsMap, body) + } + + opts := search.InjectionOpts{ + BudgetTokens: cfg.Cfg.InjectionBudgetTokens, + MaxItems: cfg.Cfg.InjectionMaxItems, + MinScore: cfg.Cfg.InjectionMinScore, + IncludeCore: cfg.IncludeCore, + CoreProvider: cfg.CoreProvider, + } + pkt, err := search.BuildContextPacket(ctx, cfg.SearchSvc, agent, query, opts) + if err != nil { + logger.Debug("build context packet failed", "tool", toolName, "error", err) + return res, nil + } + if pkt == nil { + // Empty packet → omit `relevant_context` entirely. + return res, nil + } + if len(pkt.Memories) == 0 && pkt.CoreMemory == "" { + return res, nil + } + + body["relevant_context"] = pkt + + merged, err := json.Marshal(body) + if err != nil { + logger.Debug("re-marshal failed", "tool", toolName, "error", err) + return res, nil + } + res.Content[idx] = mcplib.TextContent{Type: "text", Text: string(merged)} + + // Audit-ring write is best-effort; never blocks the response. + recordInjection(cfg.Injections, logger, agent, toolName, pkt) + + return res, nil + } +} + +// singleJSONText reports the index of the single TextContent in `res` +// when its Text is a JSON object. Anything else (multiple contents, +// non-text, non-object JSON) returns ok=false → pass through. +func singleJSONText(res *mcplib.CallToolResult) (int, string, bool) { + if res == nil || len(res.Content) != 1 { + return 0, "", false + } + tc, ok := res.Content[0].(mcplib.TextContent) + if !ok { + return 0, "", false + } + // Quick sanity check that the text starts with `{` — avoids + // allocating a map for a known-non-object payload. + for i := 0; i < len(tc.Text); i++ { + switch tc.Text[i] { + case ' ', '\t', '\n', '\r': + continue + case '{': + return 0, tc.Text, true + default: + return 0, "", false + } + } + return 0, "", false +} + +// callerAgent unpacks *agents.Agent from the request context. Uses the +// agents middleware ContextWithAgent, populated by the auth path. +func callerAgent(ctx context.Context) (*agents.Agent, bool) { + return agents.AgentFromContext(ctx) +} + +// recordInjection writes one audit-ring row. Runs in a fresh goroutine +// so it cannot block the response, but inherits a detached context with +// a short timeout via the inner call site. Failures are logged at debug +// level since they're non-fatal for the request. +func recordInjection(store *messaging.MemoryInjections, logger *slog.Logger, agent *agents.Agent, toolName string, pkt *search.ContextPacket) { + if store == nil || agent == nil || pkt == nil { + return + } + ids := make([]int64, 0, len(pkt.Memories)) + for _, m := range pkt.Memories { + ids = append(ids, m.ID) + } + rec := messaging.InjectionRecord{ + OwnerID: ownerIDString(agent.OwnerID), + AgentName: agent.Name, + ToolName: toolName, + PacketSizeChars: pkt.PacketChars, + PacketItemsCount: len(pkt.Memories), + MessageIDs: ids, + CoreBlobIncluded: pkt.CoreMemory != "", + } + go func() { + // Detached background context: the request context may already + // be canceled by the time this goroutine runs. + ctx := context.Background() + if err := store.Record(ctx, rec); err != nil { + logger.Debug("audit-ring insert failed", "tool", toolName, "error", err) + } + }() +} + +func ownerIDString(id int64) string { + if id == 0 { + return "" + } + // strconv.FormatInt is faster than fmt.Sprintf; mirror what + // agents.OwnerFor produces so comparisons in the search layer line + // up. + return formatInt64(id) +} + +// formatInt64 is a tiny helper to avoid pulling strconv into the public +// surface area for one line. +func formatInt64(v int64) string { + const digits = "0123456789" + if v == 0 { + return "0" + } + neg := false + if v < 0 { + neg = true + v = -v + } + var buf [20]byte + i := len(buf) + for v > 0 { + i-- + buf[i] = digits[v%10] + v /= 10 + } + if neg { + i-- + buf[i] = '-' + } + return string(buf[i:]) +} diff --git a/internal/mcp/injection_wrap_test.go b/internal/mcp/injection_wrap_test.go new file mode 100644 index 0000000..91058bf --- /dev/null +++ b/internal/mcp/injection_wrap_test.go @@ -0,0 +1,218 @@ +package mcp + +import ( + "context" + "encoding/json" + "testing" + + mcplib "github.com/mark3labs/mcp-go/mcp" + + "github.com/synapbus/synapbus/internal/agents" + "github.com/synapbus/synapbus/internal/messaging" + "github.com/synapbus/synapbus/internal/search" +) + +// stubBuilder lets us bypass the real search service and short-circuit +// BuildContextPacket to whatever ContextPacket we want for the test. +// The wrap layer doesn't expose a builder seam (it calls +// search.BuildContextPacket directly), so we instead drive the wrap end +// to end with a real (empty) *search.Service and a stub CoreProvider +// that emits a packet when IncludeCore is true. + +type stubCoreProvider struct{ blob string } + +func (s *stubCoreProvider) Get(_ context.Context, _, _ string) (string, error) { + return s.blob, nil +} + +func stubHandler(body map[string]any) ToolHandler { + return func(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) { + b, _ := json.Marshal(body) + return mcplib.NewToolResultText(string(b)), nil + } +} + +func extractJSON(t *testing.T, res *mcplib.CallToolResult) map[string]any { + t.Helper() + if res == nil { + t.Fatal("nil result") + } + if len(res.Content) != 1 { + t.Fatalf("expected 1 content, got %d", len(res.Content)) + } + tc, ok := res.Content[0].(mcplib.TextContent) + if !ok { + t.Fatalf("content not TextContent: %T", res.Content[0]) + } + var m map[string]any + if err := json.Unmarshal([]byte(tc.Text), &m); err != nil { + t.Fatalf("non-JSON content: %v: %q", err, tc.Text) + } + return m +} + +func TestWrapInjection_DisabledReturnsHandlerUnchanged(t *testing.T) { + inner := stubHandler(map[string]any{"hello": "world"}) + cfg := WrapConfig{ + Cfg: messaging.MemoryConfig{InjectionEnabled: false}, + } + wrapped := WrapInjection(inner, "my_status", cfg) + + res, err := wrapped(context.Background(), mcplib.CallToolRequest{}) + if err != nil { + t.Fatalf("wrapped: %v", err) + } + body := extractJSON(t, res) + if _, has := body["relevant_context"]; has { + t.Error("relevant_context attached despite disabled config") + } + if body["hello"] != "world" { + t.Errorf("inner body mutated: %v", body) + } +} + +func TestWrapInjection_EmptyMemoriesAndNoCore_OmitsField(t *testing.T) { + // IncludeCore=false and no memories in the DB → packet is nil → no + // relevant_context field on the response. + cfg := WrapConfig{ + Cfg: messaging.MemoryConfig{ + InjectionEnabled: true, + InjectionBudgetTokens: 500, + InjectionMaxItems: 5, + InjectionMinScore: 0.25, + }, + SearchSvc: nil, // BuildContextPacket short-circuits when query=="" → returns nil + QuerySource: func(_ context.Context, _ string, _ map[string]any, _ map[string]any) string { return "" }, + } + // SearchSvc nil + empty query forces BuildContextPacket through the + // "no retrieval" path. But it'll still try a core fetch (skipped: + // IncludeCore=false). With no provider and no retrieval → nil packet. + inner := stubHandler(map[string]any{"ok": true}) + wrapped := WrapInjection(inner, "send_message", cfg) + + // Inject a caller agent into the context so the wrapper does not + // bail out at the identity check. + ctx := agents.ContextWithAgent(context.Background(), &agents.Agent{Name: "a1", OwnerID: 1}) + + res, err := wrapped(ctx, mcplib.CallToolRequest{}) + if err != nil { + t.Fatalf("wrapped: %v", err) + } + body := extractJSON(t, res) + if _, has := body["relevant_context"]; has { + t.Errorf("relevant_context attached when memories+core empty: %v", body["relevant_context"]) + } +} + +func TestWrapInjection_AppendsRelevantContext_CoreOnly(t *testing.T) { + cfg := WrapConfig{ + Cfg: messaging.MemoryConfig{ + InjectionEnabled: true, + InjectionBudgetTokens: 500, + InjectionMaxItems: 5, + InjectionMinScore: 0.25, + }, + SearchSvc: nil, // query="" guarantees no retrieval attempt + IncludeCore: true, + CoreProvider: &stubCoreProvider{blob: "I am a1."}, + QuerySource: func(_ context.Context, _ string, _ map[string]any, _ map[string]any) string { return "" }, + } + inner := stubHandler(map[string]any{"agent": "a1"}) + wrapped := WrapInjection(inner, "my_status", cfg) + + ctx := agents.ContextWithAgent(context.Background(), &agents.Agent{Name: "a1", OwnerID: 1}) + res, err := wrapped(ctx, mcplib.CallToolRequest{}) + if err != nil { + t.Fatalf("wrapped: %v", err) + } + body := extractJSON(t, res) + rc, has := body["relevant_context"].(map[string]any) + if !has { + t.Fatalf("relevant_context missing: %+v", body) + } + if rc["core_memory"] != "I am a1." { + t.Errorf("core_memory wrong: %v", rc["core_memory"]) + } + if mems, ok := rc["memories"].([]any); !ok || len(mems) != 0 { + t.Errorf("memories should be empty slice when only core is set: %v", rc["memories"]) + } + // PacketChars at minimum the length of the core blob. + if got, ok := rc["packet_chars"].(float64); !ok || int(got) < len("I am a1.") { + t.Errorf("packet_chars looks wrong: %v", rc["packet_chars"]) + } +} + +func TestWrapInjection_NonJSONResultPassesThrough(t *testing.T) { + inner := func(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) { + return mcplib.NewToolResultText("not json"), nil + } + cfg := WrapConfig{ + Cfg: messaging.MemoryConfig{ + InjectionEnabled: true, + InjectionBudgetTokens: 500, + InjectionMaxItems: 5, + }, + QuerySource: func(_ context.Context, _ string, _ map[string]any, _ map[string]any) string { return "" }, + } + wrapped := WrapInjection(inner, "execute", cfg) + ctx := agents.ContextWithAgent(context.Background(), &agents.Agent{Name: "a1", OwnerID: 1}) + res, err := wrapped(ctx, mcplib.CallToolRequest{}) + if err != nil { + t.Fatalf("wrapped: %v", err) + } + if res == nil || len(res.Content) != 1 { + t.Fatalf("unexpected result shape: %+v", res) + } + tc, ok := res.Content[0].(mcplib.TextContent) + if !ok || tc.Text != "not json" { + t.Errorf("non-JSON result mutated: %+v", res.Content[0]) + } +} + +func TestWrapInjection_NoAgentInContext_PassesThrough(t *testing.T) { + cfg := WrapConfig{ + Cfg: messaging.MemoryConfig{ + InjectionEnabled: true, + InjectionBudgetTokens: 500, + InjectionMaxItems: 5, + }, + IncludeCore: true, + CoreProvider: &stubCoreProvider{blob: "blob"}, + QuerySource: func(_ context.Context, _ string, _ map[string]any, _ map[string]any) string { return "" }, + } + wrapped := WrapInjection(stubHandler(map[string]any{"x": 1}), "my_status", cfg) + res, err := wrapped(context.Background(), mcplib.CallToolRequest{}) + if err != nil { + t.Fatalf("wrapped: %v", err) + } + body := extractJSON(t, res) + if _, has := body["relevant_context"]; has { + t.Error("relevant_context attached despite no caller agent") + } +} + +// Ensure ContextPacket as the value carries through json round-trip +// (it's used directly as a map entry via body["relevant_context"] = pkt). +func TestWrapInjection_ContextPacketJSONShape(t *testing.T) { + pkt := &search.ContextPacket{ + Memories: []search.MemoryItem{}, + CoreMemory: "core", + PacketChars: 4, + PacketTokenEstimate: 1, + RetrievalQuery: "", + SearchMode: "auto", + } + b, err := json.Marshal(pkt) + if err != nil { + t.Fatalf("marshal: %v", err) + } + var rt map[string]any + if err := json.Unmarshal(b, &rt); err != nil { + t.Fatalf("unmarshal: %v", err) + } + for _, k := range []string{"memories", "core_memory", "packet_chars", "packet_token_estimate", "retrieval_query", "search_mode"} { + if _, has := rt[k]; !has { + t.Errorf("missing JSON key %q", k) + } + } +} diff --git a/internal/mcp/server.go b/internal/mcp/server.go index fd1b0df..255715a 100644 --- a/internal/mcp/server.go +++ b/internal/mcp/server.go @@ -207,8 +207,9 @@ func NewMCPServer( } // WireGoalsTools registers the spec-018 tool surface (create_goal, -// propose_task_tree, propose_agent, claim_task, request_resource, -// list_resources) on the MCP server. Must be called after NewMCPServer. +// propose_task_tree, claim_task, request_resource, list_resources, +// complete_goal) on the MCP server. Must be called after NewMCPServer. +// Note: propose_agent was removed in the internal-only mode change. func (s *MCPServer) WireGoalsTools(r *GoalsToolRegistrar) { if r == nil || s.mcpServer == nil { return diff --git a/internal/mcp/tools_hybrid.go b/internal/mcp/tools_hybrid.go index 7924d8e..2dc50f1 100644 --- a/internal/mcp/tools_hybrid.go +++ b/internal/mcp/tools_hybrid.go @@ -44,6 +44,26 @@ type HybridToolRegistrar struct { db *sql.DB queryExecutor *agentquery.Executor logger *slog.Logger + + // Injection (feature 020). When injectionCfg.Cfg.InjectionEnabled + // is false (the default), WrapInjection returns handlers unchanged + // so existing tool response shapes are preserved bit-for-bit. + injectionCfg messaging.MemoryConfig + memoryInjections *messaging.MemoryInjections + coreProvider search.CoreMemoryProvider +} + +// SetInjection wires the proactive-memory injection middleware into +// every eligible MCP tool registered by RegisterAllOnServer. Call this +// after NewHybridToolRegistrar and before RegisterAllOnServer. +// +// `coreProvider` is consulted only on session-start tools (currently +// `my_status`). May be nil when US2 has not yet been wired — the +// wrapper simply skips the core-memory hook in that case. +func (h *HybridToolRegistrar) SetInjection(cfg messaging.MemoryConfig, store *messaging.MemoryInjections, coreProvider search.CoreMemoryProvider) { + h.injectionCfg = cfg + h.memoryInjections = store + h.coreProvider = coreProvider } // SetMarketplaceService attaches the marketplace service for the 5 new @@ -93,14 +113,104 @@ func NewHybridToolRegistrar( } // RegisterAllOnServer registers all hybrid tools on an mcp-go MCPServer. +// +// When proactive-memory injection is enabled (via SetInjection), the +// eligible tool handlers are wrapped with WrapInjection so their JSON +// responses gain a `relevant_context` field per +// `contracts/mcp-injection.md`. Tools NOT in the eligible set +// (currently `get_replies`) are registered unchanged. func (h *HybridToolRegistrar) RegisterAllOnServer(s *server.MCPServer) { - s.AddTool(h.myStatusTool(), h.handleMyStatus) - s.AddTool(h.sendMessageTool(), h.handleSendMessage) - s.AddTool(h.searchTool(), h.handleSearch) - s.AddTool(h.executeTool(), h.handleExecute) + s.AddTool(h.myStatusTool(), h.wrap("my_status", h.handleMyStatus, true)) + s.AddTool(h.sendMessageTool(), h.wrap("send_message", h.handleSendMessage, false)) + s.AddTool(h.searchTool(), h.wrap("search", h.handleSearch, false)) + s.AddTool(h.executeTool(), h.wrap("execute", h.handleExecute, false)) s.AddTool(h.getRepliesTool(), h.handleGetReplies) - h.logger.Info("hybrid MCP tools registered", "count", 5) + h.logger.Info("hybrid MCP tools registered", + "count", 5, + "injection_enabled", h.injectionCfg.InjectionEnabled, + ) +} + +// wrap applies the proactive-memory WrapInjection middleware to one +// tool handler. When InjectionEnabled is false (default), wrap returns +// the original handler unchanged. `includeCore` is true only for +// session-start tools (my_status today). +func (h *HybridToolRegistrar) wrap(toolName string, inner ToolHandler, includeCore bool) ToolHandler { + if !h.injectionCfg.InjectionEnabled { + return inner + } + cfg := WrapConfig{ + Cfg: h.injectionCfg, + SearchSvc: h.searchService, + Injections: h.memoryInjections, + IncludeCore: includeCore, + CoreProvider: h.coreProvider, + QuerySource: querySourceFor(toolName), + Logger: h.logger, + } + return WrapInjection(inner, toolName, cfg) +} + +// querySourceFor returns the QuerySource closure for the given tool. +// The contract (`contracts/mcp-injection.md`) prescribes the retrieval +// query per tool: +// +// - my_status: "" (fallback when nothing else) +// - send_message: body of the sent message +// - search: the user's query argument +// - execute: stringified args (best-effort) +// +// Tools not registered as MCP tools here (`claim_messages`, +// `read_inbox`, `read_channel`) are exercised via the `execute` bridge +// and inherit the `execute` query source. +func querySourceFor(toolName string) QuerySource { + switch toolName { + case "my_status": + return func(_ context.Context, _ string, _ map[string]any, _ map[string]any) string { + // FR-009: when there's no explicit query, use recent owner + // activity. The retrieval layer interprets the empty string + // as "no query" and falls back to recency-ordered FTS. + return "" + } + case "send_message": + return func(_ context.Context, _ string, args map[string]any, _ map[string]any) string { + if args == nil { + return "" + } + if body, ok := args["body"].(string); ok { + return body + } + return "" + } + case "search": + return func(_ context.Context, _ string, args map[string]any, _ map[string]any) string { + if args == nil { + return "" + } + if q, ok := args["query"].(string); ok { + return q + } + return "" + } + case "execute": + return func(_ context.Context, _ string, args map[string]any, _ map[string]any) string { + if args == nil { + return "" + } + // Best-effort: use the `code` argument verbatim. It is the + // only required input and reliably reflects what the agent + // is about to do. + if code, ok := args["code"].(string); ok { + return code + } + return "" + } + default: + return func(_ context.Context, _ string, _ map[string]any, _ map[string]any) string { + return "" + } + } } // --- Tool Definitions --- @@ -118,7 +228,7 @@ func (h *HybridToolRegistrar) sendMessageTool() mcplib.Tool { mcplib.WithString("channel", mcplib.Description("Channel name or numeric ID for channel messages")), mcplib.WithString("body", mcplib.Description("Message body text"), mcplib.Required()), mcplib.WithString("subject", mcplib.Description("Conversation subject (optional)")), - mcplib.WithNumber("priority", mcplib.Description("Message priority (1-10, default 5)"), mcplib.Min(1), mcplib.Max(10)), + mcplib.WithNumber("priority", mcplib.Description("Message priority (1-10, default 5)")), mcplib.WithString("metadata", mcplib.Description("JSON metadata object (optional)")), mcplib.WithNumber("reply_to", mcplib.Description("ID of the parent message to reply to. Creates a threaded reply. Always use reply_to when responding to a message that is itself a thread reply, to keep conversations organized.")), mcplib.WithString("attachments", mcplib.Description("Comma-separated list of attachment hashes to link to this message. Upload attachments first using the upload_attachment action via the execute tool.")), diff --git a/internal/messaging/memory_injections.go b/internal/messaging/memory_injections.go new file mode 100644 index 0000000..42f7f85 --- /dev/null +++ b/internal/messaging/memory_injections.go @@ -0,0 +1,147 @@ +package messaging + +import ( + "context" + "database/sql" + "encoding/json" + "fmt" + "time" +) + +// InjectionRecord is one row in the `memory_injections` 24-hour audit +// ring. Each row captures what was attached to a single MCP tool +// response so the owner can later answer "why did my agent know this?" +// via the recent-injections debug surface (FR-025). +type InjectionRecord struct { + ID int64 `json:"id"` + OwnerID string `json:"owner_id"` + AgentName string `json:"agent_name"` + ToolName string `json:"tool_name"` + PacketSizeChars int `json:"packet_size_chars"` + PacketItemsCount int `json:"packet_items_count"` + MessageIDs []int64 `json:"message_ids"` + CoreBlobIncluded bool `json:"core_blob_included"` + CreatedAt time.Time `json:"created_at"` +} + +// MemoryInjections is the audit-ring store for proactive injection +// (data-model.md §`memory_injections`). Each row is best-effort: +// failures are non-fatal for the injection request — callers should +// log and continue. +type MemoryInjections struct { + db *sql.DB +} + +// NewMemoryInjections wraps a *sql.DB. +func NewMemoryInjections(db *sql.DB) *MemoryInjections { + return &MemoryInjections{db: db} +} + +// Record inserts one injection row. `row.MessageIDs` is JSON-encoded. +// `created_at` defaults to CURRENT_TIMESTAMP when zero. +func (s *MemoryInjections) Record(ctx context.Context, row InjectionRecord) error { + if s == nil || s.db == nil { + return nil + } + ids := row.MessageIDs + if ids == nil { + ids = []int64{} + } + b, err := json.Marshal(ids) + if err != nil { + return fmt.Errorf("memory_injections: marshal message_ids: %w", err) + } + + if row.CreatedAt.IsZero() { + _, err = s.db.ExecContext(ctx, + `INSERT INTO memory_injections + (owner_id, agent_name, tool_name, packet_size_chars, + packet_items_count, message_ids, core_blob_included) + VALUES (?, ?, ?, ?, ?, ?, ?)`, + row.OwnerID, row.AgentName, row.ToolName, + row.PacketSizeChars, row.PacketItemsCount, string(b), + row.CoreBlobIncluded, + ) + } else { + _, err = s.db.ExecContext(ctx, + `INSERT INTO memory_injections + (owner_id, agent_name, tool_name, packet_size_chars, + packet_items_count, message_ids, core_blob_included, created_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?)`, + row.OwnerID, row.AgentName, row.ToolName, + row.PacketSizeChars, row.PacketItemsCount, string(b), + row.CoreBlobIncluded, row.CreatedAt.UTC(), + ) + } + if err != nil { + return fmt.Errorf("memory_injections: insert: %w", err) + } + return nil +} + +// Cleanup deletes rows older than `olderThan` ago. Returns the number +// of rows removed. Safe to call from a periodic ticker. +func (s *MemoryInjections) Cleanup(ctx context.Context, olderThan time.Duration) (int64, error) { + if s == nil || s.db == nil { + return 0, nil + } + if olderThan <= 0 { + return 0, nil + } + cutoff := time.Now().Add(-olderThan).UTC() + res, err := s.db.ExecContext(ctx, + `DELETE FROM memory_injections WHERE created_at < ?`, cutoff, + ) + if err != nil { + return 0, fmt.Errorf("memory_injections: cleanup: %w", err) + } + affected, _ := res.RowsAffected() + return affected, nil +} + +// ListRecent returns the most recent injections for one owner, newest +// first, up to `limit`. +func (s *MemoryInjections) ListRecent(ctx context.Context, ownerID string, limit int) ([]InjectionRecord, error) { + if s == nil || s.db == nil { + return nil, nil + } + if limit <= 0 { + limit = 50 + } + rows, err := s.db.QueryContext(ctx, + `SELECT id, owner_id, agent_name, tool_name, packet_size_chars, + packet_items_count, message_ids, core_blob_included, created_at + FROM memory_injections + WHERE owner_id = ? + ORDER BY created_at DESC, id DESC + LIMIT ?`, ownerID, limit, + ) + if err != nil { + return nil, fmt.Errorf("memory_injections: list recent: %w", err) + } + defer rows.Close() + + var out []InjectionRecord + for rows.Next() { + var rec InjectionRecord + var idsJSON string + if err := rows.Scan( + &rec.ID, &rec.OwnerID, &rec.AgentName, &rec.ToolName, + &rec.PacketSizeChars, &rec.PacketItemsCount, &idsJSON, + &rec.CoreBlobIncluded, &rec.CreatedAt, + ); err != nil { + return nil, fmt.Errorf("memory_injections: scan: %w", err) + } + if idsJSON == "" { + rec.MessageIDs = []int64{} + } else if err := json.Unmarshal([]byte(idsJSON), &rec.MessageIDs); err != nil { + // Corrupt row — surface as empty rather than fail the whole listing. + rec.MessageIDs = []int64{} + } + out = append(out, rec) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("memory_injections: iterate: %w", err) + } + return out, nil +} diff --git a/internal/messaging/memory_injections_test.go b/internal/messaging/memory_injections_test.go new file mode 100644 index 0000000..45874e0 --- /dev/null +++ b/internal/messaging/memory_injections_test.go @@ -0,0 +1,166 @@ +package messaging + +import ( + "context" + "testing" + "time" + + _ "modernc.org/sqlite" +) + +func TestMemoryInjections_RecordAndList(t *testing.T) { + db := newTestDB(t) + store := NewMemoryInjections(db) + ctx := context.Background() + + rec := InjectionRecord{ + OwnerID: "1", + AgentName: "research-mcpproxy", + ToolName: "my_status", + PacketSizeChars: 412, + PacketItemsCount: 3, + MessageIDs: []int64{1, 2, 3}, + CoreBlobIncluded: true, + } + if err := store.Record(ctx, rec); err != nil { + t.Fatalf("Record: %v", err) + } + + got, err := store.ListRecent(ctx, "1", 10) + if err != nil { + t.Fatalf("ListRecent: %v", err) + } + if len(got) != 1 { + t.Fatalf("ListRecent returned %d rows, want 1", len(got)) + } + if got[0].PacketSizeChars != 412 || got[0].PacketItemsCount != 3 { + t.Errorf("packet counters mismatch: %+v", got[0]) + } + if len(got[0].MessageIDs) != 3 || got[0].MessageIDs[0] != 1 { + t.Errorf("MessageIDs round-trip failed: %v", got[0].MessageIDs) + } + if !got[0].CoreBlobIncluded { + t.Error("CoreBlobIncluded round-trip failed") + } +} + +func TestMemoryInjections_CleanupOnlyOldRows(t *testing.T) { + db := newTestDB(t) + store := NewMemoryInjections(db) + ctx := context.Background() + + now := time.Now().UTC() + tests := []struct { + name string + offset time.Duration + wantKept bool + }{ + {"3 days old", -72 * time.Hour, false}, + {"36 hours old", -36 * time.Hour, false}, + {"23 hours old", -23 * time.Hour, true}, + {"30 minutes old", -30 * time.Minute, true}, + {"current", 0, true}, + } + + // Seed 50 rows: 10 per bucket so we exercise the DELETE plan. + for _, tc := range tests { + for i := 0; i < 10; i++ { + rec := InjectionRecord{ + OwnerID: "1", + AgentName: "a", + ToolName: "my_status", + PacketSizeChars: 100, + PacketItemsCount: 1, + MessageIDs: []int64{int64(i)}, + CreatedAt: now.Add(tc.offset), + } + if err := store.Record(ctx, rec); err != nil { + t.Fatalf("Record %s: %v", tc.name, err) + } + } + } + + deleted, err := store.Cleanup(ctx, 24*time.Hour) + if err != nil { + t.Fatalf("Cleanup: %v", err) + } + // 2 buckets older than 24h × 10 rows = 20 expected deletions. + if deleted != 20 { + t.Errorf("Cleanup removed %d rows, want 20", deleted) + } + + remaining, err := store.ListRecent(ctx, "1", 100) + if err != nil { + t.Fatalf("ListRecent: %v", err) + } + if len(remaining) != 30 { + t.Errorf("after cleanup got %d rows, want 30", len(remaining)) + } +} + +func TestMemoryInjections_ListRecent_OwnerScoped(t *testing.T) { + db := newTestDB(t) + store := NewMemoryInjections(db) + ctx := context.Background() + + for _, owner := range []string{"1", "2", "3"} { + for i := 0; i < 5; i++ { + if err := store.Record(ctx, InjectionRecord{ + OwnerID: owner, + AgentName: "a", + ToolName: "my_status", + PacketSizeChars: 100, + PacketItemsCount: 1, + MessageIDs: []int64{int64(i)}, + }); err != nil { + t.Fatalf("Record owner=%s i=%d: %v", owner, i, err) + } + } + } + + got, err := store.ListRecent(ctx, "2", 100) + if err != nil { + t.Fatalf("ListRecent: %v", err) + } + if len(got) != 5 { + t.Fatalf("owner=2 got %d rows, want 5", len(got)) + } + for _, r := range got { + if r.OwnerID != "2" { + t.Errorf("found leaked row OwnerID=%q in owner=2 listing", r.OwnerID) + } + } +} + +func TestMemoryInjections_ListRecent_LimitOrdering(t *testing.T) { + db := newTestDB(t) + store := NewMemoryInjections(db) + ctx := context.Background() + + base := time.Now().UTC().Add(-1 * time.Hour) + for i := 0; i < 10; i++ { + if err := store.Record(ctx, InjectionRecord{ + OwnerID: "1", + AgentName: "a", + ToolName: "my_status", + PacketSizeChars: 100, + PacketItemsCount: 1, + MessageIDs: []int64{int64(i)}, + CreatedAt: base.Add(time.Duration(i) * time.Minute), + }); err != nil { + t.Fatalf("Record %d: %v", i, err) + } + } + + got, err := store.ListRecent(ctx, "1", 3) + if err != nil { + t.Fatalf("ListRecent: %v", err) + } + if len(got) != 3 { + t.Fatalf("ListRecent limit=3 returned %d", len(got)) + } + // Newest first: i=9, 8, 7. Validate MessageIDs[0] descends. + if got[0].MessageIDs[0] != 9 || got[1].MessageIDs[0] != 8 || got[2].MessageIDs[0] != 7 { + t.Errorf("ordering wrong: %d, %d, %d", got[0].MessageIDs[0], got[1].MessageIDs[0], got[2].MessageIDs[0]) + } +} diff --git a/internal/messaging/stalemate.go b/internal/messaging/stalemate.go index 5fc5fdb..d92d44f 100644 --- a/internal/messaging/stalemate.go +++ b/internal/messaging/stalemate.go @@ -14,13 +14,15 @@ import ( ) // StalemateConfig holds stalemate detection settings. +// +// The historical reminder/escalation knobs (ReminderAfter, EscalateAfter) +// were removed when SynapBus moved to internal-only mode (no human +// approval loop). See migration 027_remove_approval_noise.sql. type StalemateConfig struct { - // ProcessingTimeout is how long a message can stay in "processing" before auto-fail (default 24h). + // ProcessingTimeout is how long a message can stay in "processing" + // before auto-fail (default 24h). Protects the inbox queue from + // agents that crash after claiming a message. ProcessingTimeout time.Duration - // ReminderAfter is how long a pending DM waits before a system reminder is sent (default 4h). - ReminderAfter time.Duration - // EscalateAfter is how long a pending DM waits before escalation to #approvals (default 48h). - EscalateAfter time.Duration // Interval is how often the worker checks for stale messages (default 15m). Interval time.Duration } @@ -29,8 +31,6 @@ type StalemateConfig struct { func DefaultStalemateConfig() StalemateConfig { return StalemateConfig{ ProcessingTimeout: 24 * time.Hour, - ReminderAfter: 4 * time.Hour, - EscalateAfter: 48 * time.Hour, Interval: 15 * time.Minute, } } @@ -43,7 +43,6 @@ func parseDurationWithDays(s string) (time.Duration, error) { return 0, fmt.Errorf("empty duration string") } - // Try "Nd" format (days) if strings.HasSuffix(s, "d") { days, err := strconv.Atoi(strings.TrimSuffix(s, "d")) if err == nil && days > 0 { @@ -51,7 +50,6 @@ func parseDurationWithDays(s string) (time.Duration, error) { } } - // Try standard Go duration return time.ParseDuration(s) } @@ -64,16 +62,6 @@ func ParseStalemateConfig() StalemateConfig { cfg.ProcessingTimeout = d } } - if v := os.Getenv("SYNAPBUS_STALEMATE_REMINDER_AFTER"); v != "" { - if d, err := parseDurationWithDays(v); err == nil && d > 0 { - cfg.ReminderAfter = d - } - } - if v := os.Getenv("SYNAPBUS_STALEMATE_ESCALATE_AFTER"); v != "" { - if d, err := parseDurationWithDays(v); err == nil && d > 0 { - cfg.EscalateAfter = d - } - } if v := os.Getenv("SYNAPBUS_STALEMATE_INTERVAL"); v != "" { if d, err := parseDurationWithDays(v); err == nil && d > 0 { cfg.Interval = d @@ -83,32 +71,37 @@ func ParseStalemateConfig() StalemateConfig { return cfg } -// ChannelLookup provides channel lookup by name without importing the channels package. -type ChannelLookup interface { - // GetChannelIDByName returns a channel ID by name, or 0 if not found. - GetChannelIDByName(ctx context.Context, name string) (int64, error) -} - -// StalemateWorker periodically checks for and handles stale messages. +// StalemateWorker periodically auto-fails messages whose claim has timed out. type StalemateWorker struct { db *sql.DB msgService *MessagingService - channelLookup ChannelLookup config StalemateConfig logger *slog.Logger done chan struct{} wg sync.WaitGroup + + // memoryInjections, when non-nil, drives an hourly cleanup of the + // 24h proactive-injection audit ring (feature 020). Plumbed via + // SetMemoryInjections after construction so adding the feature is + // non-breaking for existing call sites. + memoryInjections *MemoryInjections + tickCount int +} + +// SetMemoryInjections registers the audit-ring store the worker will +// cleanup hourly. Pass nil to disable; safe to call before Start. +func (w *StalemateWorker) SetMemoryInjections(store *MemoryInjections) { + w.memoryInjections = store } // NewStalemateWorker creates a new stalemate detection worker. -func NewStalemateWorker(db *sql.DB, msgService *MessagingService, channelLookup ChannelLookup, config StalemateConfig) *StalemateWorker { +func NewStalemateWorker(db *sql.DB, msgService *MessagingService, config StalemateConfig) *StalemateWorker { return &StalemateWorker{ - db: db, - msgService: msgService, - channelLookup: channelLookup, - config: config, - logger: slog.Default().With("component", "stalemate-worker"), - done: make(chan struct{}), + db: db, + msgService: msgService, + config: config, + logger: slog.Default().With("component", "stalemate-worker"), + done: make(chan struct{}), } } @@ -120,8 +113,6 @@ func (w *StalemateWorker) Start() { w.logger.Info("stalemate worker started", "interval", w.config.Interval.String(), "processing_timeout", w.config.ProcessingTimeout.String(), - "reminder_after", w.config.ReminderAfter.String(), - "escalate_after", w.config.EscalateAfter.String(), ) ticker := time.NewTicker(w.config.Interval) @@ -132,6 +123,7 @@ func (w *StalemateWorker) Start() { case <-ticker.C: ctx, cancel := context.WithTimeout(context.Background(), 60*time.Second) w.checkStaleMessages(ctx) + w.maybeCleanupInjections(ctx) cancel() case <-w.done: w.logger.Info("stalemate worker stopped") @@ -147,23 +139,40 @@ func (w *StalemateWorker) Stop() { w.wg.Wait() } -// checkStaleMessages runs all stalemate checks. +// checkStaleMessages runs the auto-fail check for timed-out claimed messages. func (w *StalemateWorker) checkStaleMessages(ctx context.Context) { - failed := w.failTimedOutProcessing(ctx) - reminded := w.sendPendingReminders(ctx) - escalated := w.escalatePendingMessages(ctx) + if failed := w.failTimedOutProcessing(ctx); failed > 0 { + w.logger.Info("stalemate check complete", "auto_failed", failed) + } +} - // Phase 2: Workflow stalemate checks for channel messages - wfReminded, wfEscalated := w.checkWorkflowStalemates(ctx) - - if failed > 0 || reminded > 0 || escalated > 0 || wfReminded > 0 || wfEscalated > 0 { - w.logger.Info("stalemate check complete", - "auto_failed", failed, - "reminders_sent", reminded, - "escalations_sent", escalated, - "workflow_reminders", wfReminded, - "workflow_escalations", wfEscalated, - ) +// maybeCleanupInjections piggybacks an hourly cleanup of the 24h +// proactive-injection audit ring (feature 020) on the stalemate worker +// tick. With the default 15-minute Interval, the cleanup fires every +// 4th tick (i.e. ~1h). A no-op when SetMemoryInjections has not been +// called or the store is nil. +func (w *StalemateWorker) maybeCleanupInjections(ctx context.Context) { + if w.memoryInjections == nil { + return + } + w.tickCount++ + // Fire roughly hourly. With Interval=15m the modulus matches the + // 4-tick cadence called out in spec 020 T017. For non-default + // intervals it still fires hourly-ish on a best-effort basis. + ticksPerHour := int(time.Hour / w.config.Interval) + if ticksPerHour <= 0 { + ticksPerHour = 1 + } + if w.tickCount%ticksPerHour != 0 { + return + } + deleted, err := w.memoryInjections.Cleanup(ctx, 24*time.Hour) + if err != nil { + w.logger.Warn("memory_injections cleanup failed", "error", err) + return + } + if deleted > 0 { + w.logger.Info("memory_injections cleanup", "deleted", deleted) } } @@ -242,9 +251,6 @@ func (w *StalemateWorker) failTimedOutProcessing(ctx context.Context) int64 { } affected, _ := res.RowsAffected() if affected == 0 { - // Message was reclaimed, completed, or otherwise moved out of the - // stale window between SELECT and UPDATE. Skip silently — the - // next worker tick will re-evaluate. w.logger.Debug("stale processing message no longer stale; skipped", "message_id", dm.ID, "claimed_by", dm.ClaimedBy, @@ -261,646 +267,3 @@ func (w *StalemateWorker) failTimedOutProcessing(ctx context.Context) int64 { } return count } - -// sendPendingReminders sends system DM reminders for pending messages older than ReminderAfter. -func (w *StalemateWorker) sendPendingReminders(ctx context.Context) int64 { - cutoff := time.Now().Add(-w.config.ReminderAfter) - - rows, err := w.db.QueryContext(ctx, - `SELECT id, from_agent, to_agent, body, created_at - FROM messages - WHERE status = 'pending' - AND to_agent IS NOT NULL - AND to_agent != '' - AND from_agent != 'system' - AND to_agent != 'system' - AND created_at < ?`, - cutoff, - ) - if err != nil { - w.logger.Error("query pending reminder candidates failed", "error", err) - return 0 - } - defer rows.Close() - - type pendingMsg struct { - ID int64 - FromAgent string - ToAgent string - Body string - CreatedAt time.Time - } - - var pending []pendingMsg - for rows.Next() { - var pm pendingMsg - if err := rows.Scan(&pm.ID, &pm.FromAgent, &pm.ToAgent, &pm.Body, &pm.CreatedAt); err != nil { - w.logger.Error("scan pending message failed", "error", err) - continue - } - pending = append(pending, pm) - } - - count := int64(0) - for _, pm := range pending { - // Check if a reminder already exists for this message - if w.reminderExists(ctx, pm.ID, pm.ToAgent) { - continue - } - - age := formatAge(time.Since(pm.CreatedAt)) - truncBody := truncate(pm.Body, 100) - - body := fmt.Sprintf( - "**Reminder**: You have a pending message from %s (%s old). Message: \"%s\". Please claim and process it.", - pm.FromAgent, age, truncBody, - ) - - _, err := w.msgService.SendMessage(ctx, "system", pm.ToAgent, body, SendOptions{ - Subject: fmt.Sprintf("stalemate-reminder:%d", pm.ID), - Priority: 7, - Metadata: fmt.Sprintf(`{"stalemate_reminder_for":%d}`, pm.ID), - }) - if err != nil { - w.logger.Error("send stalemate reminder failed", - "message_id", pm.ID, - "to_agent", pm.ToAgent, - "error", err, - ) - continue - } - w.logger.Info("sent stalemate reminder", - "message_id", pm.ID, - "to_agent", pm.ToAgent, - "from_agent", pm.FromAgent, - "age", age, - ) - count++ - } - return count -} - -// escalatePendingMessages escalates pending messages older than EscalateAfter to #approvals. -func (w *StalemateWorker) escalatePendingMessages(ctx context.Context) int64 { - cutoff := time.Now().Add(-w.config.EscalateAfter) - - rows, err := w.db.QueryContext(ctx, - `SELECT id, from_agent, to_agent, body, created_at - FROM messages - WHERE status = 'pending' - AND to_agent IS NOT NULL - AND to_agent != '' - AND from_agent != 'system' - AND to_agent != 'system' - AND created_at < ?`, - cutoff, - ) - if err != nil { - w.logger.Error("query escalation candidates failed", "error", err) - return 0 - } - defer rows.Close() - - type pendingMsg struct { - ID int64 - FromAgent string - ToAgent string - Body string - CreatedAt time.Time - } - - var pending []pendingMsg - for rows.Next() { - var pm pendingMsg - if err := rows.Scan(&pm.ID, &pm.FromAgent, &pm.ToAgent, &pm.Body, &pm.CreatedAt); err != nil { - w.logger.Error("scan escalation candidate failed", "error", err) - continue - } - pending = append(pending, pm) - } - - if len(pending) == 0 { - return 0 - } - - // Look up #approvals channel - channelID, err := w.channelLookup.GetChannelIDByName(ctx, "approvals") - if err != nil { - w.logger.Warn("cannot escalate: #approvals channel not found", "error", err) - return 0 - } - - count := int64(0) - for _, pm := range pending { - // Check if already escalated - if w.escalationExists(ctx, pm.ID) { - continue - } - - age := formatAge(time.Since(pm.CreatedAt)) - truncBody := truncate(pm.Body, 100) - - body := fmt.Sprintf( - "**ESCALATION**: Pending message for @%s from %s has been unprocessed for %s. Message: \"%s\". Manual intervention may be required.", - pm.ToAgent, pm.FromAgent, age, truncBody, - ) - - _, err := w.msgService.SendMessage(ctx, "system", "", body, SendOptions{ - Subject: fmt.Sprintf("stalemate-escalation:%d", pm.ID), - Priority: 9, - Metadata: fmt.Sprintf(`{"stalemate_escalation_for":%d}`, pm.ID), - ChannelID: &channelID, - }) - if err != nil { - w.logger.Error("send escalation to #approvals failed", - "message_id", pm.ID, - "error", err, - ) - continue - } - w.logger.Info("escalated stale message to #approvals", - "message_id", pm.ID, - "to_agent", pm.ToAgent, - "from_agent", pm.FromAgent, - "age", age, - ) - count++ - } - return count -} - -// reminderExists checks if a system reminder already exists for a given message ID. -func (w *StalemateWorker) reminderExists(ctx context.Context, messageID int64, toAgent string) bool { - var count int - err := w.db.QueryRowContext(ctx, - `SELECT COUNT(*) FROM messages - WHERE from_agent = 'system' - AND to_agent = ? - AND metadata LIKE ?`, - toAgent, fmt.Sprintf(`%%"stalemate_reminder_for":%d%%`, messageID), - ).Scan(&count) - if err != nil { - return false - } - return count > 0 -} - -// escalationExists checks if an escalation already exists for a given message ID. -func (w *StalemateWorker) escalationExists(ctx context.Context, messageID int64) bool { - var count int - err := w.db.QueryRowContext(ctx, - `SELECT COUNT(*) FROM messages - WHERE from_agent = 'system' - AND metadata LIKE ?`, - fmt.Sprintf(`%%"stalemate_escalation_for":%d%%`, messageID), - ).Scan(&count) - if err != nil { - return false - } - return count > 0 -} - -// workflowChannel holds channel info relevant to workflow stalemate checking. -type workflowChannel struct { - ID int64 - Name string - StalemateRemindAfter string - StalemateEscalateAfter string -} - -// staleWorkflowMsg holds info about a channel message in a stale workflow state. -type staleWorkflowMsg struct { - ID int64 - Body string - FromAgent string - ChannelID int64 - Channel string - State string - StateAge time.Duration -} - -// checkWorkflowStalemates scans workflow-enabled channels for messages stuck in -// non-terminal workflow states (proposed, approved, in_progress) and sends -// reminders to channel members or escalates to #approvals. -func (w *StalemateWorker) checkWorkflowStalemates(ctx context.Context) (reminded int64, escalated int64) { - // Step 1: Find all workflow-enabled channels - channels, err := w.listWorkflowChannels(ctx) - if err != nil { - w.logger.Error("list workflow channels failed", "error", err) - return 0, 0 - } - if len(channels) == 0 { - return 0, 0 - } - - for _, ch := range channels { - remindTimeout, err := parseDurationWithDays(ch.StalemateRemindAfter) - if err != nil || remindTimeout <= 0 { - remindTimeout = 24 * time.Hour // default - } - escalateTimeout, err := parseDurationWithDays(ch.StalemateEscalateAfter) - if err != nil || escalateTimeout <= 0 { - escalateTimeout = 72 * time.Hour // default - } - - // Step 2: Find messages in non-terminal workflow states - staleMessages, err := w.findStaleWorkflowMessages(ctx, ch) - if err != nil { - w.logger.Error("find stale workflow messages failed", - "channel", ch.Name, - "error", err, - ) - continue - } - - for _, msg := range staleMessages { - // Step 3: Check escalation first (longer timeout) - if msg.StateAge >= escalateTimeout { - if w.workflowEscalationExists(ctx, msg.ID) { - continue - } - if w.sendWorkflowEscalation(ctx, msg) { - escalated++ - } - continue - } - - // Step 4: Check reminder (shorter timeout) - if msg.StateAge >= remindTimeout { - if w.workflowReminderExists(ctx, msg.ID) { - continue - } - r := w.sendWorkflowReminders(ctx, msg, ch.ID) - reminded += r - } - } - } - - return reminded, escalated -} - -// listWorkflowChannels returns all channels that have workflow_enabled = true. -func (w *StalemateWorker) listWorkflowChannels(ctx context.Context) ([]workflowChannel, error) { - rows, err := w.db.QueryContext(ctx, - `SELECT id, name, stalemate_remind_after, stalemate_escalate_after - FROM channels - WHERE workflow_enabled = 1`) - if err != nil { - return nil, fmt.Errorf("query workflow channels: %w", err) - } - defer rows.Close() - - var channels []workflowChannel - for rows.Next() { - var ch workflowChannel - if err := rows.Scan(&ch.ID, &ch.Name, &ch.StalemateRemindAfter, &ch.StalemateEscalateAfter); err != nil { - return nil, fmt.Errorf("scan workflow channel: %w", err) - } - channels = append(channels, ch) - } - return channels, rows.Err() -} - -// findStaleWorkflowMessages finds channel messages in non-terminal workflow states -// and computes how long they have been in their current state. -func (w *StalemateWorker) findStaleWorkflowMessages(ctx context.Context, ch workflowChannel) ([]staleWorkflowMsg, error) { - // Get all messages in this channel that could be in a workflow state. - // We fetch messages and their reactions, then compute state in Go. - rows, err := w.db.QueryContext(ctx, - `SELECT m.id, m.body, m.from_agent, m.created_at - FROM messages m - WHERE m.channel_id = ? - AND m.from_agent != 'system' - ORDER BY m.created_at ASC`, - ch.ID, - ) - if err != nil { - return nil, fmt.Errorf("query channel messages: %w", err) - } - defer rows.Close() - - type chanMsg struct { - ID int64 - Body string - FromAgent string - CreatedAt time.Time - } - - var msgs []chanMsg - for rows.Next() { - var m chanMsg - if err := rows.Scan(&m.ID, &m.Body, &m.FromAgent, &m.CreatedAt); err != nil { - return nil, fmt.Errorf("scan channel message: %w", err) - } - msgs = append(msgs, m) - } - if err := rows.Err(); err != nil { - return nil, err - } - - if len(msgs) == 0 { - return nil, nil - } - - // Batch-fetch reactions for all messages - msgIDs := make([]int64, len(msgs)) - for i, m := range msgs { - msgIDs[i] = m.ID - } - - reactionsMap, err := w.getReactionsByMessageIDs(ctx, msgIDs) - if err != nil { - return nil, fmt.Errorf("get reactions: %w", err) - } - - now := time.Now() - var stale []staleWorkflowMsg - for _, m := range msgs { - reactions := reactionsMap[m.ID] - state := computeWorkflowStateFromReactions(reactions) - - // Skip terminal states - if isTerminalWorkflowState(state) { - continue - } - - // Determine the "state age": how long since the state was entered. - // If reactions exist, use the most recent reaction's created_at. - // If no reactions (proposed state), use the message's created_at. - stateEnteredAt := m.CreatedAt - if len(reactions) > 0 { - // Find the most recent reaction - for _, r := range reactions { - if r.CreatedAt.After(stateEnteredAt) { - stateEnteredAt = r.CreatedAt - } - } - } - - stale = append(stale, staleWorkflowMsg{ - ID: m.ID, - Body: m.Body, - FromAgent: m.FromAgent, - ChannelID: ch.ID, - Channel: ch.Name, - State: state, - StateAge: now.Sub(stateEnteredAt), - }) - } - - return stale, nil -} - -// reactionRow holds a raw reaction row for workflow state computation. -type reactionRow struct { - Reaction string - CreatedAt time.Time -} - -// getReactionsByMessageIDs fetches reactions for a batch of message IDs. -func (w *StalemateWorker) getReactionsByMessageIDs(ctx context.Context, messageIDs []int64) (map[int64][]reactionRow, error) { - if len(messageIDs) == 0 { - return map[int64][]reactionRow{}, nil - } - - placeholders := make([]string, len(messageIDs)) - args := make([]any, len(messageIDs)) - for i, id := range messageIDs { - placeholders[i] = "?" - args[i] = id - } - - query := fmt.Sprintf( - `SELECT message_id, reaction, created_at - FROM message_reactions - WHERE message_id IN (%s) - ORDER BY created_at ASC`, - strings.Join(placeholders, ","), - ) - - rows, err := w.db.QueryContext(ctx, query, args...) - if err != nil { - return nil, fmt.Errorf("query reactions: %w", err) - } - defer rows.Close() - - result := make(map[int64][]reactionRow) - for rows.Next() { - var msgID int64 - var r reactionRow - if err := rows.Scan(&msgID, &r.Reaction, &r.CreatedAt); err != nil { - return nil, fmt.Errorf("scan reaction: %w", err) - } - result[msgID] = append(result[msgID], r) - } - return result, rows.Err() -} - -// computeWorkflowStateFromReactions derives workflow state from raw reaction rows. -// Mirrors the logic in reactions.ComputeWorkflowState without importing that package. -func computeWorkflowStateFromReactions(reactions []reactionRow) string { - if len(reactions) == 0 { - return "proposed" - } - - // Reaction priority (same as reactions.reactionPriority) - priority := map[string]int{ - "approve": 2, - "in_progress": 3, - "reject": 4, - "done": 5, - "published": 6, - } - - // Reaction-to-state mapping (same as reactions.reactionToState) - toState := map[string]string{ - "approve": "approved", - "reject": "rejected", - "in_progress": "in_progress", - "done": "done", - "published": "published", - } - - highestPriority := 0 - highestState := "proposed" - - for _, r := range reactions { - if p, ok := priority[r.Reaction]; ok && p > highestPriority { - highestPriority = p - highestState = toState[r.Reaction] - } - } - - return highestState -} - -// isTerminalWorkflowState returns true if the state should not trigger stalemate checks. -func isTerminalWorkflowState(state string) bool { - switch state { - case "rejected", "done", "published": - return true - default: - return false - } -} - -// sendWorkflowReminders sends DMs to channel members about a stale workflow message. -func (w *StalemateWorker) sendWorkflowReminders(ctx context.Context, msg staleWorkflowMsg, channelID int64) int64 { - // Get channel members - rows, err := w.db.QueryContext(ctx, - `SELECT agent_name FROM channel_members WHERE channel_id = ?`, - channelID, - ) - if err != nil { - w.logger.Error("query channel members for workflow reminder failed", - "channel_id", channelID, - "error", err, - ) - return 0 - } - defer rows.Close() - - var members []string - for rows.Next() { - var name string - if err := rows.Scan(&name); err != nil { - continue - } - members = append(members, name) - } - - age := formatAge(msg.StateAge) - truncBody := truncate(msg.Body, 100) - count := int64(0) - - for _, member := range members { - body := fmt.Sprintf( - "**STALE**: Message #%d in #%s in '%s' for %s. \"%s\" — @%s", - msg.ID, msg.Channel, msg.State, age, truncBody, msg.FromAgent, - ) - - _, err := w.msgService.SendMessage(ctx, "system", member, body, SendOptions{ - Subject: fmt.Sprintf("workflow-stalemate-reminder:%d", msg.ID), - Priority: 7, - Metadata: fmt.Sprintf(`{"workflow_stalemate_reminder_for":%d}`, msg.ID), - }) - if err != nil { - w.logger.Error("send workflow stalemate reminder failed", - "message_id", msg.ID, - "to_agent", member, - "error", err, - ) - continue - } - w.logger.Info("sent workflow stalemate reminder", - "message_id", msg.ID, - "channel", msg.Channel, - "state", msg.State, - "to_agent", member, - "age", age, - ) - count++ - } - return count -} - -// sendWorkflowEscalation posts an escalation to #approvals for a stale workflow message. -func (w *StalemateWorker) sendWorkflowEscalation(ctx context.Context, msg staleWorkflowMsg) bool { - approvalsChanID, err := w.channelLookup.GetChannelIDByName(ctx, "approvals") - if err != nil { - w.logger.Warn("cannot escalate workflow stalemate: #approvals channel not found", "error", err) - return false - } - - age := formatAge(msg.StateAge) - truncBody := truncate(msg.Body, 100) - - body := fmt.Sprintf( - "**STALE**: Message #%d in #%s in '%s' for %s. \"%s\" — @%s", - msg.ID, msg.Channel, msg.State, age, truncBody, msg.FromAgent, - ) - - _, err = w.msgService.SendMessage(ctx, "system", "", body, SendOptions{ - Subject: fmt.Sprintf("workflow-stalemate-escalation:%d", msg.ID), - Priority: 9, - Metadata: fmt.Sprintf(`{"workflow_stalemate_escalation_for":%d}`, msg.ID), - ChannelID: &approvalsChanID, - }) - if err != nil { - w.logger.Error("send workflow escalation to #approvals failed", - "message_id", msg.ID, - "channel", msg.Channel, - "error", err, - ) - return false - } - w.logger.Info("escalated stale workflow message to #approvals", - "message_id", msg.ID, - "channel", msg.Channel, - "state", msg.State, - "age", age, - ) - return true -} - -// workflowReminderExists checks if a workflow stalemate reminder already exists for a message. -func (w *StalemateWorker) workflowReminderExists(ctx context.Context, messageID int64) bool { - var count int - err := w.db.QueryRowContext(ctx, - `SELECT COUNT(*) FROM messages - WHERE from_agent = 'system' - AND metadata LIKE ?`, - fmt.Sprintf(`%%"workflow_stalemate_reminder_for":%d%%`, messageID), - ).Scan(&count) - if err != nil { - return false - } - return count > 0 -} - -// workflowEscalationExists checks if a workflow stalemate escalation already exists for a message. -func (w *StalemateWorker) workflowEscalationExists(ctx context.Context, messageID int64) bool { - var count int - err := w.db.QueryRowContext(ctx, - `SELECT COUNT(*) FROM messages - WHERE from_agent = 'system' - AND metadata LIKE ?`, - fmt.Sprintf(`%%"workflow_stalemate_escalation_for":%d%%`, messageID), - ).Scan(&count) - if err != nil { - return false - } - return count > 0 -} - -// truncate truncates a string to maxLen characters, appending "..." if truncated. -func truncate(s string, maxLen int) string { - runes := []rune(s) - if len(runes) <= maxLen { - return s - } - return string(runes[:maxLen]) + "..." -} - -// formatAge returns a human-readable age string. -func formatAge(d time.Duration) string { - if d < time.Hour { - return fmt.Sprintf("%dm", int(d.Minutes())) - } - hours := int(d.Hours()) - if hours < 24 { - return fmt.Sprintf("%dh", hours) - } - days := hours / 24 - remainingHours := hours % 24 - if remainingHours == 0 { - if days == 1 { - return "1 day" - } - return fmt.Sprintf("%d days", days) - } - if days == 1 { - return fmt.Sprintf("1 day %dh", remainingHours) - } - return fmt.Sprintf("%d days %dh", days, remainingHours) -} diff --git a/internal/messaging/stalemate_test.go b/internal/messaging/stalemate_test.go index c248757..88fe1a5 100644 --- a/internal/messaging/stalemate_test.go +++ b/internal/messaging/stalemate_test.go @@ -3,7 +3,6 @@ package messaging import ( "context" "database/sql" - "fmt" "os" "testing" "time" @@ -13,19 +12,6 @@ import ( "github.com/synapbus/synapbus/internal/trace" ) -// stubChannelLookup implements ChannelLookup for tests. -type stubChannelLookup struct { - channelID int64 - err error -} - -func (s *stubChannelLookup) GetChannelIDByName(ctx context.Context, name string) (int64, error) { - if s.err != nil { - return 0, s.err - } - return s.channelID, nil -} - // newStalemateTestService creates a MessagingService and DB for stalemate tests. func newStalemateTestService(t *testing.T) (*MessagingService, *sql.DB) { t.Helper() @@ -47,7 +33,6 @@ func newStalemateTestService(t *testing.T) (*MessagingService, *sql.DB) { func insertStaleMessage(t *testing.T, db *sql.DB, from, to, body, status string, createdAt time.Time, claimedAt *time.Time, claimedBy string) int64 { t.Helper() - // Insert conversation first result, err := db.Exec( `INSERT INTO conversations (subject, created_by, created_at, updated_at) VALUES (?, ?, ?, ?)`, @@ -83,19 +68,15 @@ func TestStalemateWorker_ProcessingTimeout(t *testing.T) { svc, db := newStalemateTestService(t) ctx := context.Background() - // Insert a message in "processing" status with old claimed_at oldClaimedAt := time.Now().Add(-25 * time.Hour) msgID := insertStaleMessage(t, db, "sender", "receiver", "stale processing task", StatusProcessing, time.Now().Add(-26*time.Hour), &oldClaimedAt, "receiver") config := DefaultStalemateConfig() config.ProcessingTimeout = 24 * time.Hour - - lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")} - worker := NewStalemateWorker(db, svc, lookup, config) + worker := NewStalemateWorker(db, svc, config) worker.checkStaleMessages(ctx) - // Verify message was auto-failed var status, metadata string err := db.QueryRowContext(ctx, `SELECT status, metadata FROM messages WHERE id = ?`, msgID).Scan(&status, &metadata) if err != nil { @@ -113,19 +94,15 @@ func TestStalemateWorker_ProcessingTimeout_NotExpired(t *testing.T) { svc, db := newStalemateTestService(t) ctx := context.Background() - // Insert a message in "processing" status with recent claimed_at (should NOT be failed) recentClaimedAt := time.Now().Add(-1 * time.Hour) msgID := insertStaleMessage(t, db, "sender", "receiver", "recent processing task", StatusProcessing, time.Now().Add(-2*time.Hour), &recentClaimedAt, "receiver") config := DefaultStalemateConfig() config.ProcessingTimeout = 24 * time.Hour - - lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")} - worker := NewStalemateWorker(db, svc, lookup, config) + worker := NewStalemateWorker(db, svc, config) worker.checkStaleMessages(ctx) - // Verify message was NOT auto-failed var status string err := db.QueryRowContext(ctx, `SELECT status FROM messages WHERE id = ?`, msgID).Scan(&status) if err != nil { @@ -136,29 +113,22 @@ func TestStalemateWorker_ProcessingTimeout_NotExpired(t *testing.T) { } } -// TestStalemateWorker_ProcessingTimeout_RaceGuard verifies that the -// auto-fail UPDATE re-checks claimed_at < cutoff and won't stomp a row that -// was legitimately re-claimed (claimed_at refreshed) between the worker's -// SELECT scan and its row-by-row UPDATE. This guards the TOCTOU window -// the stale-worker race depends on. +// TestStalemateWorker_ProcessingTimeout_RaceGuard verifies that the auto-fail +// UPDATE re-checks claimed_at < cutoff so a row re-claimed between SELECT and +// UPDATE is not stomped. Guards the TOCTOU window the stale-worker race +// depends on. func TestStalemateWorker_ProcessingTimeout_RaceGuard(t *testing.T) { svc, db := newStalemateTestService(t) ctx := context.Background() - // Insert a message that *was* stale at SELECT time. oldClaimedAt := time.Now().Add(-25 * time.Hour) msgID := insertStaleMessage(t, db, "sender", "receiver", "racing task", StatusProcessing, time.Now().Add(-26*time.Hour), &oldClaimedAt, "receiver") config := DefaultStalemateConfig() config.ProcessingTimeout = 24 * time.Hour - lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")} - worker := NewStalemateWorker(db, svc, lookup, config) + worker := NewStalemateWorker(db, svc, config) - // Simulate the race: between the worker's SELECT (which would have picked - // this row) and its UPDATE, the legitimate claimer refreshes claimed_at to - // "now". With the cutoff predicate in place the UPDATE no-ops instead of - // silently failing live work. freshClaimedAt := time.Now() if _, err := db.ExecContext(ctx, `UPDATE messages SET claimed_at = ? WHERE id = ?`, @@ -180,177 +150,6 @@ func TestStalemateWorker_ProcessingTimeout_RaceGuard(t *testing.T) { } } -func TestStalemateWorker_PendingReminder(t *testing.T) { - svc, db := newStalemateTestService(t) - ctx := context.Background() - - // Insert a pending DM that is 5 hours old - insertStaleMessage(t, db, "sender", "receiver", "please review this", StatusPending, time.Now().Add(-5*time.Hour), nil, "") - - config := DefaultStalemateConfig() - config.ReminderAfter = 4 * time.Hour - config.EscalateAfter = 48 * time.Hour // won't trigger - - lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")} - worker := NewStalemateWorker(db, svc, lookup, config) - - worker.checkStaleMessages(ctx) - - // Verify a system reminder was sent to receiver - var count int - err := db.QueryRowContext(ctx, - `SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND to_agent = 'receiver' AND body LIKE '%Reminder%'`, - ).Scan(&count) - if err != nil { - t.Fatalf("query reminder: %v", err) - } - if count != 1 { - t.Errorf("expected 1 reminder, got %d", count) - } -} - -func TestStalemateWorker_SystemMessageSkip(t *testing.T) { - svc, db := newStalemateTestService(t) - ctx := context.Background() - - // Insert a pending DM FROM system (should be skipped) - insertStaleMessage(t, db, "system", "receiver", "system notification", StatusPending, time.Now().Add(-5*time.Hour), nil, "") - - config := DefaultStalemateConfig() - config.ReminderAfter = 4 * time.Hour - - lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")} - worker := NewStalemateWorker(db, svc, lookup, config) - - worker.checkStaleMessages(ctx) - - // Verify NO reminder was sent (only the original system message should exist) - var count int - err := db.QueryRowContext(ctx, - `SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND body LIKE '%Reminder%'`, - ).Scan(&count) - if err != nil { - t.Fatalf("query reminder: %v", err) - } - if count != 0 { - t.Errorf("expected 0 reminders for system message, got %d", count) - } -} - -func TestStalemateWorker_DuplicateReminderPrevention(t *testing.T) { - svc, db := newStalemateTestService(t) - ctx := context.Background() - - // Insert a pending DM that is old enough for a reminder - insertStaleMessage(t, db, "sender", "receiver", "need your attention", StatusPending, time.Now().Add(-5*time.Hour), nil, "") - - config := DefaultStalemateConfig() - config.ReminderAfter = 4 * time.Hour - config.EscalateAfter = 48 * time.Hour - - lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no channel")} - worker := NewStalemateWorker(db, svc, lookup, config) - - // Run check twice - worker.checkStaleMessages(ctx) - worker.checkStaleMessages(ctx) - - // Verify only ONE reminder was sent - var count int - err := db.QueryRowContext(ctx, - `SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND to_agent = 'receiver' AND body LIKE '%Reminder%'`, - ).Scan(&count) - if err != nil { - t.Fatalf("query reminders: %v", err) - } - if count != 1 { - t.Errorf("expected 1 reminder (no duplicates), got %d", count) - } -} - -func TestStalemateWorker_Escalation(t *testing.T) { - svc, db := newStalemateTestService(t) - ctx := context.Background() - - // Create #approvals channel - _, err := db.Exec( - `INSERT INTO channels (id, name, description, topic, type, is_private, is_system, created_by, created_at, updated_at) - VALUES (1, 'approvals', 'Approval queue', '', 'standard', 0, 0, 'system', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`) - if err != nil { - t.Fatalf("create approvals channel: %v", err) - } - // Add system as member - _, err = db.Exec( - `INSERT INTO channel_members (channel_id, agent_name, role, joined_at) - VALUES (1, 'system', 'owner', CURRENT_TIMESTAMP)`) - if err != nil { - t.Fatalf("add system to channel: %v", err) - } - - // Insert a pending DM that is 49 hours old (beyond escalation threshold) - insertStaleMessage(t, db, "sender", "receiver", "urgent task ignored", StatusPending, time.Now().Add(-49*time.Hour), nil, "") - - config := DefaultStalemateConfig() - config.ReminderAfter = 4 * time.Hour - config.EscalateAfter = 48 * time.Hour - - lookup := &stubChannelLookup{channelID: 1} - worker := NewStalemateWorker(db, svc, lookup, config) - - worker.checkStaleMessages(ctx) - - // Verify an escalation was sent to #approvals channel - var count int - err = db.QueryRowContext(ctx, - `SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND channel_id = 1 AND body LIKE '%ESCALATION%'`, - ).Scan(&count) - if err != nil { - t.Fatalf("query escalations: %v", err) - } - if count != 1 { - t.Errorf("expected 1 escalation, got %d", count) - } -} - -func TestStalemateWorker_DuplicateEscalationPrevention(t *testing.T) { - svc, db := newStalemateTestService(t) - ctx := context.Background() - - // Create #approvals channel - db.Exec( - `INSERT INTO channels (id, name, description, topic, type, is_private, is_system, created_by, created_at, updated_at) - VALUES (1, 'approvals', 'Approval queue', '', 'standard', 0, 0, 'system', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`) - db.Exec( - `INSERT INTO channel_members (channel_id, agent_name, role, joined_at) - VALUES (1, 'system', 'owner', CURRENT_TIMESTAMP)`) - - // Insert a pending DM that is 49 hours old - insertStaleMessage(t, db, "sender", "receiver", "urgent task", StatusPending, time.Now().Add(-49*time.Hour), nil, "") - - config := DefaultStalemateConfig() - config.ReminderAfter = 4 * time.Hour - config.EscalateAfter = 48 * time.Hour - - lookup := &stubChannelLookup{channelID: 1} - worker := NewStalemateWorker(db, svc, lookup, config) - - // Run check twice - worker.checkStaleMessages(ctx) - worker.checkStaleMessages(ctx) - - // Verify only ONE escalation was sent - var count int - err := db.QueryRowContext(ctx, - `SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND channel_id = 1 AND body LIKE '%ESCALATION%'`, - ).Scan(&count) - if err != nil { - t.Fatalf("query escalations: %v", err) - } - if count != 1 { - t.Errorf("expected 1 escalation (no duplicates), got %d", count) - } -} - func TestParseStalemateConfig(t *testing.T) { tests := []struct { name string @@ -366,14 +165,10 @@ func TestParseStalemateConfig(t *testing.T) { name: "custom values with day format", envVars: map[string]string{ "SYNAPBUS_STALEMATE_PROCESSING_TIMEOUT": "7d", - "SYNAPBUS_STALEMATE_REMINDER_AFTER": "8h", - "SYNAPBUS_STALEMATE_ESCALATE_AFTER": "3d", "SYNAPBUS_STALEMATE_INTERVAL": "30m", }, expected: StalemateConfig{ ProcessingTimeout: 7 * 24 * time.Hour, - ReminderAfter: 8 * time.Hour, - EscalateAfter: 3 * 24 * time.Hour, Interval: 30 * time.Minute, }, }, @@ -381,14 +176,10 @@ func TestParseStalemateConfig(t *testing.T) { name: "standard Go duration format", envVars: map[string]string{ "SYNAPBUS_STALEMATE_PROCESSING_TIMEOUT": "48h", - "SYNAPBUS_STALEMATE_REMINDER_AFTER": "2h30m", - "SYNAPBUS_STALEMATE_ESCALATE_AFTER": "72h", "SYNAPBUS_STALEMATE_INTERVAL": "5m", }, expected: StalemateConfig{ ProcessingTimeout: 48 * time.Hour, - ReminderAfter: 2*time.Hour + 30*time.Minute, - EscalateAfter: 72 * time.Hour, Interval: 5 * time.Minute, }, }, @@ -396,28 +187,22 @@ func TestParseStalemateConfig(t *testing.T) { name: "invalid values fall back to defaults", envVars: map[string]string{ "SYNAPBUS_STALEMATE_PROCESSING_TIMEOUT": "invalid", - "SYNAPBUS_STALEMATE_REMINDER_AFTER": "bad", - "SYNAPBUS_STALEMATE_ESCALATE_AFTER": "", "SYNAPBUS_STALEMATE_INTERVAL": "-5m", }, expected: DefaultStalemateConfig(), }, } + envKeys := []string{ + "SYNAPBUS_STALEMATE_PROCESSING_TIMEOUT", + "SYNAPBUS_STALEMATE_INTERVAL", + } + for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { - // Clear all env vars first - envKeys := []string{ - "SYNAPBUS_STALEMATE_PROCESSING_TIMEOUT", - "SYNAPBUS_STALEMATE_REMINDER_AFTER", - "SYNAPBUS_STALEMATE_ESCALATE_AFTER", - "SYNAPBUS_STALEMATE_INTERVAL", - } for _, k := range envKeys { os.Unsetenv(k) } - - // Set test env vars for k, v := range tt.envVars { os.Setenv(k, v) } @@ -432,12 +217,6 @@ func TestParseStalemateConfig(t *testing.T) { if cfg.ProcessingTimeout != tt.expected.ProcessingTimeout { t.Errorf("ProcessingTimeout = %v, want %v", cfg.ProcessingTimeout, tt.expected.ProcessingTimeout) } - if cfg.ReminderAfter != tt.expected.ReminderAfter { - t.Errorf("ReminderAfter = %v, want %v", cfg.ReminderAfter, tt.expected.ReminderAfter) - } - if cfg.EscalateAfter != tt.expected.EscalateAfter { - t.Errorf("EscalateAfter = %v, want %v", cfg.EscalateAfter, tt.expected.EscalateAfter) - } if cfg.Interval != tt.expected.Interval { t.Errorf("Interval = %v, want %v", cfg.Interval, tt.expected.Interval) } @@ -447,10 +226,10 @@ func TestParseStalemateConfig(t *testing.T) { func TestParseDurationWithDays(t *testing.T) { tests := []struct { - name string - input string - want time.Duration - wantErr bool + name string + input string + want time.Duration + wantErr bool }{ {"7 days", "7d", 7 * 24 * time.Hour, false}, {"1 day", "1d", 24 * time.Hour, false}, @@ -475,392 +254,3 @@ func TestParseDurationWithDays(t *testing.T) { }) } } - -func TestTruncate(t *testing.T) { - tests := []struct { - name string - input string - maxLen int - want string - }{ - {"short string", "hello", 10, "hello"}, - {"exact length", "hello", 5, "hello"}, - {"truncated", "hello world, this is a long message", 10, "hello worl..."}, - {"empty", "", 10, ""}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := truncate(tt.input, tt.maxLen) - if got != tt.want { - t.Errorf("truncate(%q, %d) = %q, want %q", tt.input, tt.maxLen, got, tt.want) - } - }) - } -} - -func TestFormatAge(t *testing.T) { - tests := []struct { - name string - d time.Duration - want string - }{ - {"minutes", 30 * time.Minute, "30m"}, - {"hours", 5 * time.Hour, "5h"}, - {"1 day", 24 * time.Hour, "1 day"}, - {"2 days", 48 * time.Hour, "2 days"}, - {"1 day with hours", 25 * time.Hour, "1 day 1h"}, - {"2 days with hours", 50 * time.Hour, "2 days 2h"}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := formatAge(tt.d) - if got != tt.want { - t.Errorf("formatAge(%v) = %q, want %q", tt.d, got, tt.want) - } - }) - } -} - -func TestComputeWorkflowStateFromReactions(t *testing.T) { - tests := []struct { - name string - reactions []reactionRow - want string - }{ - {"no reactions = proposed", nil, "proposed"}, - {"approve only", []reactionRow{{Reaction: "approve"}}, "approved"}, - {"in_progress only", []reactionRow{{Reaction: "in_progress"}}, "in_progress"}, - {"reject only", []reactionRow{{Reaction: "reject"}}, "rejected"}, - {"done only", []reactionRow{{Reaction: "done"}}, "done"}, - {"published only", []reactionRow{{Reaction: "published"}}, "published"}, - {"approve + in_progress = in_progress (higher priority)", []reactionRow{ - {Reaction: "approve"}, - {Reaction: "in_progress"}, - }, "in_progress"}, - {"approve + done = done", []reactionRow{ - {Reaction: "approve"}, - {Reaction: "done"}, - }, "done"}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := computeWorkflowStateFromReactions(tt.reactions) - if got != tt.want { - t.Errorf("computeWorkflowStateFromReactions() = %q, want %q", got, tt.want) - } - }) - } -} - -func TestIsTerminalWorkflowState(t *testing.T) { - tests := []struct { - state string - terminal bool - }{ - {"proposed", false}, - {"approved", false}, - {"in_progress", false}, - {"rejected", true}, - {"done", true}, - {"published", true}, - } - - for _, tt := range tests { - t.Run(tt.state, func(t *testing.T) { - got := isTerminalWorkflowState(tt.state) - if got != tt.terminal { - t.Errorf("isTerminalWorkflowState(%q) = %v, want %v", tt.state, got, tt.terminal) - } - }) - } -} - -func TestStalemateWorker_WorkflowReminder(t *testing.T) { - svc, db := newStalemateTestService(t) - ctx := context.Background() - - // Create a workflow-enabled channel with short timeouts - _, err := db.Exec( - `INSERT INTO channels (id, name, description, topic, type, is_private, is_system, created_by, workflow_enabled, stalemate_remind_after, stalemate_escalate_after, created_at, updated_at) - VALUES (10, 'news-test', 'Test news channel', '', 'standard', 0, 0, 'system', 1, '1s', '72h', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`) - if err != nil { - t.Fatalf("create workflow channel: %v", err) - } - - // Add system and sender as members - db.Exec(`INSERT INTO channel_members (channel_id, agent_name, role, joined_at) VALUES (10, 'system', 'owner', CURRENT_TIMESTAMP)`) - db.Exec(`INSERT INTO channel_members (channel_id, agent_name, role, joined_at) VALUES (10, 'sender', 'member', CURRENT_TIMESTAMP)`) - db.Exec(`INSERT INTO channel_members (channel_id, agent_name, role, joined_at) VALUES (10, 'receiver', 'member', CURRENT_TIMESTAMP)`) - - // Insert a channel message with old created_at (will be in "proposed" state since no reactions) - oldTime := time.Now().Add(-2 * time.Second) - convResult, err := db.Exec( - `INSERT INTO conversations (subject, created_by, created_at, updated_at) VALUES ('wf-test', 'sender', ?, ?)`, - oldTime, oldTime, - ) - if err != nil { - t.Fatalf("insert conversation: %v", err) - } - convID, _ := convResult.LastInsertId() - - channelID := int64(10) - _, err = db.Exec( - `INSERT INTO messages (conversation_id, from_agent, to_agent, body, priority, status, metadata, channel_id, created_at, updated_at) - VALUES (?, 'sender', '', 'Draft blog post about MCP', 5, 'pending', '{}', ?, ?, ?)`, - convID, channelID, oldTime, oldTime, - ) - if err != nil { - t.Fatalf("insert channel message: %v", err) - } - - // Wait for the timeout to elapse - time.Sleep(10 * time.Millisecond) - - config := DefaultStalemateConfig() - lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no approvals channel")} - worker := NewStalemateWorker(db, svc, lookup, config) - - worker.checkStaleMessages(ctx) - - // Verify workflow stalemate reminders were sent to channel members - var count int - err = db.QueryRowContext(ctx, - `SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND body LIKE '%STALE%'`, - ).Scan(&count) - if err != nil { - t.Fatalf("query workflow reminders: %v", err) - } - // Should have sent reminders to all 3 members (system, sender, receiver) - if count < 1 { - t.Errorf("expected at least 1 workflow reminder, got %d", count) - } -} - -func TestStalemateWorker_WorkflowEscalation(t *testing.T) { - svc, db := newStalemateTestService(t) - ctx := context.Background() - - // Create a workflow-enabled channel with short escalation timeout - _, err := db.Exec( - `INSERT INTO channels (id, name, description, topic, type, is_private, is_system, created_by, workflow_enabled, stalemate_remind_after, stalemate_escalate_after, created_at, updated_at) - VALUES (10, 'news-test', 'Test news channel', '', 'standard', 0, 0, 'system', 1, '1s', '1s', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`) - if err != nil { - t.Fatalf("create workflow channel: %v", err) - } - - // Create #approvals channel - db.Exec( - `INSERT INTO channels (id, name, description, topic, type, is_private, is_system, created_by, created_at, updated_at) - VALUES (20, 'approvals', 'Approval queue', '', 'standard', 0, 0, 'system', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`) - db.Exec(`INSERT INTO channel_members (channel_id, agent_name, role, joined_at) VALUES (20, 'system', 'owner', CURRENT_TIMESTAMP)`) - - // Add members to workflow channel - db.Exec(`INSERT INTO channel_members (channel_id, agent_name, role, joined_at) VALUES (10, 'sender', 'member', CURRENT_TIMESTAMP)`) - - // Insert a channel message old enough to trigger escalation - oldTime := time.Now().Add(-2 * time.Second) - convResult, _ := db.Exec( - `INSERT INTO conversations (subject, created_by, created_at, updated_at) VALUES ('wf-esc', 'sender', ?, ?)`, - oldTime, oldTime, - ) - convID, _ := convResult.LastInsertId() - - channelID := int64(10) - _, err = db.Exec( - `INSERT INTO messages (conversation_id, from_agent, to_agent, body, priority, status, metadata, channel_id, created_at, updated_at) - VALUES (?, 'sender', '', 'Stale proposal needing attention', 5, 'pending', '{}', ?, ?, ?)`, - convID, channelID, oldTime, oldTime, - ) - if err != nil { - t.Fatalf("insert channel message: %v", err) - } - - time.Sleep(10 * time.Millisecond) - - config := DefaultStalemateConfig() - lookup := &stubChannelLookup{channelID: 20} - worker := NewStalemateWorker(db, svc, lookup, config) - - worker.checkStaleMessages(ctx) - - // Verify escalation was sent to #approvals - var count int - err = db.QueryRowContext(ctx, - `SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND channel_id = 20 AND body LIKE '%STALE%'`, - ).Scan(&count) - if err != nil { - t.Fatalf("query workflow escalation: %v", err) - } - if count != 1 { - t.Errorf("expected 1 workflow escalation, got %d", count) - } -} - -func TestStalemateWorker_WorkflowTerminalStateSkip(t *testing.T) { - svc, db := newStalemateTestService(t) - ctx := context.Background() - - // Create a workflow-enabled channel with short timeouts - _, err := db.Exec( - `INSERT INTO channels (id, name, description, topic, type, is_private, is_system, created_by, workflow_enabled, stalemate_remind_after, stalemate_escalate_after, created_at, updated_at) - VALUES (10, 'news-test', 'Test news channel', '', 'standard', 0, 0, 'system', 1, '1s', '1s', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`) - if err != nil { - t.Fatalf("create workflow channel: %v", err) - } - db.Exec(`INSERT INTO channel_members (channel_id, agent_name, role, joined_at) VALUES (10, 'sender', 'member', CURRENT_TIMESTAMP)`) - - // Insert a channel message - oldTime := time.Now().Add(-2 * time.Second) - convResult, _ := db.Exec( - `INSERT INTO conversations (subject, created_by, created_at, updated_at) VALUES ('wf-done', 'sender', ?, ?)`, - oldTime, oldTime, - ) - convID, _ := convResult.LastInsertId() - - channelID := int64(10) - msgResult, err := db.Exec( - `INSERT INTO messages (conversation_id, from_agent, to_agent, body, priority, status, metadata, channel_id, created_at, updated_at) - VALUES (?, 'sender', '', 'Completed task', 5, 'pending', '{}', ?, ?, ?)`, - convID, channelID, oldTime, oldTime, - ) - if err != nil { - t.Fatalf("insert channel message: %v", err) - } - msgID, _ := msgResult.LastInsertId() - - // Add a "done" reaction — puts it in terminal state - _, err = db.Exec( - `INSERT INTO message_reactions (message_id, agent_name, reaction, metadata, created_at) - VALUES (?, 'sender', 'done', '{}', ?)`, - msgID, oldTime, - ) - if err != nil { - t.Fatalf("insert reaction: %v", err) - } - - time.Sleep(10 * time.Millisecond) - - config := DefaultStalemateConfig() - lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no approvals")} - worker := NewStalemateWorker(db, svc, lookup, config) - - worker.checkStaleMessages(ctx) - - // Verify NO reminders were sent (message is in terminal "done" state) - var count int - err = db.QueryRowContext(ctx, - `SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND body LIKE '%STALE%'`, - ).Scan(&count) - if err != nil { - t.Fatalf("query reminders: %v", err) - } - if count != 0 { - t.Errorf("expected 0 reminders for terminal state message, got %d", count) - } -} - -func TestStalemateWorker_WorkflowDuplicateReminderPrevention(t *testing.T) { - svc, db := newStalemateTestService(t) - ctx := context.Background() - - // Create a workflow-enabled channel with short timeout - _, err := db.Exec( - `INSERT INTO channels (id, name, description, topic, type, is_private, is_system, created_by, workflow_enabled, stalemate_remind_after, stalemate_escalate_after, created_at, updated_at) - VALUES (10, 'news-test', 'Test', '', 'standard', 0, 0, 'system', 1, '1s', '72h', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`) - if err != nil { - t.Fatalf("create workflow channel: %v", err) - } - db.Exec(`INSERT INTO channel_members (channel_id, agent_name, role, joined_at) VALUES (10, 'receiver', 'member', CURRENT_TIMESTAMP)`) - - // Insert a channel message - oldTime := time.Now().Add(-2 * time.Second) - convResult, _ := db.Exec( - `INSERT INTO conversations (subject, created_by, created_at, updated_at) VALUES ('wf-dup', 'sender', ?, ?)`, - oldTime, oldTime, - ) - convID, _ := convResult.LastInsertId() - - channelID := int64(10) - _, err = db.Exec( - `INSERT INTO messages (conversation_id, from_agent, to_agent, body, priority, status, metadata, channel_id, created_at, updated_at) - VALUES (?, 'sender', '', 'Needs review', 5, 'pending', '{}', ?, ?, ?)`, - convID, channelID, oldTime, oldTime, - ) - if err != nil { - t.Fatalf("insert channel message: %v", err) - } - - time.Sleep(10 * time.Millisecond) - - config := DefaultStalemateConfig() - lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no approvals")} - worker := NewStalemateWorker(db, svc, lookup, config) - - // Run twice - worker.checkStaleMessages(ctx) - worker.checkStaleMessages(ctx) - - // Verify only one set of reminders was sent (no duplicates) - var count int - err = db.QueryRowContext(ctx, - `SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND to_agent = 'receiver' AND body LIKE '%STALE%'`, - ).Scan(&count) - if err != nil { - t.Fatalf("query reminders: %v", err) - } - if count != 1 { - t.Errorf("expected 1 reminder (no duplicates), got %d", count) - } -} - -func TestStalemateWorker_WorkflowNonWorkflowChannelSkip(t *testing.T) { - svc, db := newStalemateTestService(t) - ctx := context.Background() - - // Create a channel with workflow DISABLED - _, err := db.Exec( - `INSERT INTO channels (id, name, description, topic, type, is_private, is_system, created_by, workflow_enabled, stalemate_remind_after, stalemate_escalate_after, created_at, updated_at) - VALUES (10, 'general', 'General', '', 'standard', 0, 0, 'system', 0, '1s', '1s', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`) - if err != nil { - t.Fatalf("create channel: %v", err) - } - db.Exec(`INSERT INTO channel_members (channel_id, agent_name, role, joined_at) VALUES (10, 'sender', 'member', CURRENT_TIMESTAMP)`) - - // Insert a channel message - oldTime := time.Now().Add(-2 * time.Second) - convResult, _ := db.Exec( - `INSERT INTO conversations (subject, created_by, created_at, updated_at) VALUES ('no-wf', 'sender', ?, ?)`, - oldTime, oldTime, - ) - convID, _ := convResult.LastInsertId() - - channelID := int64(10) - db.Exec( - `INSERT INTO messages (conversation_id, from_agent, to_agent, body, priority, status, metadata, channel_id, created_at, updated_at) - VALUES (?, 'sender', '', 'No workflow here', 5, 'pending', '{}', ?, ?, ?)`, - convID, channelID, oldTime, oldTime, - ) - - time.Sleep(10 * time.Millisecond) - - config := DefaultStalemateConfig() - lookup := &stubChannelLookup{channelID: 0, err: fmt.Errorf("no approvals")} - worker := NewStalemateWorker(db, svc, lookup, config) - - worker.checkStaleMessages(ctx) - - // Verify NO reminders — channel is not workflow-enabled - var count int - err = db.QueryRowContext(ctx, - `SELECT COUNT(*) FROM messages WHERE from_agent = 'system' AND body LIKE '%STALE%'`, - ).Scan(&count) - if err != nil { - t.Fatalf("query reminders: %v", err) - } - if count != 0 { - t.Errorf("expected 0 reminders for non-workflow channel, got %d", count) - } -} diff --git a/internal/search/injection.go b/internal/search/injection.go new file mode 100644 index 0000000..d159a09 --- /dev/null +++ b/internal/search/injection.go @@ -0,0 +1,329 @@ +// Proactive-memory injection retrieval — builds the relevant-context +// packet attached to every injection-eligible MCP tool response (per +// `contracts/mcp-injection.md`). +package search + +import ( + "context" + "database/sql" + "errors" + "fmt" + "strconv" + "time" + + "github.com/synapbus/synapbus/internal/agents" +) + +// MemoryItem is one entry in the `relevant_context.memories[]` array. +// Field shape is contractual (`contracts/mcp-injection.md`). +type MemoryItem struct { + ID int64 `json:"id"` + FromAgent string `json:"from_agent"` + Channel string `json:"channel,omitempty"` + Body string `json:"body"` + CreatedAt time.Time `json:"created_at"` + Score float64 `json:"score"` + MatchType string `json:"match_type"` + Pinned bool `json:"pinned"` + Truncated bool `json:"truncated,omitempty"` +} + +// ContextPacket is the body of the `relevant_context` field. Returned +// from BuildContextPacket; rendered verbatim into the wrapped tool +// response. +type ContextPacket struct { + Memories []MemoryItem `json:"memories"` + CoreMemory string `json:"core_memory,omitempty"` + PacketChars int `json:"packet_chars"` + PacketTokenEstimate int `json:"packet_token_estimate"` + RetrievalQuery string `json:"retrieval_query"` + SearchMode string `json:"search_mode"` +} + +// CoreMemoryProvider is the seam US2 plugs into. When non-nil and +// `opts.IncludeCore` is true, BuildContextPacket calls Get() and +// includes the result verbatim in the packet. +type CoreMemoryProvider interface { + Get(ctx context.Context, ownerID, agentName string) (string, error) +} + +// InjectionOpts captures the per-call configuration for +// BuildContextPacket. Sourced from messaging.MemoryConfig at wrap time. +type InjectionOpts struct { + // BudgetTokens is the soft cap on the assembled packet. 0 disables + // injection: BuildContextPacket returns (nil, nil). + BudgetTokens int + // MaxItems caps the number of memory items in the packet. + MaxItems int + // MinScore is the relevance floor. Items below are dropped. + MinScore float64 + // IncludeCore enables the core-memory lookup. Only session-start + // tools (i.e. my_status) should set this. + IncludeCore bool + // CoreProvider is consulted when IncludeCore is true. May be nil + // (US2 not yet wired) — then no core memory is included. + CoreProvider CoreMemoryProvider + // Now is overridable for tests. Defaults to time.Now. + Now func() time.Time +} + +// EstimateTokens returns the char-based token estimate. Matches R4: +// `(chars + 3) / 4`, ceil-equivalent for positive integers. +func EstimateTokens(chars int) int { + if chars <= 0 { + return 0 + } + return (chars + 3) / 4 +} + +// BuildContextPacket retrieves owner-scoped memories matching `query`, +// applies the score floor + token budget + max items cap, optionally +// resolves the per-(owner, agent) core memory blob, and returns a +// ContextPacket ready to merge into the tool response. +// +// Returns (nil, nil) when `opts.BudgetTokens == 0` (feature disabled) +// or when no memories pass the filter AND no core memory is set — +// callers omit the `relevant_context` field entirely in that case. +func BuildContextPacket( + ctx context.Context, + svc *Service, + agent *agents.Agent, + query string, + opts InjectionOpts, +) (*ContextPacket, error) { + if opts.BudgetTokens == 0 { + return nil, nil + } + if agent == nil { + return nil, fmt.Errorf("build context packet: nil agent") + } + if opts.Now == nil { + opts.Now = time.Now + } + maxItems := opts.MaxItems + if maxItems <= 0 { + maxItems = 5 + } + + callerOwner := strconv.FormatInt(agent.OwnerID, 10) + if agent.OwnerID == 0 { + // Unowned agent: skip retrieval. We still surface a (possibly + // non-empty) core memory if the provider returns one for the + // empty owner — but that's an edge case the provider can decide + // on. + callerOwner = "" + } + + // Over-fetch x3 to absorb owner-scope filtering + score floor drop. + wantedLimit := maxItems * 3 + if wantedLimit < 15 { + wantedLimit = 15 + } + + searchMode := ModeAuto + + var memories []MemoryItem + if svc != nil && callerOwner != "" && query != "" { + // Drive retrieval through the existing hybrid path so we inherit + // access control + ranking. We still need a stricter owner filter + // on top of canAgentAccessMessage because the memory pool + // (open-brain) is broadly readable across agents within the same + // system, and we must enforce owner isolation (SC-008). + resp, err := svc.Search(ctx, agent.Name, SearchOptions{ + Query: query, + Mode: ModeAuto, + Limit: wantedLimit, + MinSimilarity: opts.MinScore, + }) + if err != nil { + return nil, fmt.Errorf("build context packet: search: %w", err) + } + if resp != nil { + searchMode = resp.SearchMode + } + if resp != nil && len(resp.Results) > 0 { + items, err := filterAndScore(ctx, svc, resp.Results, callerOwner, opts) + if err != nil { + return nil, err + } + memories = items + } + } + + // Apply token budget: greedy fill in descending score (results are + // already sorted). Truncate the last admitted item to fit when it + // would otherwise overflow. + memories = applyTokenBudget(memories, maxItems, opts.BudgetTokens) + + // Core memory lookup (US2 hook). Provider may be nil — that's fine. + var coreBlob string + if opts.IncludeCore && opts.CoreProvider != nil && callerOwner != "" { + blob, err := opts.CoreProvider.Get(ctx, callerOwner, agent.Name) + if err == nil { + coreBlob = blob + } else if !errors.Is(err, sql.ErrNoRows) { + // Surface unexpected errors so callers can log; an absent + // core blob (typical "no rows") is not an error. + return nil, fmt.Errorf("build context packet: core memory: %w", err) + } + } + + if len(memories) == 0 && coreBlob == "" { + // Empty + no core → caller should omit the relevant_context field. + return nil, nil + } + + // TODO(US3-T029): pin overlay — when ListPins lands, mark and + // always-include pinned memories regardless of the score floor. + + if memories == nil { + memories = []MemoryItem{} + } + packet := &ContextPacket{ + Memories: memories, + CoreMemory: coreBlob, + RetrievalQuery: query, + SearchMode: searchMode, + } + packet.PacketChars = packetChars(packet) + packet.PacketTokenEstimate = EstimateTokens(packet.PacketChars) + return packet, nil +} + +// filterAndScore drops non-owner messages and items below MinScore, +// then assembles MemoryItem entries in the existing RRF-sorted order. +// +// Owner of each candidate message is resolved via agents.OwnerFor on +// `from_agent`. This is the stricter filter referenced in +// `contracts/mcp-injection.md`'s cross-owner safety note. +func filterAndScore( + ctx context.Context, + svc *Service, + results []*SearchResult, + callerOwnerID string, + opts InjectionOpts, +) ([]MemoryItem, error) { + out := make([]MemoryItem, 0, len(results)) + for _, r := range results { + if r == nil || r.Message == nil { + continue + } + // Score selection: prefer SimilarityScore (semantic / hybrid), + // fall back to RelevanceScore (fulltext). + score := r.SimilarityScore + if score == 0 { + score = r.RelevanceScore + } + if opts.MinScore > 0 && score < opts.MinScore { + continue + } + + owner, err := agents.OwnerFor(ctx, svc.db, r.Message.FromAgent) + if err != nil { + // Unowned or unknown sender → exclude from injection pool. + continue + } + if owner != callerOwnerID { + continue + } + + item := MemoryItem{ + ID: r.Message.ID, + FromAgent: r.Message.FromAgent, + Body: r.Message.Body, + CreatedAt: r.Message.CreatedAt, + Score: score, + MatchType: r.MatchType, + } + if r.Message.ChannelID != nil { + if name, err := channelName(ctx, svc.db, *r.Message.ChannelID); err == nil { + item.Channel = name + } + } + out = append(out, item) + } + return out, nil +} + +// applyTokenBudget enforces both MaxItems and the token budget. Items +// are admitted greedily in input order (caller passes them already +// score-sorted). The first item that would overflow is truncated to +// fit; all later items are skipped. +func applyTokenBudget(items []MemoryItem, maxItems, budgetTokens int) []MemoryItem { + if len(items) == 0 { + return nil + } + if budgetTokens <= 0 { + return nil + } + out := make([]MemoryItem, 0, len(items)) + used := 0 + for i, it := range items { + if i >= maxItems { + break + } + cost := EstimateTokens(itemChars(it)) + if used+cost <= budgetTokens { + used += cost + out = append(out, it) + continue + } + // Truncate this item to whatever remains in the budget. + remaining := budgetTokens - used + if remaining <= 0 { + break + } + // Reserve the per-item overhead in the remaining budget so that + // post-truncate EstimateTokens(itemChars(it)) <= remaining. + overhead := itemChars(it) - len(it.Body) // = from_agent + channel + 32 + // max chars we can place into Body so that total item tokens fit. + maxBodyChars := remaining*4 - overhead + if maxBodyChars <= 0 { + break + } + if maxBodyChars >= len(it.Body) { + // Whole body still fits — admit unchanged. + used += EstimateTokens(itemChars(it)) + out = append(out, it) + continue + } + truncated := it + truncated.Body = it.Body[:maxBodyChars] + truncated.Truncated = true + used += EstimateTokens(itemChars(truncated)) + out = append(out, truncated) + break + } + return out +} + +// itemChars approximates the rendered size of one MemoryItem so the +// budget gate stays self-consistent with packetChars. +func itemChars(it MemoryItem) int { + // Body dominates; from_agent + channel + delimiters add a small + // per-item overhead we approximate at 32 chars. + return len(it.Body) + len(it.FromAgent) + len(it.Channel) + 32 +} + +func packetChars(p *ContextPacket) int { + total := 0 + for _, m := range p.Memories { + total += itemChars(m) + } + total += len(p.CoreMemory) + return total +} + +// channelName resolves a channel ID to a name for the optional +// `MemoryItem.Channel` field. Best-effort: returns ("", err) on lookup +// failure and BuildContextPacket then omits the field entirely. +func channelName(ctx context.Context, db *sql.DB, id int64) (string, error) { + var name string + err := db.QueryRowContext(ctx, + `SELECT name FROM channels WHERE id = ?`, id, + ).Scan(&name) + if err != nil { + return "", err + } + return name, nil +} diff --git a/internal/search/injection_test.go b/internal/search/injection_test.go new file mode 100644 index 0000000..368533c --- /dev/null +++ b/internal/search/injection_test.go @@ -0,0 +1,258 @@ +package search + +import ( + "context" + "database/sql" + "strings" + "testing" + + _ "modernc.org/sqlite" + + "github.com/synapbus/synapbus/internal/agents" + "github.com/synapbus/synapbus/internal/messaging" +) + +// stubCoreProvider lets a test override the core-memory blob returned to +// BuildContextPacket. +type stubCoreProvider struct { + blob string + err error +} + +func (s *stubCoreProvider) Get(ctx context.Context, ownerID, agentName string) (string, error) { + return s.blob, s.err +} + +// seedOwnedAgent inserts a users row + an agents row tied to that owner. +func seedOwnedAgent(t *testing.T, db *sql.DB, ownerID int64, ownerName, agentName string) *agents.Agent { + t.Helper() + if _, err := db.Exec( + `INSERT OR IGNORE INTO users (id, username, password_hash, display_name) + VALUES (?, ?, 'hash', ?)`, ownerID, ownerName, ownerName, + ); err != nil { + t.Fatalf("seed user: %v", err) + } + if _, err := db.Exec( + `INSERT OR IGNORE INTO agents (name, display_name, type, owner_id, api_key_hash, status) + VALUES (?, ?, 'ai', ?, ?, 'active')`, + agentName, agentName, ownerID, agentName+"hash", + ); err != nil { + t.Fatalf("seed agent: %v", err) + } + return &agents.Agent{Name: agentName, OwnerID: ownerID} +} + +// seedChannel inserts an open-brain channel with members `agentNames`. +func seedChannel(t *testing.T, db *sql.DB, channelID int64, name string, agentNames ...string) { + t.Helper() + if _, err := db.Exec( + `INSERT OR IGNORE INTO channels (id, name, description, type, created_by) + VALUES (?, ?, '', 'standard', 'system')`, channelID, name, + ); err != nil { + t.Fatalf("seed channel: %v", err) + } + for _, a := range agentNames { + if _, err := db.Exec( + `INSERT OR IGNORE INTO channel_members (channel_id, agent_name) VALUES (?, ?)`, + channelID, a, + ); err != nil { + t.Fatalf("seed channel_member: %v", err) + } + } +} + +// seedChannelMessage inserts a message directly so we don't need the +// channels.Service stack in this test. +func seedChannelMessage(t *testing.T, db *sql.DB, channelID int64, fromAgent, body string) int64 { + t.Helper() + convRes, err := db.Exec( + `INSERT INTO conversations (created_by, channel_id) VALUES (?, ?)`, + fromAgent, channelID, + ) + if err != nil { + t.Fatalf("seed conversation: %v", err) + } + convID, _ := convRes.LastInsertId() + + res, err := db.Exec( + `INSERT INTO messages (conversation_id, from_agent, channel_id, body, priority, status, metadata) + VALUES (?, ?, ?, ?, 5, 'pending', '{}')`, + convID, fromAgent, channelID, body, + ) + if err != nil { + t.Fatalf("seed message: %v", err) + } + id, _ := res.LastInsertId() + return id +} + +func TestBuildContextPacket_BudgetZeroDisables(t *testing.T) { + svc, _, db := newTestServices(t) + a := seedOwnedAgent(t, db, 1, "alice", "a1") + + pkt, err := BuildContextPacket(context.Background(), svc, a, "query", InjectionOpts{ + BudgetTokens: 0, + }) + if err != nil { + t.Fatalf("BuildContextPacket: %v", err) + } + if pkt != nil { + t.Errorf("BudgetTokens=0 should return nil packet, got %+v", pkt) + } +} + +func TestBuildContextPacket_OwnerScoping(t *testing.T) { + svc, _, db := newTestServices(t) + ctx := context.Background() + + h1Agent := seedOwnedAgent(t, db, 1, "alice", "a1") + _ = seedOwnedAgent(t, db, 2, "bob", "b1") + + seedChannel(t, db, 1, "open-brain", "a1", "b1") + seedChannelMessage(t, db, 1, "a1", "Kuzu graph DB archived 2025") + seedChannelMessage(t, db, 1, "b1", "Kuzu graph DB looks promising") + + pkt, err := BuildContextPacket(ctx, svc, h1Agent, "Kuzu", InjectionOpts{ + BudgetTokens: 500, + MaxItems: 5, + MinScore: 0.0, + }) + if err != nil { + t.Fatalf("BuildContextPacket: %v", err) + } + if pkt == nil { + t.Fatal("expected non-nil packet") + } + for _, m := range pkt.Memories { + if m.FromAgent != "a1" { + t.Errorf("leaked memory from %q (owner != caller)", m.FromAgent) + } + if !strings.Contains(m.Body, "Kuzu") { + t.Errorf("unexpected body: %q", m.Body) + } + } +} + +func TestBuildContextPacket_TokenBudgetGreedyFillAndTruncate(t *testing.T) { + items := []MemoryItem{ + {ID: 1, Body: strings.Repeat("a", 200), Score: 0.9}, + {ID: 2, Body: strings.Repeat("b", 200), Score: 0.8}, + {ID: 3, Body: strings.Repeat("c", 200), Score: 0.7}, + } + // Budget 100 tokens => 400 chars total. Each item costs ~232 chars + // → ~58 tokens. First fits cleanly; second must truncate. + out := applyTokenBudget(items, 5, 100) + + if len(out) == 0 { + t.Fatal("expected at least one admitted item") + } + if out[0].ID != 1 { + t.Errorf("first admitted item ID = %d, want 1", out[0].ID) + } + total := 0 + for _, it := range out { + total += EstimateTokens(itemChars(it)) + } + if total > 100 { + t.Errorf("total tokens admitted = %d, exceeds budget 100", total) + } + sawTruncated := false + for _, it := range out { + if it.Truncated { + sawTruncated = true + } + } + if len(out) > 1 && !sawTruncated { + t.Errorf("expected truncation when admitting a second item under tight budget") + } +} + +func TestBuildContextPacket_ScoreFloorDrops(t *testing.T) { + svc, _, db := newTestServices(t) + ctx := context.Background() + _ = seedOwnedAgent(t, db, 1, "alice", "a1") + seedChannel(t, db, 1, "open-brain", "a1") + + cidPtr := int64(1) + low := &SearchResult{ + Message: &messaging.Message{ + ID: 100, + FromAgent: "a1", + Body: "low score body", + ChannelID: &cidPtr, + }, + SimilarityScore: 0.1, + MatchType: ModeSemantic, + } + high := &SearchResult{ + Message: &messaging.Message{ + ID: 101, + FromAgent: "a1", + Body: "high score body", + ChannelID: &cidPtr, + }, + SimilarityScore: 0.8, + MatchType: ModeSemantic, + } + + items, err := filterAndScore(ctx, svc, []*SearchResult{low, high}, "1", InjectionOpts{ + MinScore: 0.5, + }) + if err != nil { + t.Fatalf("filterAndScore: %v", err) + } + if len(items) != 1 { + t.Fatalf("expected 1 admitted item, got %d", len(items)) + } + if items[0].ID != 101 { + t.Errorf("admitted ID = %d, want 101 (high score)", items[0].ID) + } +} + +func TestBuildContextPacket_CoreMemoryIncluded(t *testing.T) { + svc, _, db := newTestServices(t) + ctx := context.Background() + a := seedOwnedAgent(t, db, 1, "alice", "a1") + + provider := &stubCoreProvider{blob: "I am alice's research agent."} + pkt, err := BuildContextPacket(ctx, svc, a, "", InjectionOpts{ + BudgetTokens: 500, + MaxItems: 5, + IncludeCore: true, + CoreProvider: provider, + }) + if err != nil { + t.Fatalf("BuildContextPacket: %v", err) + } + if pkt == nil { + t.Fatal("expected non-nil packet when core memory is present") + } + if pkt.CoreMemory != provider.blob { + t.Errorf("CoreMemory = %q, want %q", pkt.CoreMemory, provider.blob) + } + if pkt.PacketChars < len(provider.blob) { + t.Errorf("PacketChars = %d, want >= %d", pkt.PacketChars, len(provider.blob)) + } + if pkt.PacketTokenEstimate != EstimateTokens(pkt.PacketChars) { + t.Errorf("PacketTokenEstimate inconsistent with PacketChars") + } +} + +func TestBuildContextPacket_EmptyReturnsNil(t *testing.T) { + svc, _, db := newTestServices(t) + ctx := context.Background() + a := seedOwnedAgent(t, db, 1, "alice", "a1") + + // No memories, no core provider → should return (nil, nil) so the + // wrapper omits the relevant_context field entirely. + pkt, err := BuildContextPacket(ctx, svc, a, "", InjectionOpts{ + BudgetTokens: 500, + MaxItems: 5, + }) + if err != nil { + t.Fatalf("BuildContextPacket: %v", err) + } + if pkt != nil { + t.Errorf("expected nil packet, got %+v", pkt) + } +} diff --git a/internal/storage/schema/027_remove_approval_noise.sql b/internal/storage/schema/027_remove_approval_noise.sql new file mode 100644 index 0000000..e0baa0a --- /dev/null +++ b/internal/storage/schema/027_remove_approval_noise.sql @@ -0,0 +1,61 @@ +-- 027_remove_approval_noise.sql — one-shot cleanup for internal-only mode. +-- Drops messages produced by the (now-deleted) stalemate reminder/escalation +-- paths and any pending agent proposals. The agent_proposals table itself +-- stays in case the propose_agent flow is ever reinstated. +-- +-- The #approvals channel row is intentionally NOT dropped: leaving it +-- keeps the migration trivially reversible (no need to recreate the +-- channel and re-grant memberships); the human can delete it via admin +-- CLI later. +-- +-- Reminder/escalation messages are identified by their metadata JSON +-- (see internal/messaging/stalemate.go in git history), which is the +-- only field the historical worker stamped reliably — the conversation +-- subject prefix was used too but is not a hard guarantee. +-- +-- FK NOTES: most refs to messages(id) are ON DELETE CASCADE / SET NULL, +-- but two are NO ACTION and would block this migration on a populated +-- DB: +-- * attachments.message_id (001_initial.sql) +-- * messages.reply_to (007_threads.sql) +-- We NULL those explicitly before deleting. Orphan attachment rows (with +-- message_id = NULL) are tolerated by the attachments service. + +-- 1. Collect the message ids we're about to drop into a temp scratch +-- table so the cascading NULL/DELETE statements all reference the +-- same set even if the metadata predicates evolve later. +CREATE TEMP TABLE _approval_noise_msgs AS +SELECT id FROM messages +WHERE metadata LIKE '%"stalemate_reminder_for":%' + OR metadata LIKE '%"stalemate_escalation_for":%' + OR metadata LIKE '%"workflow_stalemate_reminder_for":%' + OR metadata LIKE '%"workflow_stalemate_escalation_for":%' + OR channel_id IN (SELECT id FROM channels WHERE name = 'approvals'); + +-- 2. Detach NO-ACTION FKs that would otherwise block the delete. +UPDATE attachments + SET message_id = NULL + WHERE message_id IN (SELECT id FROM _approval_noise_msgs); + +UPDATE messages + SET reply_to = NULL + WHERE reply_to IN (SELECT id FROM _approval_noise_msgs); + +-- 3. Drop the messages themselves. Cascading FKs (reactions, embeddings, +-- fts triggers, agent_listings) clean up automatically; SET NULL FKs +-- (goals.*_message_id, agent_proposals.*_message_id, goal_tasks.*) +-- forget the link gracefully. +DELETE FROM messages + WHERE id IN (SELECT id FROM _approval_noise_msgs); + +DROP TABLE _approval_noise_msgs; + +-- Conversation rows for stalemate-prefixed subjects are intentionally +-- left in place. Multiple NO-ACTION FKs reference conversations(id) +-- (messages.conversation_id, inbox_state.conversation_id, possibly +-- others added by future migrations); cleaning them up reliably would +-- require chasing every ref. Empty / mostly-empty conversation rows +-- are harmless — they don't render as messages in the UI. + +-- 4. Drop pending agent proposals. The table stays for reversibility. +DELETE FROM agent_proposals;