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:
co-authored by
Claude Opus 4.7
parent
da827c03b3
commit
8a5d5e1f59
+2
-17
@@ -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.
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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:])
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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.")),
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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;
|
||||
Reference in New Issue
Block a user