diff --git a/compact/compact.go b/compact/compact.go index 209b6ad..daa754d 100644 --- a/compact/compact.go +++ b/compact/compact.go @@ -9,6 +9,8 @@ import ( "fmt" "log/slog" "strings" + "time" + "unicode/utf8" "github.com/GrayCodeAI/yaad/storage" "github.com/google/uuid" @@ -48,11 +50,24 @@ func (c *Compactor) WithSummarizer(s Summarizer) *Compactor { } // NeedsCompaction returns true if total content exceeds the token budget. +// An empty project is treated as "no compaction needed" — measuring the whole +// store would compact across unrelated projects. func (c *Compactor) NeedsCompaction(ctx context.Context, project string) (bool, int) { - nodes, _ := c.store.ListNodes(ctx, storage.NodeFilter{Project: project}) + if err := ctx.Err(); err != nil { + return false, 0 + } + if strings.TrimSpace(project) == "" { + return false, 0 + } + nodes, err := c.store.ListNodes(ctx, storage.NodeFilter{Project: project}) + if err != nil { + slog.Warn("compact: NeedsCompaction list failed", "project", project, "error", err) + return false, 0 + } totalTokens := 0 for _, n := range nodes { - totalTokens += len(n.Content) / 4 // ~4 chars per token + // ~4 chars (runes, not bytes) per token; count Summary too. + totalTokens += (utf8.RuneCountInString(n.Content) + utf8.RuneCountInString(n.Summary)) / 4 } return totalTokens > c.maxTokens, totalTokens } @@ -63,6 +78,9 @@ func (c *Compactor) Compact(ctx context.Context, project string) (int, error) { if err := ctx.Err(); err != nil { return 0, err } + if strings.TrimSpace(project) == "" { + return 0, fmt.Errorf("compact: project is required; refusing to compact across all projects") + } nodes, err := c.store.ListNodes(ctx, storage.NodeFilter{Project: project}) if err != nil { return 0, err @@ -74,6 +92,9 @@ func (c *Compactor) Compact(ctx context.Context, project string) (int, error) { if n.Type == "file" || n.Type == "entity" || n.Type == "session" { continue // don't compact anchors or sessions } + if n.Confidence <= 0 { + continue // archived nodes must never be re-compacted + } if n.Confidence < 0.5 && n.AccessCount < 3 { byType[n.Type] = append(byType[n.Type], n) } @@ -84,10 +105,13 @@ func (c *Compactor) Compact(ctx context.Context, project string) (int, error) { if len(group) < 3 { continue // not enough to compact } + if err := ctx.Err(); err != nil { + return compacted, err + } // Build summary from group - var contents []string - var ids []string + contents := make([]string, 0, len(group)) + ids := make([]string, 0, len(group)) for _, n := range group { contents = append(contents, n.Content) ids = append(ids, n.ID) @@ -103,6 +127,7 @@ func (c *Compactor) Compact(ctx context.Context, project string) (int, error) { // Create summary node hashInput := strings.Join(ids, "\x00") contentHash := fmt.Sprintf("%x", sha256.Sum256([]byte(hashInput))) + now := time.Now() summaryNode := &storage.Node{ ID: uuid.New().String(), Type: typ, @@ -114,73 +139,83 @@ func (c *Compactor) Compact(ctx context.Context, project string) (int, error) { Tier: 3, // cold Confidence: 0.6, Version: 1, - } - if err := c.store.CreateNode(ctx, summaryNode); err != nil { - // Skipping this group is safe — nothing else has been mutated yet. - slog.Warn("compact: create summary node failed, skipping group", - "type", typ, "summary_node_id", summaryNode.ID, "error", err) - continue + CreatedAt: now, + UpdatedAt: now, } - // Re-link edges: transfer edges from compacted nodes to the summary node. - // From here on the summary node exists, so failures would leave the graph - // partially re-linked — abort and propagate instead of silently continuing. - compactedIDs := make(map[string]bool, len(ids)) - for _, id := range ids { - compactedIDs[id] = true - } - for _, id := range ids { - // Outbound edges: compacted → other - outEdges, err := c.store.GetEdgesFrom(ctx, id) - if err != nil { - return compacted, fmt.Errorf("compact: list outbound edges of node %s: %w", id, err) + // The whole mutation — summary node creation, edge re-linking and + // archival — is one transaction so a mid-way failure leaves the graph + // untouched instead of half re-linked. + err = c.store.WithTx(ctx, func(tx storage.Storage) error { + if err := tx.CreateNode(ctx, summaryNode); err != nil { + return fmt.Errorf("compact: create summary node %s: %w", summaryNode.ID, err) + } + + compactedIDs := make(map[string]bool, len(ids)) + for _, id := range ids { + compactedIDs[id] = true } - for _, e := range outEdges { - if compactedIDs[e.ToID] { - continue // skip edges between compacted nodes + for _, id := range ids { + if err := ctx.Err(); err != nil { + return err } - if err := c.relinkEdge(ctx, summaryNode.ID, e.ToID, e); err != nil { - return compacted, err + outEdges, err := tx.GetEdgesFrom(ctx, id) + if err != nil { + return fmt.Errorf("compact: list outbound edges of node %s: %w", id, err) } - } - // Inbound edges: other → compacted - inEdges, err := c.store.GetEdgesTo(ctx, id) - if err != nil { - return compacted, fmt.Errorf("compact: list inbound edges of node %s: %w", id, err) - } - for _, e := range inEdges { - if compactedIDs[e.FromID] { - continue + for _, e := range outEdges { + if compactedIDs[e.ToID] { + continue // skip edges between compacted nodes + } + if err := relinkEdge(ctx, tx, summaryNode.ID, e.ToID, e); err != nil { + return err + } + } + inEdges, err := tx.GetEdgesTo(ctx, id) + if err != nil { + return fmt.Errorf("compact: list inbound edges of node %s: %w", id, err) } - if err := c.relinkEdge(ctx, e.FromID, summaryNode.ID, e); err != nil { - return compacted, err + for _, e := range inEdges { + if compactedIDs[e.FromID] { + continue + } + if err := relinkEdge(ctx, tx, e.FromID, summaryNode.ID, e); err != nil { + return err + } } } - } - // Archive compacted nodes - for _, id := range ids { - old, err := c.store.GetNode(ctx, id) - if err != nil { - if errors.Is(err, storage.ErrNodeNotFound) { - // Node disappeared concurrently — nothing to archive. - slog.Warn("compact: node vanished before archival", "node_id", id) + // Archive compacted nodes + for _, id := range ids { + if err := ctx.Err(); err != nil { + return err + } + old, err := tx.GetNode(ctx, id) + if err != nil { + if errors.Is(err, storage.ErrNodeNotFound) { + // Node disappeared concurrently — nothing to archive. + slog.Warn("compact: node vanished before archival", "node_id", id) + continue + } + return fmt.Errorf("compact: load node %s for archival: %w", id, err) + } + if old == nil { continue } - return compacted, fmt.Errorf("compact: load node %s for archival: %w", id, err) - } - if old == nil { - continue - } - if err := c.store.SaveVersion(ctx, old.ID, old.Content, "compactor", "compacted into "+summaryNode.ID[:8]); err != nil { - return compacted, fmt.Errorf("compact: save version of node %s: %w", old.ID, err) - } - old.Confidence = 0 - if err := c.store.UpdateNode(ctx, old); err != nil { - return compacted, fmt.Errorf("compact: archive node %s: %w", old.ID, err) + if err := tx.SaveVersion(ctx, old.ID, old.Content, "compactor", "compacted into "+summaryNode.ID[:8]); err != nil { + return fmt.Errorf("compact: save version of node %s: %w", old.ID, err) + } + old.Confidence = 0 + if err := tx.UpdateNode(ctx, old); err != nil { + return fmt.Errorf("compact: archive node %s: %w", old.ID, err) + } } - compacted++ + return nil + }) + if err != nil { + return compacted, fmt.Errorf("compact: commit compaction of %s: %w", typ, err) } + compacted += len(ids) } return compacted, nil } @@ -188,8 +223,8 @@ func (c *Compactor) Compact(ctx context.Context, project string) (int, error) { // relinkEdge transfers an edge onto the summary node. Duplicate edges are // benign (the link already exists) and are logged at debug level; any other // failure is propagated so the compaction pipeline can abort. -func (c *Compactor) relinkEdge(ctx context.Context, fromID, toID string, orig *storage.Edge) error { - err := c.store.CreateEdge(ctx, &storage.Edge{ +func relinkEdge(ctx context.Context, store storage.Storage, fromID, toID string, orig *storage.Edge) error { + err := store.CreateEdge(ctx, &storage.Edge{ ID: uuid.New().String(), FromID: fromID, ToID: toID, diff --git a/compact/compact_regression_test.go b/compact/compact_regression_test.go new file mode 100644 index 0000000..1c40b85 --- /dev/null +++ b/compact/compact_regression_test.go @@ -0,0 +1,212 @@ +package compact + +import ( + "context" + "crypto/sha256" + "errors" + "fmt" + "strings" + "testing" + "time" + + "github.com/GrayCodeAI/yaad/storage" +) + +func seedMemoryNode(t *testing.T, store storage.Storage, typ, project string, confidence float64, accessCount int, extra string) *storage.Node { + t.Helper() + content := "memory content " + extra + hash := fmt.Sprintf("%x", sha256.Sum256([]byte(content))) + n := &storage.Node{ + ID: fmt.Sprintf("node-%s-%s-%s", project, typ, extra), + Type: typ, + Content: content, + ContentHash: hash, + Scope: "global", + Project: project, + Confidence: confidence, + AccessCount: accessCount, + CreatedAt: time.Now().Add(-time.Hour), + UpdatedAt: time.Now().Add(-time.Hour), + } + if err := store.CreateNode(context.Background(), n); err != nil { + t.Fatal(err) + } + return n +} + +func TestCompact_EmptyProjectRefused(t *testing.T) { + t.Parallel() + store := setupStore(t) + c := New(store, 100) + ctx := context.Background() + + for _, project := range []string{"", " "} { + n, err := c.Compact(ctx, project) + if err == nil { + t.Errorf("Compact(%q): expected error, got nil (compacted=%d)", project, n) + } + } +} + +func TestCompact_SkipsArchivedNodes(t *testing.T) { + t.Parallel() + store := setupStore(t) + c := New(store, 10000) + ctx := context.Background() + + seedMemoryNode(t, store, "idea", "p1", 0.4, 0, "live-1") + seedMemoryNode(t, store, "idea", "p1", 0.4, 0, "live-2") + seedMemoryNode(t, store, "idea", "p1", 0.4, 0, "live-3") + seedMemoryNode(t, store, "idea", "p1", 0, 0, "archived-1") + seedMemoryNode(t, store, "idea", "p1", 0, 0, "archived-2") + seedMemoryNode(t, store, "idea", "p1", 0, 0, "archived-3") + + n, err := c.Compact(ctx, "p1") + if err != nil { + t.Fatalf("Compact failed: %v", err) + } + if n != 3 { + t.Errorf("expected 3 archived (live) nodes compacted, got %d", n) + } + + // Only the live nodes collapsed into the summary. + nodes, err := store.ListNodes(ctx, storage.NodeFilter{Project: "p1"}) + if err != nil { + t.Fatal(err) + } + var summary *storage.Node + for _, node := range nodes { + if strings.HasPrefix(node.Content, "Summary of") { + summary = node + } + if node.Confidence == 0 { + continue // archived stays archived + } + if node.Confidence < 0.5 { + t.Errorf("live node %s was not archived but is still low-confidence after compact", node.ID) + } + } + if summary == nil { + t.Fatal("no summary node created") + } + if !strings.Contains(summary.Content, "Summary of 3 idea memories") { + t.Errorf("summary should mention only live nodes, got: %q", summary.Content) + } + if summary.Summary != "Compacted 3 idea memories" { + t.Errorf("summary field = %q, want %q", summary.Summary, "Compacted 3 idea memories") + } + if summary.Confidence != 0.6 { + t.Errorf("summary node confidence = %f, want 0.6", summary.Confidence) + } +} + +func TestCompact_SetsSummaryNodeTimestamps(t *testing.T) { + t.Parallel() + store := setupStore(t) + c := New(store, 10000) + ctx := context.Background() + + for i := 0; i < 3; i++ { + seedMemoryNode(t, store, "idea", "p1", 0.4, 0, fmt.Sprintf("live-%d", i)) + } + + before := time.Now() + if _, err := c.Compact(ctx, "p1"); err != nil { + t.Fatalf("Compact failed: %v", err) + } + + nodes, err := store.ListNodes(ctx, storage.NodeFilter{Project: "p1"}) + if err != nil { + t.Fatal(err) + } + for _, node := range nodes { + if node.Confidence != 0.6 { + continue // summary node only + } + if node.CreatedAt.IsZero() || node.UpdatedAt.IsZero() { + t.Errorf("summary node timestamps must be set, got CreatedAt=%v UpdatedAt=%v", node.CreatedAt, node.UpdatedAt) + } + if node.CreatedAt.Before(before) { + t.Errorf("summary node CreatedAt %v is before compaction started %v", node.CreatedAt, before) + } + } +} + +// failingStore fails every CreateEdge call after the first, to force a +// mid-transaction failure while re-linking edges. +type failingStore struct { + storage.Storage + calls int +} + +func (f *failingStore) CreateEdge(ctx context.Context, e *storage.Edge) error { + f.calls++ + if f.calls > 1 { + return errors.New("injected edge failure") + } + return f.Storage.CreateEdge(ctx, e) +} + +func (f *failingStore) WithTx(ctx context.Context, fn func(storage.Storage) error) error { + return f.Storage.WithTx(ctx, func(tx storage.Storage) error { + return fn(&failingStore{Storage: tx}) + }) +} + +func TestCompact_TransactionalRollback(t *testing.T) { + t.Parallel() + store := setupStore(t) + fs := &failingStore{Storage: store} + c := New(fs, 10000) + ctx := context.Background() + + n1 := seedMemoryNode(t, store, "idea", "p1", 0.4, 0, "live-1") + n2 := seedMemoryNode(t, store, "idea", "p1", 0.4, 0, "live-2") + n3 := seedMemoryNode(t, store, "idea", "p1", 0.4, 0, "live-3") + ext := seedMemoryNode(t, store, "entity", "p1", 1.0, 5, "external") + if err := store.CreateEdge(ctx, &storage.Edge{ + ID: "edge-1", FromID: n1.ID, ToID: ext.ID, Type: "mentions", + }); err != nil { + t.Fatal(err) + } + if err := store.CreateEdge(ctx, &storage.Edge{ + ID: "edge-2", FromID: n2.ID, ToID: ext.ID, Type: "mentions", + }); err != nil { + t.Fatal(err) + } + + if _, err := c.Compact(ctx, "p1"); err == nil { + t.Fatal("expected compaction to fail on injected edge error") + } + + // Nothing may have been committed: summary node absent, original nodes + // still live, original edges untouched. + if _, err := store.GetNode(ctx, n1.ID); err != nil { + t.Errorf("node n1 lost after rollback: %v", err) + } + for _, n := range []*storage.Node{n1, n2, n3} { + got, err := store.GetNode(ctx, n.ID) + if err != nil { + t.Fatalf("GetNode(%s) failed: %v", n.ID, err) + } + if got.Confidence != 0.4 { + t.Errorf("node %s archived despite rollback (confidence=%f)", n.ID, got.Confidence) + } + } + out, err := store.GetEdgesFrom(ctx, n1.ID) + if err != nil { + t.Fatal(err) + } + if len(out) != 1 || out[0].ToID != ext.ID { + t.Errorf("original edge n1->ext not preserved after rollback: %+v", out) + } + nodes, err := store.ListNodes(ctx, storage.NodeFilter{Project: "p1"}) + if err != nil { + t.Fatal(err) + } + for _, node := range nodes { + if strings.HasPrefix(node.Content, "Summary of") { + t.Errorf("summary node survived rollback: %q", node.Content) + } + } +} diff --git a/engine/engine_mock_storage_test.go b/engine/engine_mock_storage_test.go index 57ed212..76fc939 100644 --- a/engine/engine_mock_storage_test.go +++ b/engine/engine_mock_storage_test.go @@ -123,6 +123,22 @@ func (m *mockStorage) UpdateNodeContent(ctx context.Context, id, newContent stri return storage.ErrNodeNotFound } +func (m *mockStorage) ArchiveNode(ctx context.Context, id string) (bool, error) { + if err := ctx.Err(); err != nil { + return false, err + } + m.mu.Lock() + defer m.mu.Unlock() + n, ok := m.nodes[id] + if !ok || n.Confidence <= 0 { + return false, nil + } + cp := *n + cp.Confidence = 0 + m.nodes[id] = &cp + return true, nil +} + func (m *mockStorage) DeleteNode(ctx context.Context, id string) error { if err := ctx.Err(); err != nil { return err diff --git a/engine/forget_cas_test.go b/engine/forget_cas_test.go new file mode 100644 index 0000000..2e40840 --- /dev/null +++ b/engine/forget_cas_test.go @@ -0,0 +1,68 @@ +package engine + +import ( + "context" + "testing" +) + +// TestForget_PreservesConcurrentContentUpdate pins the CAS archival contract: +// Forget must archive (confidence=0) without writing back a stale full-node +// copy, so a content update that lands between Forget's read and write is +// never clobbered. +func TestForget_PreservesConcurrentContentUpdate(t *testing.T) { + t.Parallel() + eng := newTestEngine() + + node, err := eng.Remember(context.Background(), RememberInput{Type: "decision", Content: "original content", Project: "p1"}) + if err != nil { + t.Fatalf("Remember failed: %v", err) + } + + // Simulate an async writer updating content after Forget's read but + // before its write (the ingestion path does not take e.mu). + if err := eng.store.UpdateNodeContent(context.Background(), node.ID, "updated content"); err != nil { + t.Fatalf("UpdateNodeContent failed: %v", err) + } + + if err := eng.Forget(context.Background(), node.ID); err != nil { + t.Fatalf("Forget failed: %v", err) + } + + got, err := eng.store.GetNode(context.Background(), node.ID) + if err != nil { + t.Fatalf("GetNode failed: %v", err) + } + if got.Confidence != 0 { + t.Errorf("expected confidence 0 after forget, got %f", got.Confidence) + } + if got.Content != "updated content" { + t.Errorf("forget clobbered concurrent content update: got %q, want %q", got.Content, "updated content") + } +} + +// TestForget_TwiceIsNoOp ensures re-archiving an archived node is not an error +// and does not touch the node. +func TestForget_TwiceIsNoOp(t *testing.T) { + t.Parallel() + eng := newTestEngine() + + node, err := eng.Remember(context.Background(), RememberInput{Type: "decision", Content: "archive me", Project: "p1"}) + if err != nil { + t.Fatalf("Remember failed: %v", err) + } + + if err := eng.Forget(context.Background(), node.ID); err != nil { + t.Fatalf("first Forget failed: %v", err) + } + if err := eng.Forget(context.Background(), node.ID); err != nil { + t.Fatalf("second Forget must be a no-op, got: %v", err) + } + + got, err := eng.store.GetNode(context.Background(), node.ID) + if err != nil { + t.Fatalf("GetNode failed: %v", err) + } + if got.Confidence != 0 { + t.Errorf("expected confidence 0 after double forget, got %f", got.Confidence) + } +} diff --git a/engine/memory.go b/engine/memory.go index ca5e905..9a7949f 100644 --- a/engine/memory.go +++ b/engine/memory.go @@ -30,12 +30,18 @@ func (e *Engine) Forget(ctx context.Context, id string) error { e.mu.Unlock() return fmt.Errorf("save version failed: %w", err) } - node.Confidence = 0 - err = e.store.UpdateNode(ctx, node) + // Conditional write: archive only if the node is still live, so a + // concurrent content update is never clobbered by a stale archive copy + // (and a node already archived is not re-archived). e.mu.Unlock() + archived, err := e.store.ArchiveNode(ctx, node.ID) if err != nil { return err } + if !archived { + slog.Debug("yaad: forget: node already archived or not live", "node_id", node.ID) + return nil + } // Update the MEMORY.md hot-pointer file if configured (best-effort). if rerr := e.renderMemoryFile(ctx); rerr != nil { diff --git a/go.mod b/go.mod index 53f4801..a7af6d7 100644 --- a/go.mod +++ b/go.mod @@ -33,8 +33,8 @@ require ( go.opentelemetry.io/auto/sdk v1.2.1 // indirect go.opentelemetry.io/otel/sdk v1.44.0 // indirect go.opentelemetry.io/otel/trace v1.44.0 // indirect - golang.org/x/net v0.53.0 // indirect - golang.org/x/sys v0.45.0 // indirect + golang.org/x/net v0.57.0 // indirect + golang.org/x/sys v0.47.0 // indirect google.golang.org/genproto/googleapis/rpc v0.0.0-20260414002931-afd174a4e478 // indirect modernc.org/libc v1.72.5 // indirect modernc.org/mathutil v1.7.1 // indirect diff --git a/go.sum b/go.sum index 3c8ee4a..d3d6cbd 100644 --- a/go.sum +++ b/go.sum @@ -67,12 +67,12 @@ go.opentelemetry.io/otel/trace v1.44.0 h1:jxF5CsGYCe74MCRx2X4g7WsY/VBKRqqpNvXlX/ go.opentelemetry.io/otel/trace v1.44.0/go.mod h1:oLl1jrMQAVo6v3GAggN+1VH9VIz9iUSvW53sW1Q8PIE= golang.org/x/mod v0.37.0 h1:vF1DjpVEshcIqoEaauuHebaLk1O1forxjxBaVn884JQ= golang.org/x/mod v0.37.0/go.mod h1:m8S8VeM9r4dzDwjrKO0a1sZP3YjeMamRRlD+fmR2Q/0= -golang.org/x/net v0.53.0 h1:d+qAbo5L0orcWAr0a9JweQpjXF19LMXJE8Ey7hwOdUA= -golang.org/x/net v0.53.0/go.mod h1:JvMuJH7rrdiCfbeHoo3fCQU24Lf5JJwT9W3sJFulfgs= +golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE= +golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU= golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= -golang.org/x/sys v0.45.0 h1:dO4czNzziLiiXplLQgBCEpCvXQ3dnkn0SdaZSYdQ+FY= -golang.org/x/sys v0.45.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= +golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= +golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs= golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY= golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q= diff --git a/hooks/git_test.go b/hooks/git_test.go index 8deab6d..a4ce212 100644 --- a/hooks/git_test.go +++ b/hooks/git_test.go @@ -78,7 +78,22 @@ func TestBashHookRunner_PostExecution(t *testing.T) { if len(nodes) == 0 { t.Fatal("expected stored memory node from bash hook execution") } - if nodes[0].Type != "bug" { - t.Errorf("expected node type bug for non-zero exit code, got %s", nodes[0].Type) + // ListNodes is ordered by recency (updated_at DESC), so locate the bug + // node instead of assuming a position. + var found bool + for _, n := range nodes { + if n.Type == "bug" { + found = true + break + } + } + if !found { + t.Errorf("expected a bug node for non-zero exit code, got types %v", func() []string { + var types []string + for _, n := range nodes { + types = append(types, n.Type) + } + return types + }()) } } diff --git a/internal/server/mcp_concurrency_rest_test.go b/internal/server/mcp_concurrency_rest_test.go index 1f07cbe..7f6bad0 100644 --- a/internal/server/mcp_concurrency_rest_test.go +++ b/internal/server/mcp_concurrency_rest_test.go @@ -370,7 +370,16 @@ func TestMCPCompact(t *testing.T) { defer cleanup() ctx := context.Background() + + // Without a project the tool must refuse: compaction must never run + // across all projects at once. req := toolRequest("yaad_compact", map[string]any{}) + if _, err := srv.handleCompact(ctx, req); err == nil { + t.Fatal("expected error for yaad_compact without project") + } + + // With a project it compacts that project only. + req = toolRequest("yaad_compact", map[string]any{"project": "test-project"}) res, err := srv.handleCompact(ctx, req) if err != nil { t.Fatalf("handleCompact: %v", err) diff --git a/internal/server/mcp_memory.go b/internal/server/mcp_memory.go index c7d5669..809fa5d 100644 --- a/internal/server/mcp_memory.go +++ b/internal/server/mcp_memory.go @@ -5,6 +5,7 @@ import ( "encoding/json" "fmt" "os" + "strings" "time" mcpkit "github.com/GrayCodeAI/hawk-mcpkit" @@ -121,9 +122,14 @@ func (s *MCPServer) handlePin(ctx context.Context, req mcp.CallToolRequest) (*mc func (s *MCPServer) handleCompact(ctx context.Context, req mcp.CallToolRequest) (*mcp.CallToolResult, error) { token := req.Params.Meta.ProgressToken + project := mcpkit.StrArg(req, "project") + if strings.TrimSpace(project) == "" { + return nil, fmt.Errorf("%w: yaad_compact requires a 'project' argument; refusing to compact across all projects", mcp.ErrInvalidParams) + } + sendProgress(ctx, s.mcp(), token, 1, 3, "identifying low-confidence memories") - n, err := s.eng.Compact(ctx, mcpkit.StrArg(req, "project")) + n, err := s.eng.Compact(ctx, project) if err != nil { return nil, err } diff --git a/storage/interface.go b/storage/interface.go index e8ae537..2a784da 100644 --- a/storage/interface.go +++ b/storage/interface.go @@ -25,6 +25,10 @@ type Storage interface { GetNodesBatch(ctx context.Context, ids []string) ([]*Node, error) UpdateNode(ctx context.Context, n *Node) error UpdateNodeContent(ctx context.Context, id, newContent string) error + // ArchiveNode atomically sets confidence to 0 (the archived marker) but + // only if the node is still live (confidence > 0). It returns whether the + // node was archived; a node already archived is a no-op, not an error. + ArchiveNode(ctx context.Context, id string) (bool, error) DeleteNode(ctx context.Context, id string) error ListNodes(ctx context.Context, f NodeFilter) ([]*Node, error) SearchNodes(ctx context.Context, query string, limit int) ([]*Node, error) diff --git a/storage/interface_test.go b/storage/interface_test.go index 23ef8c4..e2cf48f 100644 --- a/storage/interface_test.go +++ b/storage/interface_test.go @@ -65,6 +65,18 @@ func (m *mockStorage) UpdateNodeContent(ctx context.Context, id, newContent stri return sql.ErrNoRows } +func (m *mockStorage) ArchiveNode(ctx context.Context, id string) (bool, error) { + if err := ctx.Err(); err != nil { + return false, err + } + n, ok := m.nodes[id] + if !ok || n.Confidence <= 0 { + return false, nil + } + n.Confidence = 0 + return true, nil +} + func (m *mockStorage) DeleteNode(ctx context.Context, id string) error { delete(m.nodes, id) return nil diff --git a/storage/mock.go b/storage/mock.go index 99a7abb..76e9a0a 100644 --- a/storage/mock.go +++ b/storage/mock.go @@ -162,6 +162,20 @@ func (m *MockStorage) UpdateNodeContent(_ context.Context, id, newContent string return sql.ErrNoRows } +func (m *MockStorage) ArchiveNode(_ context.Context, id string) (bool, error) { + m.mu.Lock() + defer m.mu.Unlock() + if err := m.err(); err != nil { + return false, err + } + n, ok := m.nodes[id] + if !ok || n.Confidence <= 0 { + return false, nil + } + n.Confidence = 0 + return true, nil +} + func (m *MockStorage) DeleteNode(_ context.Context, id string) error { m.mu.Lock() defer m.mu.Unlock() diff --git a/storage/order_archive_test.go b/storage/order_archive_test.go new file mode 100644 index 0000000..6416173 --- /dev/null +++ b/storage/order_archive_test.go @@ -0,0 +1,117 @@ +package storage + +import ( + "context" + "path/filepath" + "testing" + "time" +) + +func newOrderTestStore(t *testing.T) *Store { + t.Helper() + store, err := NewStore(filepath.Join(t.TempDir(), "test.db")) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = store.Close() }) + return store +} + +func seedNode(t *testing.T, store *Store, id, project string, updatedAt time.Time) { + t.Helper() + if err := store.CreateNode(context.Background(), &Node{ + ID: id, + Type: "idea", + Content: "content " + id, + ContentHash: "hash-" + id, + Scope: "global", + Project: project, + Confidence: 0.8, + CreatedAt: updatedAt, + UpdatedAt: updatedAt, + }); err != nil { + t.Fatal(err) + } +} + +// TestListNodes_OrdersByUpdatedAtDesc pins the deterministic recency order: +// list queries with a Limit must return the most recently updated nodes, and +// pagination must be stable. +func TestListNodes_OrdersByUpdatedAtDesc(t *testing.T) { + t.Parallel() + store := newOrderTestStore(t) + ctx := context.Background() + + base := time.Now().Add(-24 * time.Hour) + seedNode(t, store, "oldest", "p1", base) + seedNode(t, store, "middle", "p1", base.Add(1*time.Hour)) + seedNode(t, store, "newest", "p1", base.Add(2*time.Hour)) + + nodes, err := store.ListNodes(ctx, NodeFilter{Project: "p1", Limit: 2}) + if err != nil { + t.Fatal(err) + } + if len(nodes) != 2 { + t.Fatalf("expected 2 nodes, got %d", len(nodes)) + } + if nodes[0].ID != "newest" || nodes[1].ID != "middle" { + t.Errorf("expected [newest middle], got [%s %s]", nodes[0].ID, nodes[1].ID) + } + + // Page 2 continues deterministically. + page2, err := store.ListNodes(ctx, NodeFilter{Project: "p1", Limit: 1, Offset: 2}) + if err != nil { + t.Fatal(err) + } + if len(page2) != 1 || page2[0].ID != "oldest" { + t.Errorf("expected page 2 = [oldest], got %+v", page2) + } +} + +// TestArchiveNode_OnlyArchivesLiveNodes pins the CAS semantics behind Forget: +// archiving is a confidence>0-guarded single-column write, and re-archiving is +// an idempotent no-op rather than a full-node write-back. +func TestArchiveNode_OnlyArchivesLiveNodes(t *testing.T) { + t.Parallel() + store := newOrderTestStore(t) + ctx := context.Background() + + seedNode(t, store, "live", "p1", time.Now()) + seedNode(t, store, "gone", "p1", time.Now()) + + archived, err := store.ArchiveNode(ctx, "live") + if err != nil { + t.Fatal(err) + } + if !archived { + t.Error("expected live node to be archived") + } + got, err := store.GetNode(ctx, "live") + if err != nil { + t.Fatal(err) + } + if got.Confidence != 0 { + t.Errorf("confidence = %f, want 0", got.Confidence) + } + if got.Content != "content live" { + t.Errorf("archiving must not touch other fields, content = %q", got.Content) + } + + // Idempotent: already archived is a no-op, not an error. + again, err := store.ArchiveNode(ctx, "live") + if err != nil { + t.Fatal(err) + } + if again { + t.Error("expected already-archived node to report not archived") + } + + // Missing node: no-op, no error. + missing, err := store.ArchiveNode(ctx, "does-not-exist") + if err != nil { + t.Fatal(err) + } + if missing { + t.Error("expected missing node to report not archived") + } +} diff --git a/storage/sqlite_nodes.go b/storage/sqlite_nodes.go index 10d9327..90a8b1f 100644 --- a/storage/sqlite_nodes.go +++ b/storage/sqlite_nodes.go @@ -156,6 +156,23 @@ func updateNodeQ(ctx context.Context, q queryable, n *Node) error { return saveNodeMetadataQ(ctx, q, n.ID, n.Metadata) } +func (s *Store) ArchiveNode(ctx context.Context, id string) (bool, error) { + return retryOnBusyVal(func() (bool, error) { + ctx, cancel := s.withTimeout(ctx) + defer cancel() + return archiveNodeQ(ctx, s.q(), id) + }, 5, 50*time.Millisecond) +} + +func archiveNodeQ(ctx context.Context, q queryable, id string) (bool, error) { + res, err := q.ExecContext(ctx, `UPDATE nodes SET confidence=0, updated_at=CURRENT_TIMESTAMP WHERE id=? AND confidence>0`, id) + if err != nil { + return false, err + } + n, err := res.RowsAffected() + return n > 0, err +} + func (s *Store) UpdateNodeContent(ctx context.Context, id, newContent string) error { return retryOnBusy(func() error { ctx, cancel := s.withTimeout(ctx) @@ -250,6 +267,10 @@ func listNodesQ(ctx context.Context, q queryable, f NodeFilter) ([]*Node, error) query += ` AND EXISTS (SELECT 1 FROM node_metadata nm WHERE nm.node_id = nodes.id AND nm.key = ? AND nm.value = ?)` args = append(args, k, v) } + // Deterministic order: most recently updated first, then insertion order. + // Pagination (Limit/Offset) is only meaningful with a stable ORDER BY, and + // recent-window queries (e.g. conflict re-detection) rely on recency. + query += " ORDER BY updated_at DESC, rowid DESC" query += " LIMIT ?" if f.Limit > 0 { args = append(args, f.Limit) diff --git a/storage/sqlite_tx.go b/storage/sqlite_tx.go index fd0efbb..b587ca7 100644 --- a/storage/sqlite_tx.go +++ b/storage/sqlite_tx.go @@ -62,6 +62,10 @@ func (t *txStore) UpdateNodeContent(ctx context.Context, id, newContent string) return updateNodeContentQ(ctx, t.tx, id, newContent) } +func (t *txStore) ArchiveNode(ctx context.Context, id string) (bool, error) { + return archiveNodeQ(ctx, t.tx, id) +} + func (t *txStore) DeleteNode(ctx context.Context, id string) error { return deleteNodeQ(ctx, t.tx, id) } func (t *txStore) ListNodes(ctx context.Context, f NodeFilter) ([]*Node, error) {