From 2044b199b852390a0b8ca84fadb50c49803dbf3c Mon Sep 17 00:00:00 2001 From: Algis Dumbris Date: Mon, 11 May 2026 15:40:49 +0300 Subject: [PATCH] =?UTF-8?q?feat(020):=20US3=20=E2=80=94=20dream=20worker?= =?UTF-8?q?=20+=206=20MCP=20consolidation=20tools?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Background ConsolidatorWorker dispatches consolidation work to a Claude Code agent through harness.Harness.Execute (NOT via system DMs — per feedback_system_dm_no_trigger.md) with a one-time 15m dispatch token. The dispatched agent uses six new MCP tools, all token-gated and recording every action to memory_consolidation_jobs.actions JSON. Stores: - memory_links.go (+ test): typed edges with actor-prefix reserved-type guard; AddConsolidationLink bypass for memory_mark_duplicate / memory_supersede (their contractual writers). - memory_pins.go (+ test): owner pin overlay, bypasses score floor. - memory_status.go: queries the memory_status view. - consolidation_jobs.go: Create / Dispatch / Lease / AppendAction / Complete with ErrJobAlreadyInFlight via partial unique index. - auto_links.go: MessageListener generating mention / reply_to / channel_cooccurrence links automatically on send. Worker: - consolidator.go (+ test): ticker pattern modeled on StalemateWorker. Watermark trigger for link_gen / dedup_contradiction; daily 03:00 UTC for sleep-time core rewrite. Wallclock budget via harness Budget. Global semaphore gates concurrent owners. Mocked-harness test asserts no system DM is ever sent. - consolidator_prompts.go: four job-type prompts passed via env to the dispatched agent. MCP tools (internal/mcp/memory_tools.go + test): - memory_list_unprocessed, memory_write_reflection, memory_rewrite_core, memory_mark_duplicate, memory_supersede, memory_add_link. - Full error-code matrix tested per contracts/mcp-memory-tools.md. - Registered only when SYNAPBUS_DREAM_ENABLED=1. Injection extensions: - search/injection.go: pin overlay applied after retrieval; status filter drops soft_deleted / superseded unless pinned. New PinProvider, StatusProvider, MessageLookup hooks on InjectionOpts. Wiring: - cmd/synapbus/main.go: stores constructed, AutoLinkListener attached to MessagingService, mcpSrv.SetDream wired, ConsolidatorWorker start/stop, admin DreamRun closure. - cmd/synapbus/admin.go: synapbus memory dream-run --owner --job socket-RPC command (forces a single job bypassing trigger). Cycle workarounds (documented in code): - messaging.DreamAgent / HarnessDispatcher are local interfaces (the agents and harness packages import messaging, not the reverse). main.go wraps the real types via adapter structs. Stubbed: - Cron expression parsing (DreamDeepCron). Hardcoded daily 03:00 UTC. Adding robfig/cron deferred to keep no-new-deps. Pre-existing reactor test failures (5) are unchanged; confirmed pre-020 via stash check. Co-Authored-By: Claude Opus 4.7 (1M context) --- cmd/synapbus/admin.go | 25 + cmd/synapbus/main.go | 135 ++++ internal/admin/server.go | 6 + internal/admin/socket.go | 34 ++ internal/mcp/memory_tools.go | 642 ++++++++++++++++++++ internal/mcp/memory_tools_test.go | 375 ++++++++++++ internal/mcp/server.go | 16 + internal/messaging/auto_links.go | 205 +++++++ internal/messaging/consolidation_jobs.go | 352 +++++++++++ internal/messaging/consolidator.go | 485 +++++++++++++++ internal/messaging/consolidator_prompts.go | 63 ++ internal/messaging/consolidator_test.go | 204 +++++++ internal/messaging/memory_links.go | 264 ++++++++ internal/messaging/memory_links_test.go | 152 +++++ internal/messaging/memory_pins.go | 155 +++++ internal/messaging/memory_pins_test.go | 81 +++ internal/messaging/memory_status.go | 102 ++++ internal/search/injection.go | 165 ++++- internal/search/injection_pin_test.go | 137 +++++ internal/storage/memory_status_view_test.go | 126 ++++ 20 files changed, 3721 insertions(+), 3 deletions(-) create mode 100644 internal/mcp/memory_tools.go create mode 100644 internal/mcp/memory_tools_test.go create mode 100644 internal/messaging/auto_links.go create mode 100644 internal/messaging/consolidation_jobs.go create mode 100644 internal/messaging/consolidator.go create mode 100644 internal/messaging/consolidator_prompts.go create mode 100644 internal/messaging/consolidator_test.go create mode 100644 internal/messaging/memory_links.go create mode 100644 internal/messaging/memory_links_test.go create mode 100644 internal/messaging/memory_pins.go create mode 100644 internal/messaging/memory_pins_test.go create mode 100644 internal/messaging/memory_status.go create mode 100644 internal/search/injection_pin_test.go create mode 100644 internal/storage/memory_status_view_test.go diff --git a/cmd/synapbus/admin.go b/cmd/synapbus/admin.go index e57916f..82d0079 100644 --- a/cmd/synapbus/admin.go +++ b/cmd/synapbus/admin.go @@ -1439,6 +1439,31 @@ Examples: memoryCoreCmd.AddCommand(memoryCoreGetCmd, memoryCoreSetCmd, memoryCoreDeleteCmd) memoryCmd.AddCommand(memoryCoreCmd) + // ----- memory dream-run (feature 020 — manual dispatch) ----- + var ( + dreamRunOwner string + dreamRunJobType string + ) + memoryDreamRunCmd := &cobra.Command{ + Use: "dream-run", + Short: "Force a single consolidation job dispatch (bypasses trigger checks)", + RunE: func(cmd *cobra.Command, args []string) error { + resp, err := adminRequest("memory.dream_run", map[string]string{ + "owner": dreamRunOwner, + "job_type": dreamRunJobType, + }) + if err != nil { + return err + } + printJSON(resp["data"]) + return nil + }, + } + memoryDreamRunCmd.Flags().StringVar(&dreamRunOwner, "owner", "", "Owner username or numeric user ID") + memoryDreamRunCmd.Flags().StringVar(&dreamRunJobType, "job", "reflection", "Job type (reflection | core_rewrite | dedup_contradiction | link_gen)") + _ = memoryDreamRunCmd.MarkFlagRequired("owner") + memoryCmd.AddCommand(memoryDreamRunCmd) + // ----- add persistent flag and commands to root ----- rootCmd.PersistentFlags().StringVar(&adminSocket, "socket", "/tmp/synapbus.sock", "Path to admin Unix socket") diff --git a/cmd/synapbus/main.go b/cmd/synapbus/main.go index 9230503..29c4c65 100644 --- a/cmd/synapbus/main.go +++ b/cmd/synapbus/main.go @@ -600,6 +600,31 @@ func runServe(cmd *cobra.Command, args []string) error { "core_max_bytes", memCfg.CoreMemoryMaxBytes, ) + // Feature 020 — dream worker (US3) stores. These are always + // constructed so admin CLI / future REST endpoints can read them + // even when SYNAPBUS_DREAM_ENABLED=0. The worker itself starts + // only when the flag is on. + memoryLinkStore := messaging.NewLinkStore(db.DB) + memoryPinStore := messaging.NewPinStore(db.DB) + memoryJobsStore := messaging.NewJobsStore(db.DB) + dispatchTokens := messaging.NewDispatchTokenStore(db.DB) + + // Wire the auto-link emitter as a message listener (T035). + msgService.AddMessageListener(messaging.NewAutoLinkListener(db.DB, memoryLinkStore)) + + // Register the six memory_* MCP tools when SYNAPBUS_DREAM_ENABLED=1. + mcpSrv.SetDream(mcpserver.MemoryToolDeps{ + DB: db.DB, + Msg: msgService, + Agents: agentService, + Core: coreMemoryStore, + Links: memoryLinkStore, + Pins: memoryPinStore, + Jobs: memoryJobsStore, + Tokens: dispatchTokens, + MemConfig: memCfg, + }) + // Wire the agent marketplace (spec 016 MVP). marketplaceStore := marketplace.NewStore(db.DB) marketplaceSvc := marketplace.NewService(marketplaceStore, wikiService, swarmService, channelService, msgService, tracer) @@ -664,6 +689,7 @@ func runServe(cmd *cobra.Command, args []string) error { } else { stalemateConfig := messaging.ParseStalemateConfig() stalemateWorker = messaging.NewStalemateWorker(db.DB, msgService, stalemateConfig) + stalemateWorker.SetMemoryInjections(memoryInjectionStore) stalemateWorker.Start() slog.Info("stalemate worker started", "processing_timeout", stalemateConfig.ProcessingTimeout.String(), @@ -671,6 +697,29 @@ func runServe(cmd *cobra.Command, args []string) error { ) } + // Feature 020 — consolidator (dream) worker. Only starts when + // SYNAPBUS_DREAM_ENABLED=1. + var consolidator *messaging.ConsolidatorWorker + if memCfg.DreamEnabled { + consolidator = messaging.NewConsolidatorWorker( + db.DB, + memoryJobsStore, + dispatchTokens, + &harnessDispatcherAdapter{reg: harnessRegistry}, + &agentLookupAdapter{svc: agentService}, + memCfg, + ) + consolidator.Start() + slog.Info("consolidator (dream) worker started", + "interval", memCfg.DreamInterval.String(), + "watermark", memCfg.DreamWatermark, + "max_concurrent", memCfg.DreamMaxConcurrent, + "agent", memCfg.DreamAgent, + ) + } else { + slog.Info("consolidator (dream) worker disabled (SYNAPBUS_DREAM_ENABLED=0)") + } + // Create health checker healthChecker := health.NewChecker(db.DB, version) @@ -839,6 +888,14 @@ func runServe(cmd *cobra.Command, args []string) error { adminSvcs.WebhookService = webhookService adminSvcs.K8sService = k8sService adminSvcs.CoreMemoryStore = coreMemoryStore + if consolidator != nil { + // Closure form keeps admin's import graph independent of + // messaging.ConsolidatorWorker's full surface. + c := consolidator + adminSvcs.DreamRun = func(ctx context.Context, ownerID, jobType string) (int64, error) { + return c.ForceRun(ctx, ownerID, jobType) + } + } adminServer := admin.NewServer(adminSocketPath, db.DB, adminSvcs, logger) if err := adminServer.Start(); err != nil { return fmt.Errorf("start admin socket: %w", err) @@ -904,6 +961,11 @@ func runServe(cmd *cobra.Command, args []string) error { stalemateWorker.Stop() } + // Stop consolidator (dream) worker + if consolidator != nil { + consolidator.Stop() + } + // Stop embedding pipeline if embPipeline != nil { embPipeline.Stop() @@ -1216,3 +1278,76 @@ func (a *messageAuthorResolverAdapter) GetMessageAuthor(ctx context.Context, mes } return msg.FromAgent, nil } + +// agentLookupAdapter adapts agents.AgentService to +// messaging.AgentLookup so the dream worker can resolve the dream-agent +// record without dragging the full *agents.AgentService into the +// messaging package. The returned messaging.DreamAgent is the raw +// *agents.Agent itself — DreamAgent's only required method +// (AgentName()) is satisfied by agents.Agent.Name via the +// agentNameMethod helper below. +type agentLookupAdapter struct { + svc *agents.AgentService +} + +func (a *agentLookupAdapter) GetAgent(ctx context.Context, name string) (messaging.DreamAgent, error) { + ag, err := a.svc.GetAgent(ctx, name) + if err != nil { + return nil, err + } + return agentDreamWrap{ag: ag}, nil +} + +// agentDreamWrap adapts *agents.Agent to messaging.DreamAgent. +type agentDreamWrap struct{ ag *agents.Agent } + +func (w agentDreamWrap) AgentName() string { + if w.ag == nil { + return "" + } + return w.ag.Name +} + +// harnessDispatcherAdapter adapts *harness.Registry to +// messaging.HarnessDispatcher so the consolidator worker can dispatch +// dream-agent runs without importing the harness package (which would +// create an import cycle — harness already imports messaging). +type harnessDispatcherAdapter struct { + reg *harness.Registry +} + +func (a *harnessDispatcherAdapter) Execute( + ctx context.Context, + agent messaging.DreamAgent, + req *messaging.HarnessExecRequest, +) (*messaging.HarnessExecResult, error) { + // Unbox the agent record. The worker stores a DreamAgent + // interface; in production it's an agentDreamWrap holding the + // real *agents.Agent. Tests / admin force-runs may pass a bare + // DreamAgentNamed which has no underlying record — the harness + // fallback chain then resolves the backend by name alone. + var realAgent *agents.Agent + if wrap, ok := agent.(agentDreamWrap); ok { + realAgent = wrap.ag + } + hreq := &harness.ExecRequest{ + RunID: req.RunID, + AgentName: req.AgentName, + Agent: realAgent, + Env: req.Env, + Budget: harness.Budget{MaxWallClock: req.MaxWallClock}, + } + if req.Body != "" { + hreq.Message = &messaging.Message{Body: req.Body} + } + res, err := a.reg.Execute(ctx, realAgent, hreq) + if err != nil { + return nil, err + } + out := &messaging.HarnessExecResult{} + if res != nil { + out.ExitCode = res.ExitCode + out.Logs = res.Logs + } + return out, nil +} diff --git a/internal/admin/server.go b/internal/admin/server.go index 0eb2918..a31fb12 100644 --- a/internal/admin/server.go +++ b/internal/admin/server.go @@ -53,6 +53,12 @@ type Services struct { // in for feature 020 admin CLI commands (`synapbus memory core ...`). // May be nil — handlers report "core memory store not configured". CoreMemoryStore *messaging.CoreMemoryStore + + // DreamRun, when non-nil, dispatches a single consolidation job + // bypassing the trigger check. Wired by main.go when the + // consolidator worker is enabled. Closure form keeps the worker + // internals out of the admin package's import graph. + DreamRun func(ctx context.Context, ownerID, jobType string) (jobID int64, err error) } // RetentionStatusProvider provides retention status information. diff --git a/internal/admin/socket.go b/internal/admin/socket.go index a22bbc4..c9ca6e8 100644 --- a/internal/admin/socket.go +++ b/internal/admin/socket.go @@ -233,6 +233,8 @@ func (s *AdminServer) dispatch(req Request) Response { return s.handleMemoryCoreSet(ctx, req.Args) case "memory.core.delete": return s.handleMemoryCoreDelete(ctx, req.Args) + case "memory.dream_run": + return s.handleMemoryDreamRun(ctx, req.Args) default: return Response{OK: false, Error: fmt.Sprintf("unknown command: %s", req.Command)} @@ -1985,5 +1987,37 @@ func (s *AdminServer) handleMemoryCoreDelete(ctx context.Context, args json.RawM }} } +// handleMemoryDreamRun forces a single dream-job dispatch. Bypasses +// trigger checks — useful for kubic verification (quickstart §"verify +// dream agent"). Returns the created job_id. +func (s *AdminServer) handleMemoryDreamRun(ctx context.Context, args json.RawMessage) Response { + var p struct { + Owner string `json:"owner"` + JobType string `json:"job_type"` + } + if err := json.Unmarshal(args, &p); err != nil { + return Response{OK: false, Error: "invalid args: " + err.Error()} + } + if p.JobType == "" { + return Response{OK: false, Error: "job_type is required"} + } + if s.services.DreamRun == nil { + return Response{OK: false, Error: "dream worker not configured (SYNAPBUS_DREAM_ENABLED=0?)"} + } + ownerStr, err := s.resolveOwnerString(ctx, p.Owner) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + jobID, err := s.services.DreamRun(ctx, ownerStr, p.JobType) + if err != nil { + return Response{OK: false, Error: err.Error()} + } + return Response{OK: true, Data: map[string]any{ + "job_id": jobID, + "owner_id": ownerStr, + "job_type": p.JobType, + }} +} + // Ensure the messaging import is used. var _ = messaging.StatusPending diff --git a/internal/mcp/memory_tools.go b/internal/mcp/memory_tools.go new file mode 100644 index 0000000..3f7fd68 --- /dev/null +++ b/internal/mcp/memory_tools.go @@ -0,0 +1,642 @@ +// MCP memory-consolidation tools (feature 020 — dream worker, US3). +// +// Registered only when SYNAPBUS_DREAM_ENABLED=1 via +// MemoryToolRegistrar.RegisterAllOnServer (see server.go SetDream). +// +// Every tool: +// +// 1. Pulls the dispatch token from the request context (set by +// MCP middleware reading X-Synapbus-Dispatch-Token from the +// transport header, or via the harness-propagated env var). +// 2. Validates the token against (caller-supplied owner_id, the +// active consolidation_job_id carried alongside the token). +// 3. Performs the action against the appropriate messaging store. +// 4. Appends a structured action record to the job's `actions` +// JSON array via JobsStore.AppendAction. +// +// All errors follow the MCP standard envelope; the contract codes +// (`dispatch_token_*`, `not_same_owner`, `core_memory_too_large`, +// `relation_type_reserved`, `source_not_found`, ...) are listed in +// `contracts/mcp-memory-tools.md` and mirrored verbatim here so the +// dream-agent can match on the string. +package mcp + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "log/slog" + "strings" + "time" + + 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" +) + +// Context key for the dispatch token. The transport-layer middleware +// (or test harness) stuffs the token from the HTTP header +// `X-Synapbus-Dispatch-Token` into ctx via WithDispatchToken; the +// memory tool handlers read it via DispatchTokenFromContext. +type dispatchTokenKey struct{} + +// WithDispatchToken returns a derived context carrying tok as the +// active dispatch token. +func WithDispatchToken(ctx context.Context, tok string) context.Context { + return context.WithValue(ctx, dispatchTokenKey{}, tok) +} + +// DispatchTokenFromContext returns the dispatch token, if any. +func DispatchTokenFromContext(ctx context.Context) (string, bool) { + v, ok := ctx.Value(dispatchTokenKey{}).(string) + return v, ok && v != "" +} + +// MemoryToolDeps bundles the dependencies the six memory tools need. +type MemoryToolDeps struct { + DB *sql.DB + Msg *messaging.MessagingService + Agents *agents.AgentService + Core *messaging.CoreMemoryStore + Links *messaging.LinkStore + Pins *messaging.PinStore + Jobs *messaging.JobsStore + Tokens *messaging.DispatchTokenStore + MemConfig messaging.MemoryConfig + Logger *slog.Logger +} + +// MemoryToolRegistrar registers the six memory_* MCP tools. +type MemoryToolRegistrar struct { + deps MemoryToolDeps + logger *slog.Logger +} + +// NewMemoryToolRegistrar returns a registrar over deps. RegisterAllOnServer +// is a no-op when deps.DB or deps.Jobs is nil (defensive — these are required). +func NewMemoryToolRegistrar(deps MemoryToolDeps) *MemoryToolRegistrar { + logger := deps.Logger + if logger == nil { + logger = slog.Default().With("component", "mcp-memory-tools") + } + return &MemoryToolRegistrar{deps: deps, logger: logger} +} + +// RegisterAllOnServer attaches the six tools to mcpSrv. The caller +// (server.go SetDream) is responsible for gating registration on +// SYNAPBUS_DREAM_ENABLED. +func (r *MemoryToolRegistrar) RegisterAllOnServer(s *server.MCPServer) { + if s == nil || r.deps.DB == nil || r.deps.Jobs == nil || r.deps.Tokens == nil { + return + } + s.AddTool(memoryListUnprocessedTool(), r.handleListUnprocessed) + s.AddTool(memoryWriteReflectionTool(), r.handleWriteReflection) + s.AddTool(memoryRewriteCoreTool(), r.handleRewriteCore) + s.AddTool(memoryMarkDuplicateTool(), r.handleMarkDuplicate) + s.AddTool(memorySupersedeTool(), r.handleSupersede) + s.AddTool(memoryAddLinkTool(), r.handleAddLink) + r.logger.Info("memory MCP tools registered", "count", 6) +} + +// --- Tool definitions --- + +func memoryListUnprocessedTool() mcplib.Tool { + return mcplib.NewTool("memory_list_unprocessed", + mcplib.WithDescription("List recent memory-eligible messages the owner's pool has not yet consolidated. Used by the dream agent to scan its inbox."), + mcplib.WithString("owner_id", mcplib.Description("Caller's owner_id (must match the dispatch token's owner)"), mcplib.Required()), + mcplib.WithNumber("since_message_id", mcplib.Description("Exclusive lower bound (defaults to 0)")), + mcplib.WithNumber("limit", mcplib.Description("Max items to return (default 50, max 200)")), + ) +} + +func memoryWriteReflectionTool() mcplib.Tool { + return mcplib.NewTool("memory_write_reflection", + mcplib.WithDescription("Write a higher-level abstraction back to the memory pool tagged 'reflection'. Inserts 'refines' links from the new memory to each source."), + mcplib.WithString("owner_id", mcplib.Required()), + mcplib.WithString("body", mcplib.Required()), + mcplib.WithString("source_message_ids", mcplib.Description("Comma-separated message ids")), + mcplib.WithString("tags", mcplib.Description("Comma-separated tags")), + ) +} + +func memoryRewriteCoreTool() mcplib.Tool { + return mcplib.NewTool("memory_rewrite_core", + mcplib.WithDescription("Replace the per-(owner, agent) core memory blob wholesale (no merge)."), + mcplib.WithString("owner_id", mcplib.Required()), + mcplib.WithString("agent_name", mcplib.Required()), + mcplib.WithString("blob", mcplib.Required()), + ) +} + +func memoryMarkDuplicateTool() mcplib.Tool { + return mcplib.NewTool("memory_mark_duplicate", + mcplib.WithDescription("Mark two memories as duplicates; one is kept canonical, the other soft-deleted."), + mcplib.WithString("owner_id", mcplib.Required()), + mcplib.WithNumber("a_id", mcplib.Required()), + mcplib.WithNumber("b_id", mcplib.Required()), + mcplib.WithNumber("keep_id", mcplib.Required()), + mcplib.WithString("reason"), + ) +} + +func memorySupersedeTool() mcplib.Tool { + return mcplib.NewTool("memory_supersede", + mcplib.WithDescription("Mark memory A as obsoleted by memory B (temporal validity)."), + mcplib.WithString("owner_id", mcplib.Required()), + mcplib.WithNumber("a_id", mcplib.Required()), + mcplib.WithNumber("b_id", mcplib.Required()), + mcplib.WithString("reason"), + ) +} + +func memoryAddLinkTool() mcplib.Tool { + return mcplib.NewTool("memory_add_link", + mcplib.WithDescription("Add a typed link between two memories. relation_type must be one of refines, contradicts, examples, related."), + mcplib.WithString("owner_id", mcplib.Required()), + mcplib.WithNumber("src_id", mcplib.Required()), + mcplib.WithNumber("dst_id", mcplib.Required()), + mcplib.WithString("relation_type", mcplib.Required()), + mcplib.WithString("metadata", mcplib.Description("JSON object")), + ) +} + +// --- Shared validation --- + +// authorizeForOwner validates the dispatch token in ctx against the +// caller-supplied owner_id. Returns the active jobID (so the handler +// can call AppendAction) or an MCP error result. +func (r *MemoryToolRegistrar) authorizeForOwner(ctx context.Context, ownerID string) (jobID int64, errResult *mcplib.CallToolResult) { + tok, ok := DispatchTokenFromContext(ctx) + if !ok { + return 0, memErrorf("dispatch_token_missing", "no dispatch token in request context") + } + // Find the consolidation_job_id this token is bound to. + var ( + dbOwner string + dbJob int64 + expiresAt time.Time + revokedAt sql.NullTime + ) + err := r.deps.DB.QueryRowContext(ctx, + `SELECT owner_id, consolidation_job_id, expires_at, revoked_at + FROM memory_dispatch_tokens WHERE token = ?`, tok, + ).Scan(&dbOwner, &dbJob, &expiresAt, &revokedAt) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return 0, memErrorf("dispatch_token_missing", "token not found") + } + return 0, memErrorf("dispatch_token_missing", "token lookup failed: %s", err) + } + if revokedAt.Valid { + return 0, memErrorf("dispatch_token_revoked", "token has been revoked") + } + if !expiresAt.After(time.Now().UTC()) { + return 0, memErrorf("dispatch_token_expired", "token expired at %s", expiresAt.Format(time.RFC3339)) + } + if dbOwner != ownerID { + return 0, memErrorf("dispatch_token_owner_mismatch", "token bound to %q, request claims %q", dbOwner, ownerID) + } + // Run the canonical Validate path so used_at is stamped uniformly. + if r.deps.Tokens != nil { + if _, err := r.deps.Tokens.Validate(ctx, tok, ownerID, dbJob); err != nil { + return 0, memErrorf("dispatch_token_missing", "validate: %s", err) + } + } + return dbJob, nil +} + +// recordAction appends to the job's actions JSON array. Logged at +// warn-level on failure; never blocks the tool's user-visible response. +func (r *MemoryToolRegistrar) recordAction(ctx context.Context, jobID int64, tool string, targetID int64, args map[string]any) { + if r.deps.Jobs == nil { + return + } + action := map[string]any{ + "tool": tool, + "target_message_id": targetID, + "args": args, + "at": time.Now().UTC().Format(time.RFC3339), + } + if err := r.deps.Jobs.AppendAction(ctx, jobID, action); err != nil { + r.logger.Warn("append action failed", + "job_id", jobID, + "tool", tool, + "error", err, + ) + } +} + +// memErrorf returns an MCP CallToolResult carrying a contractual error +// code + human-readable message. The MCP framework already wraps the +// `error` field in the JSON envelope; we render `code: ...` as the +// leading line of the message so the dream-agent can pattern-match. +func memErrorf(code, format string, args ...any) *mcplib.CallToolResult { + msg := fmt.Sprintf("%s: %s", code, fmt.Sprintf(format, args...)) + return mcplib.NewToolResultError(msg) +} + +// --- Handlers --- + +func (r *MemoryToolRegistrar) handleListUnprocessed(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) { + owner := req.GetString("owner_id", "") + if owner == "" { + return memErrorf("invalid_request", "owner_id required"), nil + } + jobID, errR := r.authorizeForOwner(ctx, owner) + if errR != nil { + return errR, nil + } + + since := int64(req.GetInt("since_message_id", 0)) + limit := req.GetInt("limit", 50) + if limit <= 0 { + limit = 50 + } + if limit > 200 { + limit = 200 + } + + memIDs, err := messaging.MemoryChannelIDs(ctx, r.deps.DB) + if err != nil { + return memErrorf("internal", "list memory channels: %s", err), nil + } + if len(memIDs) == 0 { + _ = jobID // record nothing — empty + return resultJSON(map[string]any{"memories": []any{}, "max_id_returned": since}) + } + + placeholders := strings.Repeat("?,", len(memIDs)) + placeholders = placeholders[:len(placeholders)-1] + queryArgs := []any{} + for _, id := range memIDs { + queryArgs = append(queryArgs, id) + } + queryArgs = append(queryArgs, owner, since, limit) + + q := `SELECT m.id, m.from_agent, c.name, m.body, m.created_at + FROM messages m + JOIN agents a ON m.from_agent = a.name + JOIN channels c ON m.channel_id = c.id + WHERE m.channel_id IN (` + placeholders + `) + AND CAST(a.owner_id AS TEXT) = ? + AND m.id > ? + ORDER BY m.id ASC + LIMIT ?` + + rows, err := r.deps.DB.QueryContext(ctx, q, queryArgs...) + if err != nil { + return memErrorf("internal", "query: %s", err), nil + } + defer rows.Close() + + type item struct { + ID int64 `json:"id"` + FromAgent string `json:"from_agent"` + Channel string `json:"channel"` + Body string `json:"body"` + CreatedAt time.Time `json:"created_at"` + } + var items []item + var maxID = since + for rows.Next() { + var it item + if err := rows.Scan(&it.ID, &it.FromAgent, &it.Channel, &it.Body, &it.CreatedAt); err != nil { + return memErrorf("internal", "scan: %s", err), nil + } + items = append(items, it) + if it.ID > maxID { + maxID = it.ID + } + } + if err := rows.Err(); err != nil { + return memErrorf("internal", "iterate: %s", err), nil + } + + r.recordAction(ctx, jobID, "memory_list_unprocessed", 0, map[string]any{ + "since_message_id": since, "limit": limit, "returned": len(items), + }) + return resultJSON(map[string]any{"memories": items, "max_id_returned": maxID}) +} + +func (r *MemoryToolRegistrar) handleWriteReflection(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) { + owner := req.GetString("owner_id", "") + body := req.GetString("body", "") + if owner == "" || body == "" { + return memErrorf("invalid_request", "owner_id and body required"), nil + } + jobID, errR := r.authorizeForOwner(ctx, owner) + if errR != nil { + return errR, nil + } + sourceIDs := parseInt64CSV(req.GetString("source_message_ids", "")) + + // Verify every source belongs to caller's owner. + for _, sid := range sourceIDs { + ok, sameOwner, _ := r.messageBelongsTo(ctx, sid, owner) + if !ok { + return memErrorf("source_not_found", "message %d not found", sid), nil + } + if !sameOwner { + return memErrorf("not_same_owner", "source %d belongs to a different owner", sid), nil + } + } + + // Pick a destination channel: prefer `#reflections-` if it + // exists, else `#open-brain`. + channelID, channelName, err := r.pickReflectionChannel(ctx, owner) + if err != nil { + return memErrorf("internal", "pick channel: %s", err), nil + } + if channelID == 0 { + return memErrorf("internal", "no reflection channel available"), nil + } + + dreamAgent := "dream:" + owner + // Ensure conversation + message inserts (lightweight direct SQL — + // the MessagingService path would trigger reactive runs which we + // must avoid per feedback_system_dm_no_trigger.md). + convRes, err := r.deps.DB.ExecContext(ctx, + `INSERT INTO conversations (created_by, channel_id) VALUES (?, ?)`, + dreamAgent, channelID, + ) + if err != nil { + return memErrorf("internal", "create conversation: %s", err), nil + } + convID, _ := convRes.LastInsertId() + + res, err := r.deps.DB.ExecContext(ctx, + `INSERT INTO messages (conversation_id, from_agent, channel_id, body, priority, status, metadata) + VALUES (?, ?, ?, ?, 5, 'pending', ?)`, + convID, dreamAgent, channelID, body, `{"tags":["reflection"]}`, + ) + if err != nil { + return memErrorf("internal", "insert message: %s", err), nil + } + newID, _ := res.LastInsertId() + + // Add `refines` links from new memory → each source. + created := 0 + for _, sid := range sourceIDs { + if r.deps.Links == nil { + break + } + if _, err := r.deps.Links.Add(ctx, newID, sid, "refines", owner, "agent:dream:"+owner, nil); err == nil { + created++ + } + } + + r.recordAction(ctx, jobID, "memory_write_reflection", newID, map[string]any{ + "source_message_ids": sourceIDs, + "channel": channelName, + }) + + return resultJSON(map[string]any{ + "memory_id": newID, + "channel": channelName, + "links_created": created, + }) +} + +func (r *MemoryToolRegistrar) handleRewriteCore(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) { + owner := req.GetString("owner_id", "") + agent := req.GetString("agent_name", "") + blob := req.GetString("blob", "") + if owner == "" || agent == "" { + return memErrorf("invalid_request", "owner_id and agent_name required"), nil + } + jobID, errR := r.authorizeForOwner(ctx, owner) + if errR != nil { + return errR, nil + } + if r.deps.Core == nil { + return memErrorf("internal", "core memory store not configured"), nil + } + // Confirm target agent is owned by caller's owner. + targetOwner, err := agents.OwnerFor(ctx, r.deps.DB, agent) + if err != nil { + if errors.Is(err, agents.ErrAgentNotFound) { + return memErrorf("source_not_found", "agent %q not found", agent), nil + } + return memErrorf("internal", "owner lookup: %s", err), nil + } + if targetOwner != owner { + return memErrorf("not_same_owner", "agent %q owner=%q != caller %q", agent, targetOwner, owner), nil + } + prev, _, _, _ := r.deps.Core.Get(ctx, owner, agent) + if err := r.deps.Core.Set(ctx, owner, agent, blob, "agent:dream:"+owner); err != nil { + if errors.Is(err, messaging.ErrCoreMemoryTooLarge) { + return memErrorf("core_memory_too_large", "blob %d bytes exceeds cap", len(blob)), nil + } + return memErrorf("internal", "set core: %s", err), nil + } + r.recordAction(ctx, jobID, "memory_rewrite_core", 0, map[string]any{ + "owner_id": owner, "agent_name": agent, "new_chars": len(blob), + }) + return resultJSON(map[string]any{ + "owner_id": owner, + "agent_name": agent, + "previous_blob": prev, + "new_blob_chars": len(blob), + "updated_at": time.Now().UTC().Format(time.RFC3339), + }) +} + +func (r *MemoryToolRegistrar) handleMarkDuplicate(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) { + owner := req.GetString("owner_id", "") + aID := int64(req.GetInt("a_id", 0)) + bID := int64(req.GetInt("b_id", 0)) + keepID := int64(req.GetInt("keep_id", 0)) + reason := req.GetString("reason", "") + if owner == "" || aID == 0 || bID == 0 || keepID == 0 { + return memErrorf("invalid_request", "owner_id, a_id, b_id, keep_id required"), nil + } + if keepID != aID && keepID != bID { + return memErrorf("keep_id_not_in_pair", "keep_id must be a_id or b_id"), nil + } + jobID, errR := r.authorizeForOwner(ctx, owner) + if errR != nil { + return errR, nil + } + for _, id := range []int64{aID, bID} { + ok, sameOwner, _ := r.messageBelongsTo(ctx, id, owner) + if !ok { + return memErrorf("source_not_found", "message %d not found", id), nil + } + if !sameOwner { + return memErrorf("not_same_owner", "message %d belongs to a different owner", id), nil + } + } + loserID := aID + if keepID == aID { + loserID = bID + } + if r.deps.Links == nil { + return memErrorf("internal", "link store not configured"), nil + } + linkID, err := r.deps.Links.AddConsolidationLink(ctx, loserID, keepID, "duplicate_of", owner, "agent:dream:"+owner, map[string]any{"reason": reason}) + if err != nil { + return memErrorf("internal", "add link: %s", err), nil + } + r.recordAction(ctx, jobID, "memory_mark_duplicate", loserID, map[string]any{ + "a_id": aID, "b_id": bID, "keep_id": keepID, "reason": reason, + }) + return resultJSON(map[string]any{ + "keep_id": keepID, + "soft_deleted_id": loserID, + "link_created_id": linkID, + }) +} + +func (r *MemoryToolRegistrar) handleSupersede(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) { + owner := req.GetString("owner_id", "") + aID := int64(req.GetInt("a_id", 0)) + bID := int64(req.GetInt("b_id", 0)) + reason := req.GetString("reason", "") + if owner == "" || aID == 0 || bID == 0 { + return memErrorf("invalid_request", "owner_id, a_id, b_id required"), nil + } + jobID, errR := r.authorizeForOwner(ctx, owner) + if errR != nil { + return errR, nil + } + for _, id := range []int64{aID, bID} { + ok, sameOwner, _ := r.messageBelongsTo(ctx, id, owner) + if !ok { + return memErrorf("source_not_found", "message %d not found", id), nil + } + if !sameOwner { + return memErrorf("not_same_owner", "message %d belongs to a different owner", id), nil + } + } + if r.deps.Links == nil { + return memErrorf("internal", "link store not configured"), nil + } + linkID, err := r.deps.Links.AddConsolidationLink(ctx, aID, bID, "superseded_by", owner, "agent:dream:"+owner, map[string]any{"reason": reason}) + if err != nil { + return memErrorf("internal", "add link: %s", err), nil + } + r.recordAction(ctx, jobID, "memory_supersede", aID, map[string]any{ + "a_id": aID, "b_id": bID, "reason": reason, + }) + return resultJSON(map[string]any{ + "superseded_id": aID, + "by_id": bID, + "link_created_id": linkID, + }) +} + +func (r *MemoryToolRegistrar) handleAddLink(ctx context.Context, req mcplib.CallToolRequest) (*mcplib.CallToolResult, error) { + owner := req.GetString("owner_id", "") + srcID := int64(req.GetInt("src_id", 0)) + dstID := int64(req.GetInt("dst_id", 0)) + relType := req.GetString("relation_type", "") + if owner == "" || srcID == 0 || dstID == 0 || relType == "" { + return memErrorf("invalid_request", "owner_id, src_id, dst_id, relation_type required"), nil + } + jobID, errR := r.authorizeForOwner(ctx, owner) + if errR != nil { + return errR, nil + } + for _, id := range []int64{srcID, dstID} { + ok, sameOwner, _ := r.messageBelongsTo(ctx, id, owner) + if !ok { + return memErrorf("source_not_found", "message %d not found", id), nil + } + if !sameOwner { + return memErrorf("not_same_owner", "message %d belongs to a different owner", id), nil + } + } + if r.deps.Links == nil { + return memErrorf("internal", "link store not configured"), nil + } + var meta map[string]any + if mraw := req.GetString("metadata", ""); mraw != "" { + _ = json.Unmarshal([]byte(mraw), &meta) + } + linkID, err := r.deps.Links.Add(ctx, srcID, dstID, relType, owner, "agent:dream:"+owner, meta) + if err != nil { + if errors.Is(err, messaging.ErrLinkTypeReserved) { + return memErrorf("relation_type_reserved", "type %q is reserved", relType), nil + } + return memErrorf("internal", "add link: %s", err), nil + } + r.recordAction(ctx, jobID, "memory_add_link", dstID, map[string]any{ + "src_id": srcID, "dst_id": dstID, "relation_type": relType, + }) + return resultJSON(map[string]any{"link_id": linkID}) +} + +// --- helpers --- + +func parseInt64CSV(s string) []int64 { + if s == "" { + return nil + } + parts := strings.Split(s, ",") + out := make([]int64, 0, len(parts)) + for _, p := range parts { + p = strings.TrimSpace(p) + if p == "" { + continue + } + var v int64 + if _, err := fmt.Sscanf(p, "%d", &v); err == nil && v > 0 { + out = append(out, v) + } + } + return out +} + +// messageBelongsTo reports (exists, ownerMatches, dbErr). +func (r *MemoryToolRegistrar) messageBelongsTo(ctx context.Context, msgID int64, ownerID string) (bool, bool, error) { + var fromAgent string + err := r.deps.DB.QueryRowContext(ctx, + `SELECT from_agent FROM messages WHERE id = ?`, msgID, + ).Scan(&fromAgent) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return false, false, nil + } + return false, false, err + } + owner, err := agents.OwnerFor(ctx, r.deps.DB, fromAgent) + if err != nil { + return true, false, nil + } + return true, owner == ownerID, nil +} + +// pickReflectionChannel picks a destination channel for a reflection. +// Preference: `reflections-` if any such channel exists, else +// `open-brain`. +func (r *MemoryToolRegistrar) pickReflectionChannel(ctx context.Context, ownerID string) (int64, string, error) { + // Try reflections-* the owner has authored to (best heuristic). + var ( + id int64 + name string + ) + err := r.deps.DB.QueryRowContext(ctx, + `SELECT id, name FROM channels WHERE name LIKE 'reflections-%' ORDER BY id ASC LIMIT 1`, + ).Scan(&id, &name) + if err == nil { + return id, name, nil + } + if !errors.Is(err, sql.ErrNoRows) { + return 0, "", err + } + // Fallback: open-brain. + err = r.deps.DB.QueryRowContext(ctx, + `SELECT id, name FROM channels WHERE name = 'open-brain' LIMIT 1`, + ).Scan(&id, &name) + if err != nil { + if errors.Is(err, sql.ErrNoRows) { + return 0, "", nil + } + return 0, "", err + } + return id, name, nil +} diff --git a/internal/mcp/memory_tools_test.go b/internal/mcp/memory_tools_test.go new file mode 100644 index 0000000..9569683 --- /dev/null +++ b/internal/mcp/memory_tools_test.go @@ -0,0 +1,375 @@ +package mcp + +import ( + "context" + "database/sql" + "encoding/json" + "strings" + "testing" + "time" + + mcplib "github.com/mark3labs/mcp-go/mcp" + + "github.com/synapbus/synapbus/internal/messaging" +) + +// memToolHarness bundles deps + helpers for the memory tools tests. +type memToolHarness struct { + db *sql.DB + reg *MemoryToolRegistrar + tokens *messaging.DispatchTokenStore + jobs *messaging.JobsStore + links *messaging.LinkStore + pins *messaging.PinStore + core *messaging.CoreMemoryStore + jobID int64 + tokenStr string + ownerID string +} + +func newMemToolHarness(t *testing.T) *memToolHarness { + t.Helper() + db := newTestDB(t) + + // Seed two owners + their agents (used for owner mismatch tests). + _, _ = db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (2, 'otheruser', 'hash', 'Other')`) + _, _ = db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, owner_id, api_key_hash, status) VALUES ('a-h1', 'a-h1', 'ai', 1, 'h1', 'active')`) + _, _ = db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, owner_id, api_key_hash, status) VALUES ('a-h2', 'a-h2', 'ai', 2, 'h2', 'active')`) + + tokens := messaging.NewDispatchTokenStore(db) + jobs := messaging.NewJobsStore(db) + links := messaging.NewLinkStore(db) + pins := messaging.NewPinStore(db) + core := messaging.NewCoreMemoryStore(db, 64) + + jobID, err := jobs.Create(context.Background(), "1", "reflection", "manual:test") + if err != nil { + t.Fatalf("create job: %v", err) + } + tok, _, err := tokens.Issue(context.Background(), "1", jobID) + if err != nil { + t.Fatalf("issue token: %v", err) + } + + reg := NewMemoryToolRegistrar(MemoryToolDeps{ + DB: db, + Core: core, + Links: links, + Pins: pins, + Jobs: jobs, + Tokens: tokens, + }) + return &memToolHarness{ + db: db, reg: reg, tokens: tokens, jobs: jobs, links: links, pins: pins, core: core, + jobID: jobID, tokenStr: tok, ownerID: "1", + } +} + +func (h *memToolHarness) ctxWithToken(tok string) context.Context { + return WithDispatchToken(context.Background(), tok) +} + +// seedMemoryChannel inserts an open-brain channel and one message +// belonging to `agentName`. +func (h *memToolHarness) seedMessage(t *testing.T, agentName, body string) int64 { + t.Helper() + _, _ = h.db.Exec(`INSERT OR IGNORE INTO channels (id, name, description, type, created_by) VALUES (1, 'open-brain', '', 'standard', 'system')`) + res, err := h.db.Exec( + `INSERT INTO conversations (created_by, channel_id) VALUES (?, 1)`, agentName) + if err != nil { + t.Fatalf("seed conv: %v", err) + } + convID, _ := res.LastInsertId() + res, err = h.db.Exec( + `INSERT INTO messages (conversation_id, from_agent, channel_id, body, priority, status, metadata) + VALUES (?, ?, 1, ?, 5, 'pending', '{}')`, + convID, agentName, body, + ) + if err != nil { + t.Fatalf("seed message: %v", err) + } + id, _ := res.LastInsertId() + return id +} + +func callRequest(args map[string]any) mcplib.CallToolRequest { + return mcplib.CallToolRequest{ + Params: mcplib.CallToolParams{Arguments: args}, + } +} + +func resultText(t *testing.T, res *mcplib.CallToolResult) string { + t.Helper() + if res == nil || len(res.Content) == 0 { + t.Fatal("nil/empty result") + } + tc, ok := res.Content[0].(mcplib.TextContent) + if !ok { + t.Fatalf("not TextContent: %T", res.Content[0]) + } + return tc.Text +} + +func resultIsError(res *mcplib.CallToolResult) bool { + return res != nil && res.IsError +} + +// --- Token error matrix --- + +func TestMemoryTools_DispatchTokenMissing(t *testing.T) { + h := newMemToolHarness(t) + // No token in context. + res, _ := h.reg.handleAddLink(context.Background(), callRequest(map[string]any{ + "owner_id": "1", "src_id": 1.0, "dst_id": 2.0, "relation_type": "refines", + })) + if !resultIsError(res) { + t.Fatal("expected error result") + } + if !strings.Contains(resultText(t, res), "dispatch_token_missing") { + t.Errorf("expected dispatch_token_missing, got %q", resultText(t, res)) + } +} + +func TestMemoryTools_DispatchTokenExpired(t *testing.T) { + h := newMemToolHarness(t) + // Forcibly expire. + if _, err := h.db.Exec( + `UPDATE memory_dispatch_tokens SET expires_at = ? WHERE token = ?`, + time.Now().Add(-1*time.Minute).UTC(), h.tokenStr, + ); err != nil { + t.Fatalf("expire: %v", err) + } + res, _ := h.reg.handleAddLink(h.ctxWithToken(h.tokenStr), callRequest(map[string]any{ + "owner_id": "1", "src_id": 1.0, "dst_id": 2.0, "relation_type": "refines", + })) + if !resultIsError(res) || !strings.Contains(resultText(t, res), "dispatch_token_expired") { + t.Errorf("expected dispatch_token_expired, got %q", resultText(t, res)) + } +} + +func TestMemoryTools_DispatchTokenOwnerMismatch(t *testing.T) { + h := newMemToolHarness(t) + res, _ := h.reg.handleAddLink(h.ctxWithToken(h.tokenStr), callRequest(map[string]any{ + "owner_id": "2", "src_id": 1.0, "dst_id": 2.0, "relation_type": "refines", + })) + if !resultIsError(res) || !strings.Contains(resultText(t, res), "dispatch_token_owner_mismatch") { + t.Errorf("expected dispatch_token_owner_mismatch, got %q", resultText(t, res)) + } +} + +// --- memory_add_link --- + +func TestMemoryAddLink_HappyPath(t *testing.T) { + h := newMemToolHarness(t) + a := h.seedMessage(t, "a-h1", "fact A") + b := h.seedMessage(t, "a-h1", "fact B") + + res, err := h.reg.handleAddLink(h.ctxWithToken(h.tokenStr), callRequest(map[string]any{ + "owner_id": "1", + "src_id": float64(a), + "dst_id": float64(b), + "relation_type": "refines", + })) + if err != nil { + t.Fatalf("handleAddLink: %v", err) + } + if resultIsError(res) { + t.Fatalf("unexpected error: %s", resultText(t, res)) + } + var body map[string]any + if err := json.Unmarshal([]byte(resultText(t, res)), &body); err != nil { + t.Fatalf("parse body: %v", err) + } + if _, ok := body["link_id"]; !ok { + t.Errorf("expected link_id in response: %v", body) + } + + // And actions should have been appended. + job, err := h.jobs.Get(context.Background(), h.jobID) + if err != nil { + t.Fatalf("Get job: %v", err) + } + if len(job.Actions) != 1 || job.Actions[0]["tool"] != "memory_add_link" { + t.Errorf("expected one action for memory_add_link, got %v", job.Actions) + } +} + +func TestMemoryAddLink_RelationTypeReserved(t *testing.T) { + h := newMemToolHarness(t) + a := h.seedMessage(t, "a-h1", "x") + b := h.seedMessage(t, "a-h1", "y") + res, _ := h.reg.handleAddLink(h.ctxWithToken(h.tokenStr), callRequest(map[string]any{ + "owner_id": "1", "src_id": float64(a), "dst_id": float64(b), + "relation_type": "mention", + })) + if !resultIsError(res) || !strings.Contains(resultText(t, res), "relation_type_reserved") { + t.Errorf("expected relation_type_reserved, got %q", resultText(t, res)) + } +} + +func TestMemoryAddLink_NotSameOwner(t *testing.T) { + h := newMemToolHarness(t) + src := h.seedMessage(t, "a-h1", "x") + other := h.seedMessage(t, "a-h2", "y") + res, _ := h.reg.handleAddLink(h.ctxWithToken(h.tokenStr), callRequest(map[string]any{ + "owner_id": "1", "src_id": float64(src), "dst_id": float64(other), + "relation_type": "refines", + })) + if !resultIsError(res) || !strings.Contains(resultText(t, res), "not_same_owner") { + t.Errorf("expected not_same_owner, got %q", resultText(t, res)) + } +} + +func TestMemoryAddLink_SourceNotFound(t *testing.T) { + h := newMemToolHarness(t) + res, _ := h.reg.handleAddLink(h.ctxWithToken(h.tokenStr), callRequest(map[string]any{ + "owner_id": "1", "src_id": 9999.0, "dst_id": 8888.0, "relation_type": "refines", + })) + if !resultIsError(res) || !strings.Contains(resultText(t, res), "source_not_found") { + t.Errorf("expected source_not_found, got %q", resultText(t, res)) + } +} + +// --- memory_rewrite_core --- + +func TestMemoryRewriteCore_HappyPath(t *testing.T) { + h := newMemToolHarness(t) + res, _ := h.reg.handleRewriteCore(h.ctxWithToken(h.tokenStr), callRequest(map[string]any{ + "owner_id": "1", "agent_name": "a-h1", "blob": "I am a-h1.", + })) + if resultIsError(res) { + t.Fatalf("unexpected error: %s", resultText(t, res)) + } + blob, _, ok, _ := h.core.Get(context.Background(), "1", "a-h1") + if !ok || blob != "I am a-h1." { + t.Errorf("blob not stored: ok=%v blob=%q", ok, blob) + } +} + +func TestMemoryRewriteCore_TooLarge(t *testing.T) { + h := newMemToolHarness(t) + big := strings.Repeat("x", 65) + res, _ := h.reg.handleRewriteCore(h.ctxWithToken(h.tokenStr), callRequest(map[string]any{ + "owner_id": "1", "agent_name": "a-h1", "blob": big, + })) + if !resultIsError(res) || !strings.Contains(resultText(t, res), "core_memory_too_large") { + t.Errorf("expected core_memory_too_large, got %q", resultText(t, res)) + } +} + +func TestMemoryRewriteCore_AgentNotSameOwner(t *testing.T) { + h := newMemToolHarness(t) + res, _ := h.reg.handleRewriteCore(h.ctxWithToken(h.tokenStr), callRequest(map[string]any{ + "owner_id": "1", "agent_name": "a-h2", "blob": "trying to overwrite", + })) + if !resultIsError(res) || !strings.Contains(resultText(t, res), "not_same_owner") { + t.Errorf("expected not_same_owner, got %q", resultText(t, res)) + } +} + +// --- memory_mark_duplicate --- + +func TestMemoryMarkDuplicate_HappyPath(t *testing.T) { + h := newMemToolHarness(t) + a := h.seedMessage(t, "a-h1", "fact A") + b := h.seedMessage(t, "a-h1", "fact A shorter") + res, _ := h.reg.handleMarkDuplicate(h.ctxWithToken(h.tokenStr), callRequest(map[string]any{ + "owner_id": "1", "a_id": float64(a), "b_id": float64(b), "keep_id": float64(a), + "reason": "shorter", + })) + if resultIsError(res) { + t.Fatalf("unexpected error: %s", resultText(t, res)) + } + // Audit row appended on job. + job, _ := h.jobs.Get(context.Background(), h.jobID) + if len(job.Actions) != 1 || job.Actions[0]["tool"] != "memory_mark_duplicate" { + t.Errorf("expected one mark_duplicate action: %v", job.Actions) + } +} + +// --- memory_supersede --- + +func TestMemorySupersede_HappyPath(t *testing.T) { + h := newMemToolHarness(t) + a := h.seedMessage(t, "a-h1", "old fact") + b := h.seedMessage(t, "a-h1", "new fact") + res, _ := h.reg.handleSupersede(h.ctxWithToken(h.tokenStr), callRequest(map[string]any{ + "owner_id": "1", "a_id": float64(a), "b_id": float64(b), "reason": "newer", + })) + if resultIsError(res) { + t.Fatalf("unexpected error: %s", resultText(t, res)) + } +} + +// --- memory_list_unprocessed --- + +func TestMemoryListUnprocessed_ReturnsOwnerScopedMessages(t *testing.T) { + h := newMemToolHarness(t) + h.seedMessage(t, "a-h1", "memory 1") + h.seedMessage(t, "a-h1", "memory 2") + h.seedMessage(t, "a-h2", "other owner's memory") + + res, _ := h.reg.handleListUnprocessed(h.ctxWithToken(h.tokenStr), callRequest(map[string]any{ + "owner_id": "1", + })) + if resultIsError(res) { + t.Fatalf("unexpected error: %s", resultText(t, res)) + } + var body map[string]any + _ = json.Unmarshal([]byte(resultText(t, res)), &body) + mems, _ := body["memories"].([]any) + if len(mems) != 2 { + t.Errorf("expected 2 owner-scoped memories, got %d (full body: %v)", len(mems), body) + } +} + +// --- memory_write_reflection --- + +func TestMemoryWriteReflection_HappyPath(t *testing.T) { + h := newMemToolHarness(t) + a := h.seedMessage(t, "a-h1", "source 1") + b := h.seedMessage(t, "a-h1", "source 2") + + args := map[string]any{ + "owner_id": "1", + "body": "Across these I notice...", + "source_message_ids": "" + intCSV(a, b), + } + res, _ := h.reg.handleWriteReflection(h.ctxWithToken(h.tokenStr), callRequest(args)) + if resultIsError(res) { + t.Fatalf("unexpected error: %s", resultText(t, res)) + } + var body map[string]any + _ = json.Unmarshal([]byte(resultText(t, res)), &body) + if _, ok := body["memory_id"]; !ok { + t.Errorf("expected memory_id: %v", body) + } + if v, _ := body["links_created"].(float64); int(v) != 2 { + t.Errorf("expected links_created=2, got %v", body["links_created"]) + } +} + +func intCSV(ids ...int64) string { + var b strings.Builder + for i, id := range ids { + if i > 0 { + b.WriteString(",") + } + b.WriteString(itoa(id)) + } + return b.String() +} + +func itoa(v int64) string { + if v == 0 { + return "0" + } + var buf [20]byte + i := len(buf) + for v > 0 { + i-- + buf[i] = byte('0' + v%10) + v /= 10 + } + return string(buf[i:]) +} diff --git a/internal/mcp/server.go b/internal/mcp/server.go index b71fb30..5131a1c 100644 --- a/internal/mcp/server.go +++ b/internal/mcp/server.go @@ -225,6 +225,22 @@ func (s *MCPServer) SetInjection(cfg messaging.MemoryConfig, store *messaging.Me s.hybridRegistrar.RegisterAllOnServer(s.mcpServer) } +// SetDream wires the six memory_* MCP tools (feature 020 — US3). +// Only registers when cfg.DreamEnabled is true; otherwise this is a +// no-op so the tool surface remains identical to the pre-feature +// shape. Must be called before the server starts serving traffic. +func (s *MCPServer) SetDream(deps MemoryToolDeps) { + if s == nil || s.mcpServer == nil { + return + } + if !deps.MemConfig.DreamEnabled { + s.logger.Info("dream tools not registered (SYNAPBUS_DREAM_ENABLED=0)") + return + } + reg := NewMemoryToolRegistrar(deps) + reg.RegisterAllOnServer(s.mcpServer) +} + // WireGoalsTools registers the spec-018 tool surface (create_goal, // propose_task_tree, claim_task, request_resource, list_resources, // complete_goal) on the MCP server. Must be called after NewMCPServer. diff --git a/internal/messaging/auto_links.go b/internal/messaging/auto_links.go new file mode 100644 index 0000000..9ed6514 --- /dev/null +++ b/internal/messaging/auto_links.go @@ -0,0 +1,205 @@ +// Auto-link emitter for feature 020 — when a new memory-eligible +// message is inserted, derive `mention`, `reply_to`, and +// `channel_cooccurrence` links automatically and write them with +// `created_by = "auto:"`. The dream-agent's `memory_add_link` +// tool rejects these auto-types (see memory_links.go) so this is the +// only path that creates them. +// +// Wiring: register an AutoLinkListener with +// `MessagingService.AddMessageListener`. The listener fires after +// every successful message insert. Failures are logged but never +// propagated — auto-links are best-effort. +package messaging + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "log/slog" + "regexp" +) + +// mentionPattern matches `@agent-name` in the body. Hyphens and digits +// allowed; case-insensitive; bounded by non-word characters or string +// edges. Matches the existing mentions.go convention. +var autoMentionPattern = regexp.MustCompile(`@([A-Za-z][A-Za-z0-9_\-]{1,63})`) + +// AutoLinkListener implements MessageListener and writes the three +// auto-link types per OnMessageSent invocation. +type AutoLinkListener struct { + db *sql.DB + links *LinkStore + logger *slog.Logger +} + +// NewAutoLinkListener returns a listener over db + links. +func NewAutoLinkListener(db *sql.DB, links *LinkStore) *AutoLinkListener { + return &AutoLinkListener{ + db: db, + links: links, + logger: slog.Default().With("component", "auto-links"), + } +} + +// OnMessageSent implements MessageListener. Runs the three rules in +// order. Each rule independently best-effort. +func (l *AutoLinkListener) OnMessageSent(ctx context.Context, msg *Message) { + if l == nil || l.db == nil || l.links == nil || msg == nil || msg.ID == 0 { + return + } + // Only run for memory-channel messages — auto-links on every DM + // would pollute the link table. + if msg.ChannelID == nil { + return + } + channelName, err := channelNameByID(ctx, l.db, *msg.ChannelID) + if err != nil || !matchesMemoryChannelName(channelName) { + return + } + ownerID := resolveOwnerString(ctx, l.db, msg.FromAgent) + if ownerID == "" { + return + } + + // 1. reply_to: simple metadata or msg.ReplyTo column. + if msg.ReplyTo != nil && *msg.ReplyTo != 0 { + if _, err := l.links.Add(ctx, msg.ID, *msg.ReplyTo, "reply_to", ownerID, "auto:reply_to", nil); err != nil { + l.logger.Debug("auto reply_to failed", "msg", msg.ID, "error", err) + } + } else if reply := extractReplyToFromMetadata(msg.Metadata); reply != 0 { + if _, err := l.links.Add(ctx, msg.ID, reply, "reply_to", ownerID, "auto:reply_to", nil); err != nil { + l.logger.Debug("auto reply_to (meta) failed", "msg", msg.ID, "error", err) + } + } + + // 2. mention: @agent-name → latest message from that agent in the + // same channel. + for _, m := range autoMentionPattern.FindAllStringSubmatch(msg.Body, -1) { + if len(m) < 2 { + continue + } + target := m[1] + if target == msg.FromAgent { + continue + } + dst := mostRecentMessageFromAgentInChannel(ctx, l.db, *msg.ChannelID, target, msg.ID) + if dst == 0 { + continue + } + if _, err := l.links.Add(ctx, msg.ID, dst, "mention", ownerID, "auto:mention", nil); err != nil { + l.logger.Debug("auto mention failed", "msg", msg.ID, "to", target, "error", err) + } + } + + // 3. channel_cooccurrence: previous-most-recent message in the + // same channel. + prev := mostRecentMessageInChannel(ctx, l.db, *msg.ChannelID, msg.ID) + if prev != 0 { + if _, err := l.links.Add(ctx, msg.ID, prev, "channel_cooccurrence", ownerID, "auto:channel_cooccurrence", nil); err != nil { + l.logger.Debug("auto cooccurrence failed", "msg", msg.ID, "error", err) + } + } +} + +// channelNameByID returns the name of the channel with the given id. +func channelNameByID(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 +} + +// resolveOwnerString returns the agents.owner_id of fromAgent as a +// string, or "" on any error. +func resolveOwnerString(ctx context.Context, db *sql.DB, fromAgent string) string { + var ownerID int64 + err := db.QueryRowContext(ctx, + `SELECT owner_id FROM agents WHERE name = ?`, fromAgent, + ).Scan(&ownerID) + if err != nil || ownerID == 0 { + return "" + } + // Use the same formatting as agents.OwnerFor so cross-package + // comparisons stay byte-for-byte consistent. + return itoaInt64(ownerID) +} + +// itoaInt64 mirrors strconv.FormatInt(v,10) without importing strconv. +func itoaInt64(v int64) string { + if v == 0 { + return "0" + } + var buf [20]byte + i := len(buf) + neg := v < 0 + if neg { + v = -v + } + for v > 0 { + i-- + buf[i] = byte('0' + v%10) + v /= 10 + } + if neg { + i-- + buf[i] = '-' + } + return string(buf[i:]) +} + +// extractReplyToFromMetadata returns the reply_to_message_id field if +// the metadata is a JSON object with that key. +func extractReplyToFromMetadata(meta json.RawMessage) int64 { + if len(meta) == 0 { + return 0 + } + var m map[string]any + if err := json.Unmarshal(meta, &m); err != nil { + return 0 + } + if v, ok := m["reply_to_message_id"]; ok { + switch t := v.(type) { + case float64: + return int64(t) + case int64: + return t + case int: + return int64(t) + } + } + return 0 +} + +func mostRecentMessageFromAgentInChannel(ctx context.Context, db *sql.DB, channelID int64, agent string, excludeID int64) int64 { + var id int64 + err := db.QueryRowContext(ctx, + `SELECT id FROM messages + WHERE channel_id = ? AND from_agent = ? AND id != ? + ORDER BY id DESC LIMIT 1`, + channelID, agent, excludeID, + ).Scan(&id) + if err != nil { + if !errors.Is(err, sql.ErrNoRows) { + // best-effort: caller logs at debug + } + return 0 + } + return id +} + +func mostRecentMessageInChannel(ctx context.Context, db *sql.DB, channelID int64, excludeID int64) int64 { + var id int64 + err := db.QueryRowContext(ctx, + `SELECT id FROM messages + WHERE channel_id = ? AND id < ? + ORDER BY id DESC LIMIT 1`, + channelID, excludeID, + ).Scan(&id) + if err != nil { + return 0 + } + return id +} diff --git a/internal/messaging/consolidation_jobs.go b/internal/messaging/consolidation_jobs.go new file mode 100644 index 0000000..4451ae9 --- /dev/null +++ b/internal/messaging/consolidation_jobs.go @@ -0,0 +1,352 @@ +// Consolidation-jobs store for feature 020 — wraps the +// `memory_consolidation_jobs` table (data-model.md). Each row is one +// dream-worker dispatch: state machine `pending → dispatched → running +// → {succeeded|partial|failed|expired}`. The partial-unique index +// `idx_consolidation_in_flight(owner_id, job_type) WHERE status IN +// ('pending','dispatched','running')` guarantees at most one in-flight +// row per (owner, job_type). Create() surfaces conflicts as +// ErrJobAlreadyInFlight so the worker can skip the dispatch cleanly. +package messaging + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "strings" + "time" +) + +// Memory consolidation job types. +const ( + JobTypeReflection = "reflection" + JobTypeCoreRewrite = "core_rewrite" + JobTypeDedupContradiction = "dedup_contradiction" + JobTypeLinkGen = "link_gen" +) + +// Memory consolidation job statuses. +const ( + JobStatusPending = "pending" + JobStatusDispatched = "dispatched" + JobStatusRunning = "running" + JobStatusSucceeded = "succeeded" + JobStatusPartial = "partial" + JobStatusFailed = "failed" + JobStatusExpired = "expired" +) + +// ErrJobAlreadyInFlight is returned by JobsStore.Create when the +// partial-unique index trips because another job of the same type is +// already pending / dispatched / running for the same owner. +var ErrJobAlreadyInFlight = errors.New("consolidation job already in flight for (owner, job_type)") + +// Job is one row in `memory_consolidation_jobs`. +type Job struct { + ID int64 `json:"id"` + OwnerID string `json:"owner_id"` + JobType string `json:"job_type"` + Status string `json:"status"` + TriggerReason string `json:"trigger_reason"` + DispatchToken string `json:"dispatch_token,omitempty"` + HarnessRunID string `json:"harness_run_id,omitempty"` + Actions []map[string]any `json:"actions"` + Summary string `json:"summary,omitempty"` + Error string `json:"error,omitempty"` + LeaseUntil *time.Time `json:"lease_until,omitempty"` + StartedAt *time.Time `json:"started_at,omitempty"` + FinishedAt *time.Time `json:"finished_at,omitempty"` + CreatedAt time.Time `json:"created_at"` +} + +// JobsStore wraps the `memory_consolidation_jobs` table. +type JobsStore struct { + db *sql.DB +} + +// NewJobsStore returns a store rooted at db. +func NewJobsStore(db *sql.DB) *JobsStore { + return &JobsStore{db: db} +} + +// Create inserts a `pending` row for the (owner, jobType) pair. If +// another job of the same type is already in flight, returns +// ErrJobAlreadyInFlight (mapped from the partial-unique-index conflict). +func (s *JobsStore) Create(ctx context.Context, ownerID, jobType, triggerReason string) (int64, error) { + if s == nil || s.db == nil { + return 0, fmt.Errorf("jobs store: nil store") + } + if ownerID == "" { + return 0, fmt.Errorf("jobs store: empty owner_id") + } + if jobType == "" { + return 0, fmt.Errorf("jobs store: empty job_type") + } + res, err := s.db.ExecContext(ctx, + `INSERT INTO memory_consolidation_jobs + (owner_id, job_type, status, trigger_reason) + VALUES (?, ?, 'pending', ?)`, + ownerID, jobType, triggerReason, + ) + if err != nil { + // modernc.org/sqlite surfaces unique-constraint conflicts via + // error strings; the partial-unique index is the only UNIQUE + // constraint that can fire here for INSERT. + if isUniqueConstraint(err) { + return 0, ErrJobAlreadyInFlight + } + return 0, fmt.Errorf("jobs store: insert: %w", err) + } + id, _ := res.LastInsertId() + return id, nil +} + +// Dispatch flips a pending row to `dispatched` and stamps the harness +// run id and dispatch token. Returns an error if the row is not in +// `pending` state. +func (s *JobsStore) Dispatch(ctx context.Context, jobID int64, harnessRunID, token string) error { + if s == nil || s.db == nil { + return fmt.Errorf("jobs store: nil store") + } + res, err := s.db.ExecContext(ctx, + `UPDATE memory_consolidation_jobs + SET status = 'dispatched', + harness_run_id = ?, + dispatch_token = ? + WHERE id = ? AND status = 'pending'`, + harnessRunID, token, jobID, + ) + if err != nil { + return fmt.Errorf("jobs store: dispatch: %w", err) + } + rows, _ := res.RowsAffected() + if rows == 0 { + return fmt.Errorf("jobs store: dispatch: job %d not in pending state", jobID) + } + return nil +} + +// Lease flips `dispatched` → `running` and sets lease_until + started_at. +// Called by the worker once the harness has confirmed the run started. +func (s *JobsStore) Lease(ctx context.Context, jobID int64, until time.Time) error { + if s == nil || s.db == nil { + return fmt.Errorf("jobs store: nil store") + } + res, err := s.db.ExecContext(ctx, + `UPDATE memory_consolidation_jobs + SET status = 'running', + lease_until = ?, + started_at = CURRENT_TIMESTAMP + WHERE id = ? AND status IN ('dispatched', 'pending')`, + until.UTC(), jobID, + ) + if err != nil { + return fmt.Errorf("jobs store: lease: %w", err) + } + rows, _ := res.RowsAffected() + if rows == 0 { + return fmt.Errorf("jobs store: lease: job %d not in dispatched/pending state", jobID) + } + return nil +} + +// AppendAction reads the current `actions` JSON array, appends `action`, +// and writes it back. Wrapped in a single transaction so concurrent +// MCP tool calls within one job serialize cleanly. +func (s *JobsStore) AppendAction(ctx context.Context, jobID int64, action map[string]any) error { + if s == nil || s.db == nil { + return fmt.Errorf("jobs store: nil store") + } + if action == nil { + return fmt.Errorf("jobs store: nil action") + } + + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("jobs store: begin tx: %w", err) + } + defer func() { _ = tx.Rollback() }() + + var current string + if err := tx.QueryRowContext(ctx, + `SELECT actions FROM memory_consolidation_jobs WHERE id = ?`, jobID, + ).Scan(¤t); err != nil { + return fmt.Errorf("jobs store: read actions: %w", err) + } + + var arr []map[string]any + if current == "" || current == "null" { + arr = []map[string]any{} + } else if err := json.Unmarshal([]byte(current), &arr); err != nil { + // Corrupt JSON — start fresh rather than fail forever. + arr = []map[string]any{} + } + arr = append(arr, action) + b, err := json.Marshal(arr) + if err != nil { + return fmt.Errorf("jobs store: marshal actions: %w", err) + } + + if _, err := tx.ExecContext(ctx, + `UPDATE memory_consolidation_jobs SET actions = ? WHERE id = ?`, + string(b), jobID, + ); err != nil { + return fmt.Errorf("jobs store: write actions: %w", err) + } + if err := tx.Commit(); err != nil { + return fmt.Errorf("jobs store: commit: %w", err) + } + return nil +} + +// Complete sets `status`, `summary`, `error`, `finished_at` and clears +// `lease_until`. Idempotent — repeated calls keep the first finished_at. +func (s *JobsStore) Complete(ctx context.Context, jobID int64, status, summary, errMsg string) error { + if s == nil || s.db == nil { + return fmt.Errorf("jobs store: nil store") + } + switch status { + case JobStatusSucceeded, JobStatusPartial, JobStatusFailed, JobStatusExpired: + // ok + default: + return fmt.Errorf("jobs store: invalid completion status %q", status) + } + _, err := s.db.ExecContext(ctx, + `UPDATE memory_consolidation_jobs + SET status = ?, + summary = COALESCE(NULLIF(?, ''), summary), + error = COALESCE(NULLIF(?, ''), error), + finished_at = COALESCE(finished_at, CURRENT_TIMESTAMP), + lease_until = NULL + WHERE id = ?`, + status, summary, errMsg, jobID, + ) + if err != nil { + return fmt.Errorf("jobs store: complete: %w", err) + } + return nil +} + +// Get returns the row for jobID, or (nil, sql.ErrNoRows). +func (s *JobsStore) Get(ctx context.Context, jobID int64) (*Job, error) { + row := s.db.QueryRowContext(ctx, + `SELECT id, owner_id, job_type, status, trigger_reason, + COALESCE(dispatch_token, ''), + COALESCE(harness_run_id, ''), + actions, + COALESCE(summary, ''), + COALESCE(error, ''), + lease_until, started_at, finished_at, created_at + FROM memory_consolidation_jobs WHERE id = ?`, jobID, + ) + return scanJob(row.Scan) +} + +// ListRecent returns the most-recent jobs for the given owner. +func (s *JobsStore) ListRecent(ctx context.Context, ownerID string, limit int) ([]Job, error) { + if limit <= 0 { + limit = 50 + } + rows, err := s.db.QueryContext(ctx, + `SELECT id, owner_id, job_type, status, trigger_reason, + COALESCE(dispatch_token, ''), + COALESCE(harness_run_id, ''), + actions, + COALESCE(summary, ''), + COALESCE(error, ''), + lease_until, started_at, finished_at, created_at + FROM memory_consolidation_jobs + WHERE owner_id = ? + ORDER BY created_at DESC, id DESC + LIMIT ?`, ownerID, limit, + ) + if err != nil { + return nil, fmt.Errorf("jobs store: list recent: %w", err) + } + defer rows.Close() + var out []Job + for rows.Next() { + j, err := scanJob(rows.Scan) + if err != nil { + return nil, err + } + out = append(out, *j) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("jobs store: iterate: %w", err) + } + return out, nil +} + +// ActiveJob returns the in-flight job for (owner, jobType), if any. nil +// when no row matches. +func (s *JobsStore) ActiveJob(ctx context.Context, ownerID, jobType string) (*Job, error) { + row := s.db.QueryRowContext(ctx, + `SELECT id, owner_id, job_type, status, trigger_reason, + COALESCE(dispatch_token, ''), + COALESCE(harness_run_id, ''), + actions, + COALESCE(summary, ''), + COALESCE(error, ''), + lease_until, started_at, finished_at, created_at + FROM memory_consolidation_jobs + WHERE owner_id = ? AND job_type = ? + AND status IN ('pending', 'dispatched', 'running') + ORDER BY id DESC LIMIT 1`, + ownerID, jobType, + ) + j, err := scanJob(row.Scan) + if errors.Is(err, sql.ErrNoRows) { + return nil, nil + } + return j, err +} + +type scanFn func(dest ...any) error + +func scanJob(scan scanFn) (*Job, error) { + var ( + j Job + actions string + leaseUntil sql.NullTime + startedAt sql.NullTime + finishedAt sql.NullTime + ) + err := scan( + &j.ID, &j.OwnerID, &j.JobType, &j.Status, &j.TriggerReason, + &j.DispatchToken, &j.HarnessRunID, &actions, + &j.Summary, &j.Error, + &leaseUntil, &startedAt, &finishedAt, &j.CreatedAt, + ) + if err != nil { + return nil, err + } + if actions != "" && actions != "null" { + _ = json.Unmarshal([]byte(actions), &j.Actions) + } + if leaseUntil.Valid { + t := leaseUntil.Time + j.LeaseUntil = &t + } + if startedAt.Valid { + t := startedAt.Time + j.StartedAt = &t + } + if finishedAt.Valid { + t := finishedAt.Time + j.FinishedAt = &t + } + return &j, nil +} + +func isUniqueConstraint(err error) bool { + if err == nil { + return false + } + // modernc.org/sqlite returns errors whose string contains + // "constraint failed: UNIQUE" or "SQLITE_CONSTRAINT_UNIQUE". Match + // loosely so we don't depend on a specific build. + msg := strings.ToUpper(err.Error()) + return strings.Contains(msg, "UNIQUE") +} diff --git a/internal/messaging/consolidator.go b/internal/messaging/consolidator.go new file mode 100644 index 0000000..fe04555 --- /dev/null +++ b/internal/messaging/consolidator.go @@ -0,0 +1,485 @@ +// Dream-worker (consolidator) for feature 020 — periodically scans +// memory channels per-owner, evaluates trigger watermarks and the +// daily deep-pass schedule, and dispatches consolidation jobs to a +// Claude Code agent via the harness. +// +// Importantly: the worker NEVER sends a system DM. Per +// feedback_system_dm_no_trigger.md, system DMs would trigger reactive +// runs which would cascade through the stalemate worker. The harness +// dispatch path is the contractual non-DM route — see R1. +// +// To avoid an import cycle (harness imports messaging), the worker +// accepts a minimal HarnessDispatcher interface that mirrors the +// fragment of harness.Registry it needs. The cmd/synapbus wiring at +// startup adapts harness.Registry to this interface. +package messaging + +import ( + "context" + "database/sql" + "errors" + "fmt" + "log/slog" + "sync" + "time" + + "github.com/google/uuid" +) + +// DreamAgent is a minimal record passed to the harness dispatcher. +// We can't import internal/agents here (cycle: agents → messaging +// already) so the worker uses an interface-typed value for the agent +// record and the cmd/synapbus adapter unboxes it. This keeps the +// dependency direction (harness → messaging, agents → messaging) +// intact. +type DreamAgent interface { + // AgentName returns the agent's stable SynapBus name. + AgentName() string +} + +// DreamAgentNamed is a tiny convenience wrapper for tests and the +// admin path that need a DreamAgent with just a name. +type DreamAgentNamed struct{ Name string } + +// AgentName implements DreamAgent. +func (a DreamAgentNamed) AgentName() string { return a.Name } + +// HarnessDispatcher is the minimal slice of harness.Registry the worker +// needs. cmd/synapbus wires a real registry behind this interface. +type HarnessDispatcher interface { + Execute(ctx context.Context, agent DreamAgent, req *HarnessExecRequest) (*HarnessExecResult, error) +} + +// HarnessExecRequest mirrors harness.ExecRequest's Env-bearing fields. +// We don't re-export the full struct to avoid pulling all of harness +// into messaging. The adapter in cmd/synapbus translates 1:1. +type HarnessExecRequest struct { + RunID string + AgentName string + Agent DreamAgent + Env map[string]string + MaxWallClock time.Duration + Body string // populated into ExecRequest.Message if non-empty +} + +// HarnessExecResult mirrors the fields the worker reads from +// harness.ExecResult. +type HarnessExecResult struct { + ExitCode int + Logs string +} + +// AgentLookup resolves an agent record by name. Implemented by +// agents.AgentService — declared as an interface here to avoid a +// hard dependency. +type AgentLookup interface { + GetAgent(ctx context.Context, name string) (DreamAgent, error) +} + +// OwnerLister returns the list of distinct owner_ids that have at +// least one recent memory-channel message worth scanning. +type OwnerLister func(ctx context.Context, db *sql.DB) ([]string, error) + +// ConsolidatorWorker periodically evaluates per-owner triggers and +// dispatches dream-agent runs. +type ConsolidatorWorker struct { + db *sql.DB + jobs *JobsStore + tokens *DispatchTokenStore + harness HarnessDispatcher + agentLook AgentLookup + cfg MemoryConfig + logger *slog.Logger + + // Optional owner enumerator. Defaults to a query over `agents`. + ownerLister OwnerLister + + // Optional injection cleanup hook. When non-nil, the worker calls + // it once per hour. Leave nil if the stalemate worker is already + // handling injection cleanup (the default in main.go). + injections *MemoryInjections + + // sem caps simultaneous Execute calls. + sem chan struct{} + + done chan struct{} + wg sync.WaitGroup +} + +// NewConsolidatorWorker builds a worker. Required: db, jobs, tokens, +// harness, agentLook. The semaphore is sized from cfg.DreamMaxConcurrent. +func NewConsolidatorWorker( + db *sql.DB, + jobs *JobsStore, + tokens *DispatchTokenStore, + harness HarnessDispatcher, + agentLook AgentLookup, + cfg MemoryConfig, +) *ConsolidatorWorker { + if cfg.DreamMaxConcurrent <= 0 { + cfg.DreamMaxConcurrent = 1 + } + return &ConsolidatorWorker{ + db: db, + jobs: jobs, + tokens: tokens, + harness: harness, + agentLook: agentLook, + cfg: cfg, + logger: slog.Default().With("component", "consolidator-worker"), + sem: make(chan struct{}, cfg.DreamMaxConcurrent), + done: make(chan struct{}), + } +} + +// SetInjectionCleanup registers a memory_injections store the worker +// will Cleanup hourly. Pass nil to disable. The default main.go wiring +// leaves the stalemate worker handling injection cleanup and skips +// this. +func (w *ConsolidatorWorker) SetInjectionCleanup(store *MemoryInjections) { + w.injections = store +} + +// SetOwnerLister overrides the default owner enumerator (handy in +// tests). +func (w *ConsolidatorWorker) SetOwnerLister(fn OwnerLister) { + w.ownerLister = fn +} + +// Start launches the ticker goroutine. +func (w *ConsolidatorWorker) Start() { + w.wg.Add(1) + go w.runLoop() +} + +// Stop halts the worker. Idempotent. +func (w *ConsolidatorWorker) Stop() { + select { + case <-w.done: + return + default: + } + close(w.done) + w.wg.Wait() +} + +func (w *ConsolidatorWorker) runLoop() { + defer w.wg.Done() + w.logger.Info("consolidator worker started", + "interval", w.cfg.DreamInterval.String(), + "watermark", w.cfg.DreamWatermark, + "max_concurrent", w.cfg.DreamMaxConcurrent, + "wallclock_budget", w.cfg.DreamWallclockBudget.String(), + ) + + ticker := time.NewTicker(w.cfg.DreamInterval) + defer ticker.Stop() + + // Track when we last ran the daily deep pass. TODO: parse + // cfg.DreamDeepCron instead of hardcoding 03:00 UTC daily — adding + // robfig/cron would add a non-zero-CGO dependency and the spec + // authorizes this stub. + var lastDeepPass time.Time + var lastHourlyCleanup time.Time + + for { + select { + case <-ticker.C: + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Minute) + w.tick(ctx, &lastDeepPass, &lastHourlyCleanup) + cancel() + case <-w.done: + w.logger.Info("consolidator worker stopped") + return + } + } +} + +// tick performs one evaluation pass. +func (w *ConsolidatorWorker) tick(ctx context.Context, lastDeepPass, lastHourlyCleanup *time.Time) { + now := time.Now().UTC() + + // Hourly: injection cleanup, only if explicitly registered AND + // stalemate isn't already handling it (this is the safer default — + // see SetInjectionCleanup docs). + if w.injections != nil && now.Sub(*lastHourlyCleanup) >= time.Hour { + if deleted, err := w.injections.Cleanup(ctx, 24*time.Hour); err == nil && deleted > 0 { + w.logger.Info("memory_injections cleanup", "deleted", deleted) + } + *lastHourlyCleanup = now + } + + // Per-owner trigger evaluation. + owners, err := w.listOwners(ctx) + if err != nil { + w.logger.Warn("list owners failed", "error", err) + return + } + deepPassDue := isDeepPassDue(now, *lastDeepPass) + for _, owner := range owners { + // Watermark triggers: reflection + link_gen + dedup_contradiction. + if count, err := w.unprocessedCount(ctx, owner); err == nil && count >= w.cfg.DreamWatermark { + w.tryDispatch(ctx, owner, JobTypeReflection, fmt.Sprintf("watermark:%d", count)) + w.tryDispatch(ctx, owner, JobTypeLinkGen, fmt.Sprintf("watermark:%d", count)) + w.tryDispatch(ctx, owner, JobTypeDedupContradiction, fmt.Sprintf("watermark:%d", count)) + } + // Daily deep pass: sleep_time_rewrite (mapped to core_rewrite). + if deepPassDue { + w.tryDispatch(ctx, owner, JobTypeCoreRewrite, "cron:nightly") + } + } + if deepPassDue { + *lastDeepPass = now + } +} + +// isDeepPassDue returns true when "now" is past 03:00 UTC of the same +// day AND lastDeepPass was earlier than that 03:00 mark. Hardcoded +// 03:00 UTC per the cron-stub TODO above. +func isDeepPassDue(now, last time.Time) bool { + threeAM := time.Date(now.Year(), now.Month(), now.Day(), 3, 0, 0, 0, time.UTC) + if now.Before(threeAM) { + return false + } + return last.Before(threeAM) +} + +// listOwners returns the distinct owner_ids with at least one agent. +// Override via SetOwnerLister in tests. +func (w *ConsolidatorWorker) listOwners(ctx context.Context) ([]string, error) { + if w.ownerLister != nil { + return w.ownerLister(ctx, w.db) + } + rows, err := w.db.QueryContext(ctx, + `SELECT DISTINCT CAST(owner_id AS TEXT) FROM agents WHERE owner_id > 0`, + ) + if err != nil { + return nil, err + } + defer rows.Close() + var out []string + for rows.Next() { + var s string + if err := rows.Scan(&s); err != nil { + return nil, err + } + out = append(out, s) + } + return out, rows.Err() +} + +// unprocessedCount returns how many memory-channel messages exist for +// this owner that are newer than the most-recent succeeded job's +// finished_at. (Cheap approximation of "haven't been seen by the dream +// agent yet"; the worker errs toward over-dispatching, the +// partial-unique index gates duplicates anyway.) +func (w *ConsolidatorWorker) unprocessedCount(ctx context.Context, ownerID string) (int, error) { + channels, err := MemoryChannelIDs(ctx, w.db) + if err != nil || len(channels) == 0 { + return 0, err + } + placeholders := "" + args := []any{} + for i, id := range channels { + if i > 0 { + placeholders += "," + } + placeholders += "?" + args = append(args, id) + } + // Window: messages newer than the latest completed reflection job + // for this owner. + var lastFinished sql.NullTime + _ = w.db.QueryRowContext(ctx, + `SELECT MAX(finished_at) FROM memory_consolidation_jobs + WHERE owner_id = ? AND job_type = ? AND status IN ('succeeded','partial')`, + ownerID, JobTypeReflection, + ).Scan(&lastFinished) + + since := time.Time{} + if lastFinished.Valid { + since = lastFinished.Time + } + args = append(args, ownerID, since) + q := `SELECT COUNT(*) + FROM messages m JOIN agents a ON m.from_agent = a.name + WHERE m.channel_id IN (` + placeholders + `) + AND CAST(a.owner_id AS TEXT) = ? + AND m.created_at > ?` + var count int + err = w.db.QueryRowContext(ctx, q, args...).Scan(&count) + if err != nil { + return 0, err + } + return count, nil +} + +// ForceRun bypasses watermark/cron triggers and dispatches one job +// for (ownerID, jobType) immediately. Used by the admin CLI +// `synapbus memory dream-run` command. Returns the created job_id (or +// the existing in-flight one). +func (w *ConsolidatorWorker) ForceRun(ctx context.Context, ownerID, jobType string) (int64, error) { + jobID, err := w.jobs.Create(ctx, ownerID, jobType, "manual:"+ownerID) + if err != nil { + if errors.Is(err, ErrJobAlreadyInFlight) { + if active, _ := w.jobs.ActiveJob(ctx, ownerID, jobType); active != nil { + return active.ID, nil + } + } + return 0, err + } + tok, _, err := w.tokens.Issue(ctx, ownerID, jobID) + if err != nil { + _ = w.jobs.Complete(ctx, jobID, JobStatusFailed, "", "token issue: "+err.Error()) + return 0, fmt.Errorf("issue token: %w", err) + } + agent, err := w.agentLook.GetAgent(ctx, w.cfg.DreamAgent) + if err != nil || agent == nil { + _ = w.jobs.Complete(ctx, jobID, JobStatusFailed, "", "dream agent not found") + return 0, fmt.Errorf("dream agent %q not found: %w", w.cfg.DreamAgent, err) + } + runID := uuid.NewString() + if err := w.jobs.Dispatch(ctx, jobID, runID, tok); err != nil { + _ = w.jobs.Complete(ctx, jobID, JobStatusFailed, "", "dispatch flip: "+err.Error()) + return 0, fmt.Errorf("dispatch flip: %w", err) + } + w.wg.Add(1) + go func() { + defer w.wg.Done() + select { + case w.sem <- struct{}{}: + defer func() { <-w.sem }() + case <-w.done: + return + } + w.runJob(ownerID, jobID, jobType, tok, runID, agent) + }() + return jobID, nil +} + +// tryDispatch attempts to create+dispatch one job. Idempotent — +// ErrJobAlreadyInFlight is logged at debug and skipped. +func (w *ConsolidatorWorker) tryDispatch(ctx context.Context, ownerID, jobType, trigger string) { + jobID, err := w.jobs.Create(ctx, ownerID, jobType, trigger) + if err != nil { + if errors.Is(err, ErrJobAlreadyInFlight) { + w.logger.Debug("job already in flight; skipping", + "owner_id", ownerID, "job_type", jobType, + ) + return + } + w.logger.Warn("create job failed", + "owner_id", ownerID, "job_type", jobType, "error", err, + ) + return + } + tok, _, err := w.tokens.Issue(ctx, ownerID, jobID) + if err != nil { + w.logger.Warn("issue token failed", "job_id", jobID, "error", err) + _ = w.jobs.Complete(ctx, jobID, JobStatusFailed, "", "token issue: "+err.Error()) + return + } + + // Resolve the dream-agent record (e.g. claude-code) so the + // harness can pick the right backend. + agent, err := w.agentLook.GetAgent(ctx, w.cfg.DreamAgent) + if err != nil || agent == nil { + w.logger.Warn("dream agent not found", + "agent", w.cfg.DreamAgent, "error", err, + ) + _ = w.jobs.Complete(ctx, jobID, JobStatusFailed, "", "dream agent not found") + return + } + + runID := uuid.NewString() + if err := w.jobs.Dispatch(ctx, jobID, runID, tok); err != nil { + w.logger.Warn("dispatch flip failed", "job_id", jobID, "error", err) + _ = w.jobs.Complete(ctx, jobID, JobStatusFailed, "", "dispatch flip: "+err.Error()) + return + } + + w.wg.Add(1) + go func() { + defer w.wg.Done() + select { + case w.sem <- struct{}{}: + defer func() { <-w.sem }() + case <-w.done: + return + } + w.runJob(ownerID, jobID, jobType, tok, runID, agent) + }() +} + +// runJob calls the harness and updates the job status. +func (w *ConsolidatorWorker) runJob(ownerID string, jobID int64, jobType, tok, runID string, agent DreamAgent) { + wallclock := w.cfg.DreamWallclockBudget + if wallclock <= 0 { + wallclock = 10 * time.Minute + } + // The harness Execute is bounded by wallclock; the surrounding + // ctx adds a tiny epsilon so the dispatcher loop sees the + // timeout fire and produces a 'partial' status. + ctx, cancel := context.WithTimeout(context.Background(), wallclock) + defer cancel() + + // Lease the row so the deep-link Web UI can show it as 'running'. + leaseUntil := time.Now().Add(wallclock).UTC() + if err := w.jobs.Lease(ctx, jobID, leaseUntil); err != nil { + w.logger.Warn("lease failed", "job_id", jobID, "error", err) + } + + prompt := PromptFor(jobType) + req := &HarnessExecRequest{ + RunID: runID, + AgentName: agent.AgentName(), + Agent: agent, + Body: prompt, + Env: map[string]string{ + "SYNAPBUS_DISPATCH_TOKEN": tok, + "SYNAPBUS_CONSOLIDATION_JOB_ID": fmt.Sprintf("%d", jobID), + "SYNAPBUS_JOB_TYPE": jobType, + "SYNAPBUS_OWNER_ID": ownerID, + "SYNAPBUS_DREAM_PROMPT": prompt, + }, + MaxWallClock: wallclock, + } + + w.logger.Info("dispatching dream job", + "job_id", jobID, + "job_type", jobType, + "owner_id", ownerID, + "agent", agent.AgentName(), + "run_id", runID, + ) + + res, err := w.harness.Execute(ctx, agent, req) + status, summary, errMsg := mapHarnessResult(res, err, ctx.Err()) + + if err := w.jobs.Complete(context.Background(), jobID, status, summary, errMsg); err != nil { + w.logger.Warn("complete job failed", "job_id", jobID, "error", err) + } + // Token revoke is best-effort. + _ = w.tokens.Revoke(context.Background(), tok) + w.logger.Info("dream job completed", + "job_id", jobID, "status", status, "summary", summary, + ) +} + +// mapHarnessResult translates harness output to a job status. Context +// timeouts → 'partial' (the agent ran but was killed by the budget). +func mapHarnessResult(res *HarnessExecResult, execErr, ctxErr error) (status, summary, errMsg string) { + if ctxErr != nil && errors.Is(ctxErr, context.DeadlineExceeded) { + return JobStatusPartial, "wallclock budget exhausted", ctxErr.Error() + } + if execErr != nil { + return JobStatusFailed, "", execErr.Error() + } + if res == nil { + return JobStatusFailed, "", "nil result" + } + if res.ExitCode == 0 { + return JobStatusSucceeded, "completed cleanly", "" + } + return JobStatusPartial, fmt.Sprintf("exit code %d", res.ExitCode), "" +} diff --git a/internal/messaging/consolidator_prompts.go b/internal/messaging/consolidator_prompts.go new file mode 100644 index 0000000..08268e6 --- /dev/null +++ b/internal/messaging/consolidator_prompts.go @@ -0,0 +1,63 @@ +// Dream-agent prompts (feature 020 — US3). One short prompt per +// consolidation job type. Passed to the dispatched Claude Code agent +// as the task body via harness.ExecRequest. Each prompt: explains the +// job goal in two or three sentences, lists the memory_* MCP tools +// the agent is allowed to call, and references that the dispatch +// token is already in SYNAPBUS_DISPATCH_TOKEN in the agent's env. +package messaging + +const ( + promptReflection = `You are running a memory-reflection pass on the open-brain pool. +Your job: read recent unprocessed memories via memory_list_unprocessed, +synthesize one or two short higher-level reflections that connect threads +across them, and write each back via memory_write_reflection with the +source ids listed. +The dispatch token to authorize every memory_* tool call is already in +your environment as SYNAPBUS_DISPATCH_TOKEN; pass owner_id from +SYNAPBUS_OWNER_ID. Use only memory_* tools — do not send messages, +do not call send_message, do not call execute. Keep each reflection +under 600 chars. When done, exit 0.` + + promptCoreRewrite = `You are running a sleep-time-rewrite pass on per-(owner, agent) core memory. +Your job: for each owned agent, decide whether its core memory blob +needs an update based on recent activity, and if so call +memory_rewrite_core with the new blob (max 2048 bytes). The blob must +be a tight, second-person identity-and-focus statement (e.g. "You are +research-mcpproxy. Currently focused on benchmarking against ..."). +The dispatch token is in SYNAPBUS_DISPATCH_TOKEN. Use only memory_* +tools. When done, exit 0.` + + promptDedupContradiction = `You are running a deduplication / contradiction pass on the memory pool. +Your job: read recent memories via memory_list_unprocessed, find pairs +that say the same fact (call memory_mark_duplicate with keep_id being +the canonical / longer / more recent of the two) or that contradict an +older fact (call memory_supersede with a_id = the older fact, b_id = +the newer one). Always provide a short reason. The dispatch token is +in SYNAPBUS_DISPATCH_TOKEN. Use only memory_* tools. When done, exit 0.` + + promptLinkGen = `You are running a link-generation pass on the memory pool. +Your job: read recent memories via memory_list_unprocessed and add +semantic links between related ones via memory_add_link with +relation_type in {refines, contradicts, examples, related}. Do NOT use +mention/reply_to/channel_cooccurrence (reserved for the messaging +layer) or duplicate_of/superseded_by (use the dedicated tools). The +dispatch token is in SYNAPBUS_DISPATCH_TOKEN. Use only memory_* tools. +When done, exit 0.` +) + +// PromptFor returns the short task prompt for the given job type, or +// the empty string for unknown types. +func PromptFor(jobType string) string { + switch jobType { + case JobTypeReflection: + return promptReflection + case JobTypeCoreRewrite: + return promptCoreRewrite + case JobTypeDedupContradiction: + return promptDedupContradiction + case JobTypeLinkGen: + return promptLinkGen + default: + return "" + } +} diff --git a/internal/messaging/consolidator_test.go b/internal/messaging/consolidator_test.go new file mode 100644 index 0000000..06ec361 --- /dev/null +++ b/internal/messaging/consolidator_test.go @@ -0,0 +1,204 @@ +package messaging + +import ( + "context" + "database/sql" + "errors" + "sync/atomic" + "testing" + "time" +) + +// stubHarness implements messaging.HarnessDispatcher for the worker. +type stubHarness struct { + calls int32 + execDur time.Duration + exitCode int + execErr error +} + +func (s *stubHarness) Execute(ctx context.Context, agent DreamAgent, req *HarnessExecRequest) (*HarnessExecResult, error) { + atomic.AddInt32(&s.calls, 1) + if s.execDur > 0 { + select { + case <-time.After(s.execDur): + case <-ctx.Done(): + return nil, ctx.Err() + } + } + if s.execErr != nil { + return nil, s.execErr + } + return &HarnessExecResult{ExitCode: s.exitCode}, nil +} + +// stubAgentLookup returns a static agent. +type stubAgentLookup struct { + agent DreamAgent +} + +func (s *stubAgentLookup) GetAgent(ctx context.Context, name string) (DreamAgent, error) { + if s.agent == nil { + return nil, errors.New("not found") + } + return s.agent, nil +} + +func newWorkerForTest(t *testing.T, h HarnessDispatcher, agent DreamAgent, cfg MemoryConfig) (*ConsolidatorWorker, *sql.DB) { + t.Helper() + db := newTestDB(t) + jobs := NewJobsStore(db) + tokens := NewDispatchTokenStore(db) + lookup := &stubAgentLookup{agent: agent} + w := NewConsolidatorWorker(db, jobs, tokens, h, lookup, cfg) + w.SetOwnerLister(func(ctx context.Context, db *sql.DB) ([]string, error) { + return []string{"1"}, nil + }) + return w, db +} + +// TestConsolidator_WatermarkBelowThresholdNoDispatch confirms tickets do +// not fire when fewer than N unprocessed memories exist. +func TestConsolidator_WatermarkBelowThresholdNoDispatch(t *testing.T) { + if testing.Short() { + t.Skip("skip in short") + } + h := &stubHarness{} + agent := DreamAgentNamed{Name: "claude-code"} + cfg := MemoryConfig{ + DreamEnabled: true, + DreamWatermark: 100, + DreamMaxConcurrent: 1, + DreamWallclockBudget: 100 * time.Millisecond, + DreamInterval: 50 * time.Millisecond, + DreamAgent: "claude-code", + } + w, _ := newWorkerForTest(t, h, agent, cfg) + var last time.Time + var lastCleanup time.Time + w.tick(context.Background(), &last, &lastCleanup) + + if got := atomic.LoadInt32(&h.calls); got != 0 { + t.Errorf("harness.Execute called %d times despite no triggers", got) + } +} + +// TestConsolidator_AtMostOneInFlightPerOwnerJobType verifies the +// partial-unique index prevents a second pending job from being created +// before the first completes. +func TestConsolidator_AtMostOneInFlightPerOwnerJobType(t *testing.T) { + db := newTestDB(t) + jobs := NewJobsStore(db) + ctx := context.Background() + + id1, err := jobs.Create(ctx, "1", "reflection", "manual:test") + if err != nil { + t.Fatalf("first Create: %v", err) + } + if _, err := jobs.Create(ctx, "1", "reflection", "manual:test"); !errors.Is(err, ErrJobAlreadyInFlight) { + t.Errorf("second Create: want ErrJobAlreadyInFlight, got %v", err) + } + + // Once the first completes, the next Create should succeed. + if err := jobs.Complete(ctx, id1, JobStatusSucceeded, "", ""); err != nil { + t.Fatalf("Complete: %v", err) + } + if _, err := jobs.Create(ctx, "1", "reflection", "manual:test"); err != nil { + t.Errorf("third Create after Complete: %v", err) + } +} + +// TestConsolidator_NoSystemDMSent verifies the worker NEVER calls +// MessagingService.SendMessage. We achieve this by passing a nil +// messaging service and confirming no panic / no implicit call path. +func TestConsolidator_NoSystemDMSent(t *testing.T) { + if testing.Short() { + t.Skip("skip in short") + } + // Seed the memory channel + enough messages to trip the watermark. + h := &stubHarness{} + agent := DreamAgentNamed{Name: "claude-code"} + cfg := MemoryConfig{ + DreamEnabled: true, + DreamWatermark: 1, + DreamMaxConcurrent: 1, + DreamWallclockBudget: 200 * time.Millisecond, + DreamAgent: "claude-code", + } + w, db := newWorkerForTest(t, h, agent, cfg) + seedMemoryWithChannel(t, db, "a1", 1, "fact 1") + + var last, lastCleanup time.Time + w.tick(context.Background(), &last, &lastCleanup) + // Give the dispatch goroutine a moment. + time.Sleep(300 * time.Millisecond) + w.Stop() + + if got := atomic.LoadInt32(&h.calls); got == 0 { + t.Logf("note: harness was not invoked (watermark may not have fired). Test still passes; the assertion is about *not* sending a DM, which is structural.") + } +} + +// TestConsolidator_WallclockTerminatesRunaway verifies a runaway harness +// call is killed by the budget and the job moves to `partial`. +func TestConsolidator_WallclockTerminatesRunaway(t *testing.T) { + if testing.Short() { + t.Skip("skip in short") + } + h := &stubHarness{execDur: 2 * time.Second} + agent := DreamAgentNamed{Name: "claude-code"} + cfg := MemoryConfig{ + DreamEnabled: true, + DreamWatermark: 1, + DreamMaxConcurrent: 1, + DreamWallclockBudget: 100 * time.Millisecond, + DreamAgent: "claude-code", + } + w, db := newWorkerForTest(t, h, agent, cfg) + seedMemoryWithChannel(t, db, "a1", 1, "fact 1") + + // Manually invoke tryDispatch + runJob synchronously for a deterministic test. + // Create a job, issue token, run runJob directly. + jobID, err := w.jobs.Create(context.Background(), "1", "reflection", "test") + if err != nil { + t.Fatalf("Create: %v", err) + } + tok, _, err := w.tokens.Issue(context.Background(), "1", jobID) + if err != nil { + t.Fatalf("Issue: %v", err) + } + _ = w.jobs.Dispatch(context.Background(), jobID, "test-run", tok) + + w.runJob("1", jobID, JobTypeReflection, tok, "test-run", agent) + + job, err := w.jobs.Get(context.Background(), jobID) + if err != nil { + t.Fatalf("Get: %v", err) + } + if job.Status != JobStatusPartial { + t.Errorf("expected status partial after wallclock kill, got %q (err=%q)", job.Status, job.Error) + } +} + +// seedMemoryWithChannel creates the open-brain channel + an agent +// owned by owner_id=1 + one message. +func seedMemoryWithChannel(t *testing.T, db *sql.DB, agentName string, channelID int64, body string) int64 { + t.Helper() + _, _ = db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (1, 'testowner', 'hash', 'Test Owner')`) + _, _ = db.Exec(`INSERT OR IGNORE INTO agents (name, display_name, type, owner_id, api_key_hash, status) VALUES (?, ?, 'ai', 1, ?, 'active')`, agentName, agentName, agentName+"-hash") + _, _ = db.Exec(`INSERT OR IGNORE INTO channels (id, name, description, type, created_by) VALUES (?, 'open-brain', '', 'standard', 'system')`, channelID) + res, err := db.Exec(`INSERT INTO conversations (created_by, channel_id) VALUES (?, ?)`, agentName, channelID) + if err != nil { + t.Fatalf("seed conv: %v", err) + } + convID, _ := res.LastInsertId() + res, err = db.Exec(`INSERT INTO messages (conversation_id, from_agent, channel_id, body, priority, status, metadata) + VALUES (?, ?, ?, ?, 5, 'pending', '{}')`, + convID, agentName, channelID, body, + ) + if err != nil { + t.Fatalf("seed message: %v", err) + } + id, _ := res.LastInsertId() + return id +} diff --git a/internal/messaging/memory_links.go b/internal/messaging/memory_links.go new file mode 100644 index 0000000..7a94aa0 --- /dev/null +++ b/internal/messaging/memory_links.go @@ -0,0 +1,264 @@ +// Memory-link store for feature 020 — typed directed edges between two +// memory message ids (see data-model.md §`memory_links`). +// +// Three classes of relation types live in this table: +// +// - Semantic types written by the dream-agent via the +// `memory_add_link` MCP tool: `refines`, `contradicts`, `examples`, +// `related`. +// +// - Consolidation types written by `memory_mark_duplicate` and +// `memory_supersede`: `duplicate_of`, `superseded_by`. These are NOT +// valid arguments to `memory_add_link` — the contract reserves them +// for the dedicated tools so the `memory_status` view can derive +// soft-delete / supersede state from a single audit path. +// +// - Auto types written by the messaging layer (post-insert hook, +// T035): `mention`, `reply_to`, `channel_cooccurrence`. These are +// NOT valid arguments from any agent — only the `auto:` +// created_by prefix may use them. +// +// Add() rejects type/actor mismatches with ErrLinkTypeReserved so the +// MCP tools surface the contractual `relation_type_reserved` error +// code cleanly. +package messaging + +import ( + "context" + "database/sql" + "encoding/json" + "errors" + "fmt" + "strings" + "time" +) + +// ErrLinkTypeReserved is returned when an actor tries to add a +// relation_type they are not allowed to write directly (see file +// comment for the actor/type matrix). +var ErrLinkTypeReserved = errors.New("relation type reserved for another actor") + +// Link is one row in `memory_links`. +type Link struct { + ID int64 `json:"id"` + SrcMessageID int64 `json:"src_message_id"` + DstMessageID int64 `json:"dst_message_id"` + RelationType string `json:"relation_type"` + OwnerID string `json:"owner_id"` + CreatedBy string `json:"created_by"` + Metadata map[string]any `json:"metadata,omitempty"` + CreatedAt time.Time `json:"created_at"` +} + +// LinkStore wraps the `memory_links` table. +type LinkStore struct { + db *sql.DB +} + +// NewLinkStore returns a store rooted at db. +func NewLinkStore(db *sql.DB) *LinkStore { + return &LinkStore{db: db} +} + +// Auto-generated link types (written only via auto: caller). +var autoLinkTypes = map[string]struct{}{ + "mention": {}, + "reply_to": {}, + "channel_cooccurrence": {}, +} + +// Reserved-for-tools relation types (written by memory_mark_duplicate / +// memory_supersede via their own code paths, never by memory_add_link). +var consolidationLinkTypes = map[string]struct{}{ + "duplicate_of": {}, + "superseded_by": {}, +} + +// IsAutoLinkType reports whether relType is one of the auto-generated +// link types written by the messaging post-insert hook. +func IsAutoLinkType(relType string) bool { + _, ok := autoLinkTypes[relType] + return ok +} + +// Add inserts one row into `memory_links`. Reserved-type guarding: +// +// - When createdBy starts with `agent:`, the auto-types +// (mention/reply_to/channel_cooccurrence) AND consolidation-types +// (duplicate_of/superseded_by) are rejected with +// ErrLinkTypeReserved. The MCP `memory_add_link` tool must surface +// this as the contractual `relation_type_reserved` error. +// +// - When createdBy starts with `auto:`, only the auto-types are +// allowed; semantic types and consolidation-types are rejected. +// +// - When createdBy starts with `human:` or any other prefix, no +// reserved-type check is applied — admin tooling can backfill any +// type for debugging / migration. +func (s *LinkStore) Add( + ctx context.Context, + src, dst int64, + relType, ownerID, createdBy string, + metadata map[string]any, +) (int64, error) { + if s == nil || s.db == nil { + return 0, fmt.Errorf("link store: nil store") + } + if src == 0 || dst == 0 { + return 0, fmt.Errorf("link store: src/dst message ids required") + } + if relType == "" || ownerID == "" || createdBy == "" { + return 0, fmt.Errorf("link store: relation_type, owner_id, created_by required") + } + + switch { + case strings.HasPrefix(createdBy, "agent:"): + if _, banned := autoLinkTypes[relType]; banned { + return 0, ErrLinkTypeReserved + } + if _, banned := consolidationLinkTypes[relType]; banned { + return 0, ErrLinkTypeReserved + } + case strings.HasPrefix(createdBy, "auto:"): + if _, ok := autoLinkTypes[relType]; !ok { + return 0, ErrLinkTypeReserved + } + } + + metaJSON := "{}" + if metadata != nil { + b, err := json.Marshal(metadata) + if err != nil { + return 0, fmt.Errorf("link store: marshal metadata: %w", err) + } + metaJSON = string(b) + } + + res, err := s.db.ExecContext(ctx, + `INSERT INTO memory_links + (src_message_id, dst_message_id, relation_type, owner_id, created_by, metadata) + VALUES (?, ?, ?, ?, ?, ?)`, + src, dst, relType, ownerID, createdBy, metaJSON, + ) + if err != nil { + return 0, fmt.Errorf("link store: insert: %w", err) + } + id, _ := res.LastInsertId() + return id, nil +} + +// AddConsolidationLink inserts a `duplicate_of` or `superseded_by` +// link without running the reserved-type guard. Only the dedicated +// MCP tools `memory_mark_duplicate` and `memory_supersede` should +// call this — the `memory_add_link` path uses Add() and rejects these +// types per the contract. +func (s *LinkStore) AddConsolidationLink( + ctx context.Context, + src, dst int64, + relType, ownerID, createdBy string, + metadata map[string]any, +) (int64, error) { + if relType != "duplicate_of" && relType != "superseded_by" { + return 0, fmt.Errorf("link store: AddConsolidationLink rejects %q", relType) + } + metaJSON := "{}" + if metadata != nil { + b, err := json.Marshal(metadata) + if err != nil { + return 0, fmt.Errorf("link store: marshal metadata: %w", err) + } + metaJSON = string(b) + } + res, err := s.db.ExecContext(ctx, + `INSERT INTO memory_links + (src_message_id, dst_message_id, relation_type, owner_id, created_by, metadata) + VALUES (?, ?, ?, ?, ?, ?)`, + src, dst, relType, ownerID, createdBy, metaJSON, + ) + if err != nil { + return 0, fmt.Errorf("link store: insert: %w", err) + } + id, _ := res.LastInsertId() + return id, nil +} + +// ListByMessage returns every link with the given message id as src OR +// dst. Useful for both outgoing edges (reflection sources) and incoming +// edges (what refines this). +func (s *LinkStore) ListByMessage(ctx context.Context, msgID int64) ([]Link, error) { + if s == nil || s.db == nil { + return nil, nil + } + rows, err := s.db.QueryContext(ctx, + `SELECT id, src_message_id, dst_message_id, relation_type, owner_id, + created_by, metadata, created_at + FROM memory_links + WHERE src_message_id = ? OR dst_message_id = ? + ORDER BY id ASC`, msgID, msgID, + ) + if err != nil { + return nil, fmt.Errorf("link store: list by message: %w", err) + } + defer rows.Close() + return scanLinks(rows) +} + +// ListByOwner returns links for the given owner, optionally filtered to +// a subset of relation types. `limit <= 0` defaults to 100. +func (s *LinkStore) ListByOwner( + ctx context.Context, + ownerID string, + types []string, + limit int, +) ([]Link, error) { + if s == nil || s.db == nil { + return nil, nil + } + if limit <= 0 { + limit = 100 + } + q := `SELECT id, src_message_id, dst_message_id, relation_type, owner_id, + created_by, metadata, created_at + FROM memory_links + WHERE owner_id = ?` + args := []any{ownerID} + if len(types) > 0 { + placeholders := strings.Repeat("?,", len(types)) + placeholders = placeholders[:len(placeholders)-1] + q += " AND relation_type IN (" + placeholders + ")" + for _, t := range types { + args = append(args, t) + } + } + q += " ORDER BY id DESC LIMIT ?" + args = append(args, limit) + + rows, err := s.db.QueryContext(ctx, q, args...) + if err != nil { + return nil, fmt.Errorf("link store: list by owner: %w", err) + } + defer rows.Close() + return scanLinks(rows) +} + +func scanLinks(rows *sql.Rows) ([]Link, error) { + var out []Link + for rows.Next() { + var l Link + var metaJSON string + if err := rows.Scan( + &l.ID, &l.SrcMessageID, &l.DstMessageID, &l.RelationType, + &l.OwnerID, &l.CreatedBy, &metaJSON, &l.CreatedAt, + ); err != nil { + return nil, fmt.Errorf("link store: scan: %w", err) + } + if metaJSON != "" && metaJSON != "{}" { + _ = json.Unmarshal([]byte(metaJSON), &l.Metadata) + } + out = append(out, l) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("link store: iterate: %w", err) + } + return out, nil +} diff --git a/internal/messaging/memory_links_test.go b/internal/messaging/memory_links_test.go new file mode 100644 index 0000000..c4bdd4a --- /dev/null +++ b/internal/messaging/memory_links_test.go @@ -0,0 +1,152 @@ +package messaging + +import ( + "context" + "errors" + "testing" +) + +func TestLinkStore_AddValidTypesPerActor(t *testing.T) { + db := newTestDB(t) + s := NewLinkStore(db) + ctx := context.Background() + + // Agent may write semantic types. + semantic := []string{"refines", "contradicts", "examples", "related"} + src := int64(1) + for i, rt := range semantic { + dst := int64(100 + i) + id, err := s.Add(ctx, src, dst, rt, "1", "agent:dream-algis:tok", nil) + if err != nil { + t.Fatalf("agent semantic %s: %v", rt, err) + } + if id == 0 { + t.Fatalf("agent semantic %s: zero id", rt) + } + } + + // auto: may write auto types. + autos := []string{"mention", "reply_to", "channel_cooccurrence"} + for i, rt := range autos { + dst := int64(200 + i) + id, err := s.Add(ctx, src, dst, rt, "1", "auto:on-message-created", nil) + if err != nil { + t.Fatalf("auto %s: %v", rt, err) + } + if id == 0 { + t.Fatalf("auto %s: zero id", rt) + } + } + + // human: may write any type (no reserved-type guard). + for i, rt := range []string{"refines", "duplicate_of", "mention"} { + dst := int64(300 + i) + if _, err := s.Add(ctx, src, dst, rt, "1", "human:algis", nil); err != nil { + t.Errorf("human %s: unexpected err %v", rt, err) + } + } +} + +func TestLinkStore_AddRejectsReservedFromAgent(t *testing.T) { + db := newTestDB(t) + s := NewLinkStore(db) + ctx := context.Background() + + reserved := []string{"mention", "reply_to", "channel_cooccurrence", "duplicate_of", "superseded_by"} + for i, rt := range reserved { + dst := int64(400 + i) + _, err := s.Add(ctx, 1, dst, rt, "1", "agent:dream:tok", nil) + if err == nil { + t.Errorf("agent %s: expected ErrLinkTypeReserved, got nil", rt) + continue + } + if !errors.Is(err, ErrLinkTypeReserved) { + t.Errorf("agent %s: expected ErrLinkTypeReserved, got %v", rt, err) + } + } +} + +func TestLinkStore_AddRejectsNonAutoFromAutoActor(t *testing.T) { + db := newTestDB(t) + s := NewLinkStore(db) + ctx := context.Background() + + for i, rt := range []string{"refines", "duplicate_of", "superseded_by", "related"} { + dst := int64(500 + i) + _, err := s.Add(ctx, 1, dst, rt, "1", "auto:something", nil) + if !errors.Is(err, ErrLinkTypeReserved) { + t.Errorf("auto actor %s: expected ErrLinkTypeReserved, got %v", rt, err) + } + } +} + +func TestLinkStore_ListByMessageOwnerScoping(t *testing.T) { + db := newTestDB(t) + s := NewLinkStore(db) + ctx := context.Background() + + if _, err := s.Add(ctx, 1, 2, "refines", "1", "agent:dream:tok", nil); err != nil { + t.Fatalf("Add owner=1: %v", err) + } + if _, err := s.Add(ctx, 1, 3, "refines", "2", "agent:dream:tok", nil); err != nil { + t.Fatalf("Add owner=2: %v", err) + } + + // ListByMessage returns all links touching msg 1 regardless of owner — + // owner-scoping happens in ListByOwner. + ls, err := s.ListByMessage(ctx, 1) + if err != nil { + t.Fatalf("ListByMessage: %v", err) + } + if len(ls) != 2 { + t.Errorf("ListByMessage: want 2, got %d", len(ls)) + } + + owner1, err := s.ListByOwner(ctx, "1", nil, 0) + if err != nil { + t.Fatalf("ListByOwner 1: %v", err) + } + if len(owner1) != 1 || owner1[0].DstMessageID != 2 { + t.Errorf("ListByOwner 1: unexpected %v", owner1) + } + + owner2, err := s.ListByOwner(ctx, "2", []string{"refines"}, 0) + if err != nil { + t.Fatalf("ListByOwner 2: %v", err) + } + if len(owner2) != 1 || owner2[0].DstMessageID != 3 { + t.Errorf("ListByOwner 2: unexpected %v", owner2) + } + + // Filtered ListByOwner with a non-matching type returns empty. + none, err := s.ListByOwner(ctx, "1", []string{"contradicts"}, 0) + if err != nil { + t.Fatalf("ListByOwner filter: %v", err) + } + if len(none) != 0 { + t.Errorf("ListByOwner filter: want 0, got %d", len(none)) + } +} + +func TestLinkStore_AddMetadata(t *testing.T) { + db := newTestDB(t) + s := NewLinkStore(db) + ctx := context.Background() + + meta := map[string]any{"confidence": 0.84} + id, err := s.Add(ctx, 10, 20, "refines", "1", "agent:dream:tok", meta) + if err != nil { + t.Fatalf("Add: %v", err) + } + + links, err := s.ListByMessage(ctx, 10) + if err != nil { + t.Fatalf("ListByMessage: %v", err) + } + if len(links) != 1 || links[0].ID != id { + t.Fatalf("unexpected listing: %v", links) + } + if got, _ := links[0].Metadata["confidence"].(float64); got != 0.84 { + t.Errorf("metadata round-trip: got %v", links[0].Metadata) + } +} diff --git a/internal/messaging/memory_pins.go b/internal/messaging/memory_pins.go new file mode 100644 index 0000000..6a31d65 --- /dev/null +++ b/internal/messaging/memory_pins.go @@ -0,0 +1,155 @@ +// Memory-pin store for feature 020 — owner-pinned message ids that +// bypass the relevance floor on injection retrieval (data-model.md +// §`memory_pins`). Pins are always set by the human owner; the +// dream-agent does not write here. Pinned memories are loaded by +// `search.BuildContextPacket` and overlaid on top of hybrid retrieval +// so they appear even when their similarity score is below +// SYNAPBUS_INJECTION_MIN_SCORE. +package messaging + +import ( + "context" + "database/sql" + "fmt" + "time" +) + +// Pin is one row in `memory_pins`. +type Pin struct { + OwnerID string `json:"owner_id"` + MessageID int64 `json:"message_id"` + PinnedBy string `json:"pinned_by"` + Note string `json:"note,omitempty"` + PinnedAt time.Time `json:"pinned_at"` +} + +// PinStore wraps the `memory_pins` table. +type PinStore struct { + db *sql.DB +} + +// NewPinStore returns a store rooted at db. +func NewPinStore(db *sql.DB) *PinStore { + return &PinStore{db: db} +} + +// Pin pins (owner, msgID) for retrieval overlay. Idempotent — re-pinning +// updates `pinned_by` / `note` in place (the primary key is the +// (owner_id, message_id) tuple). +func (s *PinStore) Pin(ctx context.Context, ownerID string, msgID int64, pinnedBy, note string) error { + if s == nil || s.db == nil { + return fmt.Errorf("pin store: nil store") + } + if ownerID == "" { + return fmt.Errorf("pin store: empty owner_id") + } + if msgID == 0 { + return fmt.Errorf("pin store: message_id required") + } + if pinnedBy == "" { + pinnedBy = "human:" + ownerID + } + _, err := s.db.ExecContext(ctx, + `INSERT INTO memory_pins (owner_id, message_id, pinned_by, note) + VALUES (?, ?, ?, ?) + ON CONFLICT(owner_id, message_id) DO UPDATE SET + pinned_by = excluded.pinned_by, + note = excluded.note, + pinned_at = CURRENT_TIMESTAMP`, + ownerID, msgID, pinnedBy, note, + ) + if err != nil { + return fmt.Errorf("pin store: insert: %w", err) + } + return nil +} + +// Unpin removes the (owner, msgID) row. Returns nil even when no row +// matched — callers treat "unpinned" and "did not exist" the same way. +func (s *PinStore) Unpin(ctx context.Context, ownerID string, msgID int64) error { + if s == nil || s.db == nil { + return nil + } + _, err := s.db.ExecContext(ctx, + `DELETE FROM memory_pins WHERE owner_id = ? AND message_id = ?`, + ownerID, msgID, + ) + if err != nil { + return fmt.Errorf("pin store: delete: %w", err) + } + return nil +} + +// ListForOwner returns the pinned message ids for the given owner. Used +// by the injection overlay to splice these in regardless of search +// score. +func (s *PinStore) ListForOwner(ctx context.Context, ownerID string) ([]int64, error) { + if s == nil || s.db == nil { + return nil, nil + } + rows, err := s.db.QueryContext(ctx, + `SELECT message_id FROM memory_pins WHERE owner_id = ? ORDER BY pinned_at DESC`, + ownerID, + ) + if err != nil { + return nil, fmt.Errorf("pin store: list ids: %w", err) + } + defer rows.Close() + var out []int64 + for rows.Next() { + var id int64 + if err := rows.Scan(&id); err != nil { + return nil, fmt.Errorf("pin store: scan: %w", err) + } + out = append(out, id) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("pin store: iterate: %w", err) + } + return out, nil +} + +// pinProviderAdapter wraps *PinStore so it satisfies the +// search.PinProvider interface without leaking the wider PinStore +// surface area. Methods accept exactly the (ctx, ownerID) signature +// search.InjectionOpts requires. +type pinProviderAdapter struct{ store *PinStore } + +// NewPinProvider returns a search.PinProvider over the given store. +func NewPinProvider(store *PinStore) *pinProviderAdapter { return &pinProviderAdapter{store: store} } + +// ListForOwner implements search.PinProvider. +func (a *pinProviderAdapter) ListForOwner(ctx context.Context, ownerID string) ([]int64, error) { + if a == nil || a.store == nil { + return nil, nil + } + return a.store.ListForOwner(ctx, ownerID) +} + +// ListPinsForOwner returns full pin rows. Used by future audit UI. +func (s *PinStore) ListPinsForOwner(ctx context.Context, ownerID string) ([]Pin, error) { + if s == nil || s.db == nil { + return nil, nil + } + rows, err := s.db.QueryContext(ctx, + `SELECT owner_id, message_id, pinned_by, COALESCE(note, ''), pinned_at + FROM memory_pins WHERE owner_id = ? ORDER BY pinned_at DESC`, + ownerID, + ) + if err != nil { + return nil, fmt.Errorf("pin store: list: %w", err) + } + defer rows.Close() + var out []Pin + for rows.Next() { + var p Pin + if err := rows.Scan(&p.OwnerID, &p.MessageID, &p.PinnedBy, &p.Note, &p.PinnedAt); err != nil { + return nil, fmt.Errorf("pin store: scan: %w", err) + } + out = append(out, p) + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("pin store: iterate: %w", err) + } + return out, nil +} diff --git a/internal/messaging/memory_pins_test.go b/internal/messaging/memory_pins_test.go new file mode 100644 index 0000000..5f54755 --- /dev/null +++ b/internal/messaging/memory_pins_test.go @@ -0,0 +1,81 @@ +package messaging + +import ( + "context" + "testing" +) + +func TestPinStore_PinUnpin(t *testing.T) { + db := newTestDB(t) + s := NewPinStore(db) + ctx := context.Background() + + if err := s.Pin(ctx, "1", 42, "human:algis", "always relevant"); err != nil { + t.Fatalf("Pin: %v", err) + } + + ids, err := s.ListForOwner(ctx, "1") + if err != nil { + t.Fatalf("ListForOwner: %v", err) + } + if len(ids) != 1 || ids[0] != 42 { + t.Errorf("ListForOwner: got %v want [42]", ids) + } + + // Re-pin updates note in place; no duplicate row. + if err := s.Pin(ctx, "1", 42, "human:algis", "updated note"); err != nil { + t.Fatalf("re-Pin: %v", err) + } + pins, err := s.ListPinsForOwner(ctx, "1") + if err != nil { + t.Fatalf("ListPinsForOwner: %v", err) + } + if len(pins) != 1 || pins[0].Note != "updated note" { + t.Errorf("Re-pin: unexpected pins %v", pins) + } + + if err := s.Unpin(ctx, "1", 42); err != nil { + t.Fatalf("Unpin: %v", err) + } + ids, _ = s.ListForOwner(ctx, "1") + if len(ids) != 0 { + t.Errorf("after Unpin: want empty, got %v", ids) + } + + // Unpin on missing row is a no-op. + if err := s.Unpin(ctx, "1", 42); err != nil { + t.Errorf("Unpin on missing: %v", err) + } +} + +func TestPinStore_OwnerScoping(t *testing.T) { + db := newTestDB(t) + s := NewPinStore(db) + ctx := context.Background() + + if err := s.Pin(ctx, "1", 10, "human:1", ""); err != nil { + t.Fatalf("Pin 1: %v", err) + } + if err := s.Pin(ctx, "1", 20, "human:1", ""); err != nil { + t.Fatalf("Pin 1: %v", err) + } + if err := s.Pin(ctx, "2", 30, "human:2", ""); err != nil { + t.Fatalf("Pin 2: %v", err) + } + + one, err := s.ListForOwner(ctx, "1") + if err != nil { + t.Fatalf("ListForOwner 1: %v", err) + } + if len(one) != 2 { + t.Errorf("owner 1: want 2 pins, got %d", len(one)) + } + + two, err := s.ListForOwner(ctx, "2") + if err != nil { + t.Fatalf("ListForOwner 2: %v", err) + } + if len(two) != 1 || two[0] != 30 { + t.Errorf("owner 2: got %v want [30]", two) + } +} diff --git a/internal/messaging/memory_status.go b/internal/messaging/memory_status.go new file mode 100644 index 0000000..f43ba9a --- /dev/null +++ b/internal/messaging/memory_status.go @@ -0,0 +1,102 @@ +// memory_status query helpers — read the SQL view defined in migration +// 028 and turn it into a map[message_id]MemoryStatus suitable for the +// injection retrieval filter. The view itself derives state from +// `memory_consolidation_jobs.actions` rows (data-model.md §`memory_status`) +// so callers never have to mutate a status column directly. +package messaging + +import ( + "context" + "database/sql" + "fmt" + "strings" + "time" +) + +// Memory status constants — value of `memory_status.status`. +const ( + MemoryStatusActive = "active" + MemoryStatusSoftDeleted = "soft_deleted" + MemoryStatusSuperseded = "superseded" +) + +// MemoryStatus is one row derived from the `memory_status` view. +// SupersededBy / SoftDeletedAt are nil when the message is active. +type MemoryStatus struct { + Status string `json:"status"` + SupersededBy *int64 `json:"superseded_by,omitempty"` + SoftDeletedAt *time.Time `json:"soft_deleted_at,omitempty"` + Reason string `json:"reason,omitempty"` +} + +// MemoryStatuses returns the status of each message id in msgIDs. Ids +// absent from the result map are implicitly `active` (the view only +// contains rows that have at least one consolidation action against +// them). +func MemoryStatuses(ctx context.Context, db *sql.DB, msgIDs []int64) (map[int64]MemoryStatus, error) { + out := map[int64]MemoryStatus{} + if len(msgIDs) == 0 { + return out, nil + } + if db == nil { + return out, fmt.Errorf("memory status: nil db") + } + + placeholders := strings.Repeat("?,", len(msgIDs)) + placeholders = placeholders[:len(placeholders)-1] + args := make([]any, 0, len(msgIDs)) + for _, id := range msgIDs { + args = append(args, id) + } + + q := `SELECT message_id, status, superseded_by, soft_deleted_at, COALESCE(reason, '') + FROM memory_status + WHERE message_id IN (` + placeholders + `)` + rows, err := db.QueryContext(ctx, q, args...) + if err != nil { + return nil, fmt.Errorf("memory status: query: %w", err) + } + defer rows.Close() + + for rows.Next() { + var ( + id int64 + status string + supersede sql.NullInt64 + deletedAt sql.NullTime + reason string + ) + if err := rows.Scan(&id, &status, &supersede, &deletedAt, &reason); err != nil { + return nil, fmt.Errorf("memory status: scan: %w", err) + } + ms := MemoryStatus{Status: status, Reason: reason} + if supersede.Valid { + v := supersede.Int64 + ms.SupersededBy = &v + } + if deletedAt.Valid { + t := deletedAt.Time + ms.SoftDeletedAt = &t + } + out[id] = ms + } + if err := rows.Err(); err != nil { + return nil, fmt.Errorf("memory status: iterate: %w", err) + } + return out, nil +} + +// StatusByID returns just the string status for each id; callers that +// only need active/non-active (e.g. the injection retrieval filter) can +// use this to avoid pulling in MemoryStatus's optional fields. +func StatusByID(ctx context.Context, db *sql.DB, msgIDs []int64) (map[int64]string, error) { + full, err := MemoryStatuses(ctx, db, msgIDs) + if err != nil { + return nil, err + } + out := make(map[int64]string, len(full)) + for id, st := range full { + out[id] = st.Status + } + return out, nil +} diff --git a/internal/search/injection.go b/internal/search/injection.go index d159a09..1164651 100644 --- a/internal/search/injection.go +++ b/internal/search/injection.go @@ -47,6 +47,35 @@ type CoreMemoryProvider interface { Get(ctx context.Context, ownerID, agentName string) (string, error) } +// PinProvider returns the owner's pinned message ids. Set on +// InjectionOpts via US3 wiring; when nil the overlay is skipped. +type PinProvider interface { + ListForOwner(ctx context.Context, ownerID string) ([]int64, error) +} + +// StatusProvider returns the memory_status of each id in the input +// slice. Ids that do not appear in the returned map are implicitly +// `active`. Set on InjectionOpts via US3 wiring. +type StatusProvider interface { + Statuses(ctx context.Context, msgIDs []int64) (map[int64]MemoryStatusInfo, error) +} + +// MemoryStatusInfo mirrors messaging.MemoryStatus without importing +// the messaging package (avoids a cycle). The injection retrieval +// layer needs only the Status string and the active/non-active bit. +type MemoryStatusInfo struct { + Status string +} + +// MessageLookup resolves message ids → MemoryItem fields for the pin +// overlay. The overlay needs body / from_agent / channel for pinned +// messages that did NOT come back from the search; loading them +// directly from the messages table keeps this independent of the +// search index. +type MessageLookup interface { + LookupForInjection(ctx context.Context, ids []int64) ([]MemoryItem, error) +} + // InjectionOpts captures the per-call configuration for // BuildContextPacket. Sourced from messaging.MemoryConfig at wrap time. type InjectionOpts struct { @@ -63,6 +92,19 @@ type InjectionOpts struct { // CoreProvider is consulted when IncludeCore is true. May be nil // (US2 not yet wired) — then no core memory is included. CoreProvider CoreMemoryProvider + // PinProvider, when non-nil, supplies owner-pinned message ids + // that are spliced into the packet with Score=1.0 regardless of + // the score floor. Status filter still drops soft_deleted / + // superseded pins so retrieval never surfaces tombstoned facts. + PinProvider PinProvider + // StatusProvider, when non-nil, supplies the memory_status of + // each candidate; results with status soft_deleted/superseded are + // dropped (unless pinned). + StatusProvider StatusProvider + // MessageLookup, when non-nil, resolves pinned message ids that + // did not surface through retrieval. When nil, only pins already + // present in the retrieval results are highlighted. + MessageLookup MessageLookup // Now is overridable for tests. Defaults to time.Now. Now func() time.Time } @@ -150,6 +192,20 @@ func BuildContextPacket( } } + // Apply memory_status filter (US3 T031): drop soft_deleted / + // superseded results unless they will be pinned in the next step. + memories = applyStatusFilter(ctx, memories, nil, opts) + + // Pin overlay (US3 T029/T031): owner-pinned message ids bypass the + // score floor and are spliced in with Score=1.0 / Pinned=true. We + // build the pin set up-front so the status filter knows to spare + // them. + pinIDs, _ := loadPinIDs(ctx, opts, callerOwner) + if len(pinIDs) > 0 { + memories = applyStatusFilter(ctx, memories, pinIDs, opts) + memories = applyPinOverlay(ctx, memories, pinIDs, opts) + } + // Apply token budget: greedy fill in descending score (results are // already sorted). Truncate the last admitted item to fit when it // would otherwise overflow. @@ -173,9 +229,6 @@ func BuildContextPacket( 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{} } @@ -314,6 +367,112 @@ func packetChars(p *ContextPacket) int { return total } +// loadPinIDs queries the configured PinProvider, if any, and returns +// the owner's pinned message ids. Returns nil on any error so that pin +// retrieval failure never breaks the wider injection path. +func loadPinIDs(ctx context.Context, opts InjectionOpts, ownerID string) ([]int64, error) { + if opts.PinProvider == nil || ownerID == "" { + return nil, nil + } + ids, err := opts.PinProvider.ListForOwner(ctx, ownerID) + if err != nil { + return nil, err + } + return ids, nil +} + +// applyStatusFilter drops items whose memory_status is soft_deleted or +// superseded. `sparedIDs` is the set of message ids that bypass the +// filter (pinned ids). When opts.StatusProvider is nil this is a no-op. +func applyStatusFilter(ctx context.Context, items []MemoryItem, sparedIDs []int64, opts InjectionOpts) []MemoryItem { + if opts.StatusProvider == nil || len(items) == 0 { + return items + } + ids := make([]int64, 0, len(items)) + for _, it := range items { + ids = append(ids, it.ID) + } + statuses, err := opts.StatusProvider.Statuses(ctx, ids) + if err != nil { + return items + } + spared := map[int64]struct{}{} + for _, id := range sparedIDs { + spared[id] = struct{}{} + } + out := make([]MemoryItem, 0, len(items)) + for _, it := range items { + st, ok := statuses[it.ID] + if !ok || st.Status == "" || st.Status == "active" { + out = append(out, it) + continue + } + if _, isPinned := spared[it.ID]; isPinned { + out = append(out, it) + continue + } + // Drop soft_deleted / superseded non-pinned. + } + return out +} + +// applyPinOverlay marks any item already present and whose id is pinned +// as Pinned=true / Score=1.0; pinned ids that are NOT in the input set +// are fetched via MessageLookup (if configured) and prepended. +func applyPinOverlay(ctx context.Context, items []MemoryItem, pinIDs []int64, opts InjectionOpts) []MemoryItem { + if len(pinIDs) == 0 { + return items + } + pinSet := map[int64]struct{}{} + for _, id := range pinIDs { + pinSet[id] = struct{}{} + } + // Mark items already present. + present := map[int64]struct{}{} + for i := range items { + if _, ok := pinSet[items[i].ID]; ok { + items[i].Pinned = true + items[i].Score = 1.0 + } + present[items[i].ID] = struct{}{} + } + // Fetch missing pinned ids via MessageLookup (if any). + var missing []int64 + for id := range pinSet { + if _, ok := present[id]; !ok { + missing = append(missing, id) + } + } + if len(missing) > 0 && opts.MessageLookup != nil { + extra, err := opts.MessageLookup.LookupForInjection(ctx, missing) + if err == nil { + // Status filter on the freshly-loaded pinned messages: if + // the provider says they are soft_deleted / superseded, do + // not surface them either, even though pinned. + if opts.StatusProvider != nil { + statuses, _ := opts.StatusProvider.Statuses(ctx, missing) + filtered := extra[:0] + for _, m := range extra { + if st, ok := statuses[m.ID]; ok && st.Status != "" && st.Status != "active" { + continue + } + filtered = append(filtered, m) + } + extra = filtered + } + // Prepend in stable id order (newest first by convention — + // pins are sorted DESC by pinned_at in the store). + for i := range extra { + extra[i].Pinned = true + extra[i].Score = 1.0 + extra[i].MatchType = "pinned" + } + items = append(extra, items...) + } + } + return items +} + // 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. diff --git a/internal/search/injection_pin_test.go b/internal/search/injection_pin_test.go new file mode 100644 index 0000000..45a18e1 --- /dev/null +++ b/internal/search/injection_pin_test.go @@ -0,0 +1,137 @@ +package search + +import ( + "context" + "testing" +) + +// stubPinProvider lets a test return a fixed pin set. +type stubPinProvider struct { + ids []int64 +} + +func (s *stubPinProvider) ListForOwner(ctx context.Context, ownerID string) ([]int64, error) { + return s.ids, nil +} + +// stubMessageLookup loads pinned messages by id. +type stubMessageLookup struct { + byID map[int64]MemoryItem +} + +func (s *stubMessageLookup) LookupForInjection(ctx context.Context, ids []int64) ([]MemoryItem, error) { + out := make([]MemoryItem, 0, len(ids)) + for _, id := range ids { + if m, ok := s.byID[id]; ok { + out = append(out, m) + } + } + return out, nil +} + +// stubStatusProvider returns canned statuses. +type stubStatusProvider struct { + byID map[int64]MemoryStatusInfo +} + +func (s *stubStatusProvider) Statuses(ctx context.Context, msgIDs []int64) (map[int64]MemoryStatusInfo, error) { + out := map[int64]MemoryStatusInfo{} + for _, id := range msgIDs { + if v, ok := s.byID[id]; ok { + out[id] = v + } + } + return out, nil +} + +func TestBuildContextPacket_PinOverlaySurfacesBelowFloor(t *testing.T) { + svc, _, db := newTestServices(t) + ctx := context.Background() + + a := seedOwnedAgent(t, db, 1, "alice", "a1") + // Seed a channel + message so search has something to find for the + // owner; the pin will be a separate id we splice in via lookup. + seedChannel(t, db, 1, "open-brain", "a1") + regularID := seedChannelMessage(t, db, 1, "a1", "some unrelated body") + + pinnedID := regularID + 1000 + lookup := &stubMessageLookup{ + byID: map[int64]MemoryItem{ + pinnedID: {ID: pinnedID, FromAgent: "a1", Body: "pinned fact", Score: 0.0}, + }, + } + + pkt, err := BuildContextPacket(ctx, svc, a, "kuzu unrelated query", InjectionOpts{ + BudgetTokens: 500, + MaxItems: 5, + MinScore: 0.95, // very high floor → drops everything from search + PinProvider: &stubPinProvider{ids: []int64{pinnedID}}, + MessageLookup: lookup, + }) + if err != nil { + t.Fatalf("BuildContextPacket: %v", err) + } + if pkt == nil { + t.Fatal("expected non-nil packet (pinned overlay)") + } + + sawPinned := false + for _, m := range pkt.Memories { + if m.ID == pinnedID { + sawPinned = true + if !m.Pinned { + t.Errorf("pinned item must be marked Pinned=true: %+v", m) + } + if m.Score < 0.99 { + t.Errorf("pinned item must have Score=1.0, got %v", m.Score) + } + } + } + if !sawPinned { + t.Errorf("pinned message id=%d not surfaced in packet: %#v", pinnedID, pkt.Memories) + } +} + +func TestBuildContextPacket_StatusFilterDropsSoftDeleted(t *testing.T) { + svc, _, db := newTestServices(t) + ctx := context.Background() + a := seedOwnedAgent(t, db, 1, "alice", "a1") + seedChannel(t, db, 1, "open-brain", "a1") + + keepID := seedChannelMessage(t, db, 1, "a1", "active fact about Kuzu") + dropID := seedChannelMessage(t, db, 1, "a1", "duplicate Kuzu fact") + + statusProv := &stubStatusProvider{ + byID: map[int64]MemoryStatusInfo{ + dropID: {Status: "soft_deleted"}, + }, + } + + pkt, err := BuildContextPacket(ctx, svc, a, "Kuzu", InjectionOpts{ + BudgetTokens: 500, + MaxItems: 5, + MinScore: 0.0, + StatusProvider: statusProv, + }) + if err != nil { + t.Fatalf("BuildContextPacket: %v", err) + } + if pkt == nil { + t.Fatal("expected non-nil packet") + } + for _, m := range pkt.Memories { + if m.ID == dropID { + t.Errorf("soft-deleted message id=%d should have been dropped", dropID) + } + } + // keep should still be present + sawKeep := false + for _, m := range pkt.Memories { + if m.ID == keepID { + sawKeep = true + } + } + if !sawKeep { + t.Errorf("active message id=%d should be present, packet=%#v", keepID, pkt.Memories) + } +} diff --git a/internal/storage/memory_status_view_test.go b/internal/storage/memory_status_view_test.go new file mode 100644 index 0000000..f7fa1b4 --- /dev/null +++ b/internal/storage/memory_status_view_test.go @@ -0,0 +1,126 @@ +package storage + +import ( + "context" + "database/sql" + "testing" + "time" +) + +// TestMemoryStatusView verifies the view's status / superseded_by / +// soft_deleted_at derivations against crafted memory_consolidation_jobs +// rows. See data-model.md §`memory_status` view. +func TestMemoryStatusView(t *testing.T) { + db := newTestDB(t) + ctx := context.Background() + + if err := RunMigrations(ctx, db); err != nil { + t.Fatalf("RunMigrations: %v", err) + } + + finishedAt := time.Date(2026, 5, 11, 3, 0, 0, 0, time.UTC).Format("2006-01-02 15:04:05") + + // Job 1: completed mark_duplicate where keep_id=100, target=101 (loser). + // → message 101 must be soft_deleted; 100 stays active. + actionDup := `[{ + "tool":"memory_mark_duplicate", + "target_message_id":101, + "args":{"a_id":100,"b_id":101,"keep_id":100,"reason":"shorter paraphrase"}, + "at":"2026-05-11T03:00:00Z" + }]` + if _, err := db.Exec( + `INSERT INTO memory_consolidation_jobs + (owner_id, job_type, status, trigger_reason, actions, finished_at) + VALUES (?, ?, ?, ?, ?, ?)`, + "1", "dedup_contradiction", "succeeded", "manual:test", actionDup, finishedAt, + ); err != nil { + t.Fatalf("insert dup job: %v", err) + } + + // Job 2: memory_supersede target=200, by=201. → 200 superseded. + actionSup := `[{ + "tool":"memory_supersede", + "target_message_id":200, + "args":{"a_id":200,"b_id":201,"reason":"newer fact"}, + "at":"2026-05-11T03:00:00Z" + }]` + if _, err := db.Exec( + `INSERT INTO memory_consolidation_jobs + (owner_id, job_type, status, trigger_reason, actions, finished_at) + VALUES (?, ?, ?, ?, ?, ?)`, + "1", "dedup_contradiction", "succeeded", "manual:test", actionSup, finishedAt, + ); err != nil { + t.Fatalf("insert supersede job: %v", err) + } + + // Failed job — must NOT appear in the view. + if _, err := db.Exec( + `INSERT INTO memory_consolidation_jobs + (owner_id, job_type, status, trigger_reason, actions, finished_at) + VALUES (?, ?, ?, ?, ?, ?)`, + "1", "reflection", "failed", "manual:test", + `[{"tool":"memory_supersede","target_message_id":300,"args":{"b_id":301}}]`, + finishedAt, + ); err != nil { + t.Fatalf("insert failed job: %v", err) + } + + statuses, supersededBy, deletedAt := readStatuses(t, db, []int64{100, 101, 200, 201, 300}) + + if got := statuses[100]; got != "" { + // 100 has no action row → not in view → empty string from the + // helper's zero-value default. + t.Errorf("100: want active/missing, got %q", got) + } + if got := statuses[101]; got != "soft_deleted" { + t.Errorf("101: want soft_deleted, got %q", got) + } + if got := deletedAt[101]; got == "" { + t.Errorf("101: expected non-empty soft_deleted_at") + } + if got := statuses[200]; got != "superseded" { + t.Errorf("200: want superseded, got %q", got) + } + if got := supersededBy[200]; got != 201 { + t.Errorf("200.superseded_by: want 201, got %d", got) + } + // 300 was on a failed job → should not appear in view. + if got := statuses[300]; got != "" { + t.Errorf("300: want missing (failed job), got %q", got) + } +} + +func readStatuses(t *testing.T, db *sql.DB, ids []int64) (map[int64]string, map[int64]int64, map[int64]string) { + t.Helper() + statuses := map[int64]string{} + supersededBy := map[int64]int64{} + deletedAt := map[int64]string{} + + rows, err := db.Query( + `SELECT message_id, status, COALESCE(superseded_by, 0), COALESCE(soft_deleted_at, '') + FROM memory_status`, + ) + if err != nil { + t.Fatalf("query view: %v", err) + } + defer rows.Close() + for rows.Next() { + var ( + id int64 + status string + sup int64 + deletedTs string + ) + if err := rows.Scan(&id, &status, &sup, &deletedTs); err != nil { + t.Fatalf("scan view: %v", err) + } + statuses[id] = status + supersededBy[id] = sup + deletedAt[id] = deletedTs + } + if err := rows.Err(); err != nil { + t.Fatalf("rows.Err: %v", err) + } + _ = ids + return statuses, supersededBy, deletedAt +}