feat(020): US1 — proactive injection on MCP tool responses

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) <noreply@anthropic.com>
This commit is contained in:
Algis Dumbris
2026-05-11 15:09:38 +03:00
co-authored by Claude Opus 4.7
parent da827c03b3
commit 8a5d5e1f59
15 changed files with 1934 additions and 1430 deletions
+2 -17
View File
@@ -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
@@ -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.
+8 -82
View File
@@ -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 {
+197
View File
@@ -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
}
+251
View File
@@ -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: <packet>` 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:])
}
+218
View File
@@ -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)
}
}
}
+3 -2
View File
@@ -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
+116 -6
View File
@@ -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: "<recent activity>" (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.")),
+147
View File
@@ -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
}
@@ -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])
}
}
+60 -697
View File
@@ -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)
}
+16 -626
View File
@@ -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)
}
}
+329
View File
@@ -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
}
+258
View File
@@ -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)
}
}
@@ -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;