From 2dd54a82696fa912be5631428b4eefafa8800671 Mon Sep 17 00:00:00 2001 From: ZanzyTHEbar Date: Sun, 22 Feb 2026 22:48:36 +0000 Subject: [PATCH] refactor(memory): delegate sqlite, store, retrieval policy updates --- pkg/memory/delegate/sqlite.go | 78 +++++++++- pkg/memory/memory.go | 81 +++++----- pkg/memory/migrate_sessions_test.go | 9 ++ pkg/memory/store/memory_store.go | 207 +++++++++++++++++++++----- pkg/memory/store/memory_store_test.go | 10 +- pkg/memory/store/retrieval_policy.go | 23 +-- 6 files changed, 318 insertions(+), 90 deletions(-) diff --git a/pkg/memory/delegate/sqlite.go b/pkg/memory/delegate/sqlite.go index f7f694de0..52633288e 100644 --- a/pkg/memory/delegate/sqlite.go +++ b/pkg/memory/delegate/sqlite.go @@ -13,7 +13,6 @@ import ( "github.com/ZanzyTHEbar/dragonscale/pkg/memory/migrations" memsqlc "github.com/ZanzyTHEbar/dragonscale/pkg/memory/sqlc" "github.com/pressly/goose/v3" - libsql "github.com/tursodatabase/go-libsql" ) @@ -268,6 +267,24 @@ func (d *LibSQLDelegate) GetRecallItem(ctx context.Context, agentID string, id i return sqlcRecallToMemory(row), nil } +func (d *LibSQLDelegate) GetRecallItemsByIDs(ctx context.Context, agentID string, itemIDs []ids.UUID) (map[ids.UUID]*memory.RecallItem, error) { + if len(itemIDs) == 0 { + return make(map[ids.UUID]*memory.RecallItem), nil + } + rows, err := d.queries.GetRecallItemsByIDs(ctx, memsqlc.GetRecallItemsByIDsParams{ + Ids: itemIDs, + AgentID: agentID, + }) + if err != nil { + return nil, err + } + result := make(map[ids.UUID]*memory.RecallItem, len(rows)) + for _, row := range rows { + result[row.ID] = sqlcRecallToMemory(row) + } + return result, nil +} + func (d *LibSQLDelegate) UpdateRecallItem(ctx context.Context, item *memory.RecallItem) error { return d.queries.UpdateRecallItem(ctx, memsqlc.UpdateRecallItemParams{ ID: item.ID, @@ -344,6 +361,31 @@ func archivalChunkToParams(chunk *memory.ArchivalChunk) memsqlc.InsertArchivalCh } } +func (d *LibSQLDelegate) InsertArchivalChunkBatch(ctx context.Context, chunks []*memory.ArchivalChunk) error { + if len(chunks) == 0 { + return nil + } + if len(chunks) == 1 { + return d.InsertArchivalChunk(ctx, chunks[0]) + } + + tx, err := d.db.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("begin tx: %w", err) + } + defer tx.Rollback() + + qtx := d.queries.WithTx(tx) + for _, chunk := range chunks { + row, err := qtx.InsertArchivalChunk(ctx, archivalChunkToParams(chunk)) + if err != nil { + return err + } + chunk.CreatedAt = row.CreatedAt + } + return tx.Commit() +} + func (d *LibSQLDelegate) GetArchivalChunk(ctx context.Context, agentID string, id ids.UUID) (*memory.ArchivalChunk, error) { row, err := d.queries.GetArchivalChunk(ctx, memsqlc.GetArchivalChunkParams{ID: id, AgentID: agentID}) if err == sql.ErrNoRows { @@ -621,6 +663,40 @@ func (d *LibSQLDelegate) InsertAuditEntry(ctx context.Context, entry *memory.Aud return nil } +func (d *LibSQLDelegate) InsertAuditEntryBatch(ctx context.Context, entries []*memory.AuditEntry) error { + if len(entries) == 0 { + return nil + } + if len(entries) == 1 { + return d.InsertAuditEntry(ctx, entries[0]) + } + + tx, err := d.db.BeginTx(ctx, nil) + if err != nil { + return fmt.Errorf("begin tx: %w", err) + } + defer tx.Rollback() + + qtx := d.queries.WithTx(tx) + for _, entry := range entries { + row, err := qtx.InsertAuditEntry(ctx, memsqlc.InsertAuditEntryParams{ + ID: entry.ID, + AgentID: entry.AgentID, + SessionKey: entry.SessionKey, + Action: entry.Action, + Target: entry.Target, + Input: &entry.Input, + Output: &entry.Output, + DurationMs: ptrInt64(int64(entry.DurationMS)), + }) + if err != nil { + return err + } + entry.CreatedAt = row.CreatedAt + } + return tx.Commit() +} + func (d *LibSQLDelegate) ListAuditEntries(ctx context.Context, agentID string, limit int) ([]*memory.AuditEntry, error) { rows, err := d.queries.ListAuditEntries(ctx, memsqlc.ListAuditEntriesParams{ AgentID: agentID, diff --git a/pkg/memory/memory.go b/pkg/memory/memory.go index ab639ca58..4d8ffc3f7 100644 --- a/pkg/memory/memory.go +++ b/pkg/memory/memory.go @@ -244,74 +244,67 @@ type AuditEntry struct { // MemoryDelegate is the pure storage backend for the memory system. // Implementations wrap sqlc-generated queries. All persistence goes through here. // The Memory logic layer composes a MemoryDelegate for its backend. -type MemoryDelegate interface { - // Init creates tables and runs migrations. - Init(ctx context.Context) error - - // Close releases database resources. - Close() error - - // --- Working Context --- +// MemoryReader groups all read-only methods on the memory store. Enables +// independent optimization of the read path (caching, read replicas). +type MemoryReader interface { GetWorkingContext(ctx context.Context, agentID, sessionKey string) (*WorkingContext, error) - UpsertWorkingContext(ctx context.Context, agentID, sessionKey, content string) error - - // --- Recall Items --- - InsertRecallItem(ctx context.Context, item *RecallItem) error GetRecallItem(ctx context.Context, agentID string, id ids.UUID) (*RecallItem, error) - UpdateRecallItem(ctx context.Context, item *RecallItem) error - DeleteRecallItem(ctx context.Context, agentID string, id ids.UUID) error + GetRecallItemsByIDs(ctx context.Context, agentID string, itemIDs []ids.UUID) (map[ids.UUID]*RecallItem, error) ListRecallItems(ctx context.Context, agentID, sessionKey string, limit, offset int) ([]*RecallItem, error) SearchRecallByKeyword(ctx context.Context, query, agentID string, limit int) ([]*RecallItem, error) - - // --- Advanced Search --- - - // SearchRecallByFTS performs full-text search using FTS5 MATCH with BM25 ranking. - // Returns nil (not error) if FTS5 is not available -- caller should fall back to keyword search. SearchRecallByFTS(ctx context.Context, query, agentID string, limit int) ([]*RecallItem, error) - - // SearchArchivalByVector performs DB-side vector similarity search. - // Returns nil (not error) if vector search is not available -- caller should fall back to Go-side. SearchArchivalByVector(ctx context.Context, queryVec Embedding, limit, offset int) ([]SearchResult, error) - - // --- Archival Chunks --- - InsertArchivalChunk(ctx context.Context, chunk *ArchivalChunk) error GetArchivalChunk(ctx context.Context, agentID string, id ids.UUID) (*ArchivalChunk, error) ListArchivalChunks(ctx context.Context, agentID string, recallID ids.UUID) ([]*ArchivalChunk, error) ListAllArchivalChunks(ctx context.Context, agentID string, limit, offset int) ([]*ArchivalChunk, error) - DeleteArchivalChunks(ctx context.Context, recallID ids.UUID) error - - // --- Summaries --- - InsertSummary(ctx context.Context, summary *MemorySummary) error ListSummaries(ctx context.Context, agentID, sessionKey string, limit int) ([]*MemorySummary, error) - - // --- Stats --- CountRecallItems(ctx context.Context, agentID, sessionKey string) (int, error) CountArchivalChunks(ctx context.Context, agentID string) (int, error) - - // --- Key-Value Store --- GetKV(ctx context.Context, agentID, key string) (string, error) - UpsertKV(ctx context.Context, agentID, key, value string) error - DeleteKV(ctx context.Context, agentID, key string) error ListKVByPrefix(ctx context.Context, agentID, prefix string, limit int) (map[string]string, error) - - // --- Documents --- GetDocument(ctx context.Context, agentID, name string) (*AgentDocument, error) - UpsertDocument(ctx context.Context, doc *AgentDocument) error - DeleteDocument(ctx context.Context, agentID, name string) error ListDocumentsByCategory(ctx context.Context, agentID, category string) ([]*AgentDocument, error) ListAllDocuments(ctx context.Context, agentID string) ([]*AgentDocument, error) - - // --- Audit Log --- - InsertAuditEntry(ctx context.Context, entry *AuditEntry) error ListAuditEntries(ctx context.Context, agentID string, limit int) ([]*AuditEntry, error) ListAuditEntriesByAction(ctx context.Context, agentID, action string, limit int) ([]*AuditEntry, error) CountAuditEntries(ctx context.Context, agentID string) (int, error) - - // --- Capability Detection --- HasVectorSearch() bool HasFTS() bool } +// MemoryWriter groups all mutating methods on the memory store. Enables +// write batching, buffering, and independent scaling of the write path. +type MemoryWriter interface { + UpsertWorkingContext(ctx context.Context, agentID, sessionKey, content string) error + InsertRecallItem(ctx context.Context, item *RecallItem) error + UpdateRecallItem(ctx context.Context, item *RecallItem) error + DeleteRecallItem(ctx context.Context, agentID string, id ids.UUID) error + InsertArchivalChunk(ctx context.Context, chunk *ArchivalChunk) error + InsertArchivalChunkBatch(ctx context.Context, chunks []*ArchivalChunk) error + DeleteArchivalChunks(ctx context.Context, recallID ids.UUID) error + InsertSummary(ctx context.Context, summary *MemorySummary) error + UpsertKV(ctx context.Context, agentID, key, value string) error + DeleteKV(ctx context.Context, agentID, key string) error + UpsertDocument(ctx context.Context, doc *AgentDocument) error + DeleteDocument(ctx context.Context, agentID, name string) error + InsertAuditEntry(ctx context.Context, entry *AuditEntry) error + InsertAuditEntryBatch(ctx context.Context, entries []*AuditEntry) error +} + +// MemoryDelegate is the full-capability interface for memory operations. +// Embeds MemoryReader and MemoryWriter for CQRS-compatible usage: callers +// that only need reads can accept MemoryReader, enabling independent +// optimization, caching, and testing of each path. +type MemoryDelegate interface { + MemoryReader + MemoryWriter + + // Init creates tables and runs migrations. + Init(ctx context.Context) error + // Close releases database resources. + Close() error +} + // --- Embedding interface --- // EmbeddingProvider generates vector embeddings from text. diff --git a/pkg/memory/migrate_sessions_test.go b/pkg/memory/migrate_sessions_test.go index 55e92d3e6..b6672706d 100644 --- a/pkg/memory/migrate_sessions_test.go +++ b/pkg/memory/migrate_sessions_test.go @@ -43,6 +43,9 @@ func (m *mockDelegate) InsertRecallItem(_ context.Context, item *RecallItem) err func (m *mockDelegate) GetRecallItem(_ context.Context, _ string, _ ids.UUID) (*RecallItem, error) { return nil, nil } +func (m *mockDelegate) GetRecallItemsByIDs(_ context.Context, _ string, _ []ids.UUID) (map[ids.UUID]*RecallItem, error) { + return make(map[ids.UUID]*RecallItem), nil +} func (m *mockDelegate) UpdateRecallItem(_ context.Context, _ *RecallItem) error { return nil } func (m *mockDelegate) DeleteRecallItem(_ context.Context, _ string, _ ids.UUID) error { return nil } func (m *mockDelegate) ListRecallItems(_ context.Context, _, _ string, _, _ int) ([]*RecallItem, error) { @@ -58,6 +61,9 @@ func (m *mockDelegate) SearchArchivalByVector(_ context.Context, _ Embedding, _, return nil, nil } func (m *mockDelegate) InsertArchivalChunk(_ context.Context, _ *ArchivalChunk) error { return nil } +func (m *mockDelegate) InsertArchivalChunkBatch(_ context.Context, _ []*ArchivalChunk) error { + return nil +} func (m *mockDelegate) GetArchivalChunk(_ context.Context, _ string, _ ids.UUID) (*ArchivalChunk, error) { return nil, nil } @@ -96,6 +102,9 @@ func (m *mockDelegate) ListAllDocuments(_ context.Context, _ string) ([]*AgentDo return nil, nil } func (m *mockDelegate) InsertAuditEntry(_ context.Context, _ *AuditEntry) error { return nil } +func (m *mockDelegate) InsertAuditEntryBatch(_ context.Context, _ []*AuditEntry) error { + return nil +} func (m *mockDelegate) ListAuditEntries(_ context.Context, _ string, _ int) ([]*AuditEntry, error) { return nil, nil } diff --git a/pkg/memory/store/memory_store.go b/pkg/memory/store/memory_store.go index a778c29e8..d0dc3ddc7 100644 --- a/pkg/memory/store/memory_store.go +++ b/pkg/memory/store/memory_store.go @@ -44,6 +44,30 @@ type MemoryStore struct { cfg Config agentID string retrievalPolicyMu sync.Mutex + + policyCache policyCacheState + vecCache vectorCache +} + +// vectorCache holds an in-memory copy of archival chunk embeddings so that +// vectorSearchGoSide doesn't re-scan the entire DB on every search call. +// Populated lazily on first search; updated incrementally in StoreArchival. +type vectorCache struct { + mu sync.RWMutex + items []VectorSearchInput + loaded bool +} + +// policyCacheState holds in-memory copies of retrieval policy data to avoid +// 5 KV round-trips per search call. Flushed on Sync/Close or every N queries. +type policyCacheState struct { + state retrievalPolicyState + gates retrievalPromotionGates + metrics retrievalShadowMetrics + loaded bool + dirty bool + flushEvery int + queryCount int } // New creates a MemoryStore. @@ -63,12 +87,59 @@ func New(delegate memory.MemoryDelegate, chunker memory.Chunker, embedder memory embedder: embedder, chunker: chunker, cfg: cfg, + policyCache: policyCacheState{ + flushEvery: 10, + }, } } // SetAgentID sets the agent identity used to scope all memory operations. +// Invalidates the vector cache since chunks are agent-scoped. func (m *MemoryStore) SetAgentID(agentID string) { m.agentID = agentID + m.invalidateVecCache() +} + +// invalidateVecCache forces the next vector search to reload from DB. +func (m *MemoryStore) invalidateVecCache() { + m.vecCache.mu.Lock() + m.vecCache.items = nil + m.vecCache.loaded = false + m.vecCache.mu.Unlock() +} + +// ensurePolicyCacheLoaded initializes the in-memory retrieval policy cache +// from KV on the first access. Must be called with retrievalPolicyMu held. +func (m *MemoryStore) ensurePolicyCacheLoaded(ctx context.Context) { + if m.policyCache.loaded { + return + } + m.policyCache.state = m.loadRetrievalPolicyState(ctx) + m.policyCache.gates = m.loadRetrievalPromotionGates(ctx) + m.policyCache.metrics = m.loadRetrievalShadowMetrics(ctx) + m.policyCache.loaded = true + m.policyCache.dirty = false + m.policyCache.queryCount = 0 +} + +// flushPolicyCacheLocked persists cached policy state to KV. +// Must be called with retrievalPolicyMu held. +func (m *MemoryStore) flushPolicyCacheLocked(ctx context.Context) { + if !m.policyCache.dirty { + return + } + _ = m.persistRetrievalPolicyState(ctx, m.policyCache.state) + _ = m.persistRetrievalShadowMetrics(ctx, m.policyCache.metrics) + m.policyCache.dirty = false + m.policyCache.queryCount = 0 +} + +// FlushPolicyCache persists any dirty retrieval policy state. Safe to call +// from Sync/Close paths. +func (m *MemoryStore) FlushPolicyCache(ctx context.Context) { + m.retrievalPolicyMu.Lock() + defer m.retrievalPolicyMu.Unlock() + m.flushPolicyCacheLocked(ctx) } // --- Working Context (hot tier) --- @@ -109,6 +180,7 @@ func (m *MemoryStore) DeleteRecall(ctx context.Context, id ids.UUID) error { if err := m.delegate.DeleteArchivalChunks(ctx, id); err != nil { return fmt.Errorf("delete archival chunks: %w", err) } + m.invalidateVecCache() return m.delegate.DeleteRecallItem(ctx, m.agentID, id) } @@ -159,13 +231,14 @@ func (m *MemoryStore) StoreArchival(ctx context.Context, content, source string, } } - // Store each chunk + // Build all chunks, then batch-insert in a single transaction. + stored := make([]*memory.ArchivalChunk, len(chunks)) for i, chunk := range chunks { var emb memory.Embedding if i < len(embeddings) { emb = embeddings[i] } - archChunk := &memory.ArchivalChunk{ + stored[i] = &memory.ArchivalChunk{ ID: ids.New(), RecallID: recallID, ChunkIndex: chunk.Index, @@ -174,11 +247,13 @@ func (m *MemoryStore) StoreArchival(ctx context.Context, content, source string, Source: source, Hash: hashContent(chunk.Text), } - if err := m.delegate.InsertArchivalChunk(ctx, archChunk); err != nil { - return recallID, fmt.Errorf("insert chunk %d: %w", i, err) - } } + if err := m.delegate.InsertArchivalChunkBatch(ctx, stored); err != nil { + return recallID, fmt.Errorf("insert chunks: %w", err) + } + + m.appendToVecCache(stored) return recallID, nil } @@ -292,11 +367,15 @@ func (m *MemoryStore) Search(ctx context.Context, query string, opts memory.Sear } m.retrievalPolicyMu.Lock() - state := m.loadRetrievalPolicyState(ctx) - gates := m.loadRetrievalPromotionGates(ctx) - metrics := m.loadRetrievalShadowMetrics(ctx) + m.ensurePolicyCacheLoaded(ctx) parity := computeRetrievalParity(baselineFinal, augmentedFinal, limit) - state = m.updateRetrievalPolicy(ctx, state, gates, metrics, parity, augmentedUsed, baselineFinal) + m.policyCache.state, m.policyCache.metrics = m.updateRetrievalPolicy(ctx, m.policyCache.state, m.policyCache.gates, m.policyCache.metrics, parity, augmentedUsed, baselineFinal) + m.policyCache.dirty = true + m.policyCache.queryCount++ + if m.policyCache.flushEvery > 0 && m.policyCache.queryCount >= m.policyCache.flushEvery { + m.flushPolicyCacheLocked(ctx) + } + state := m.policyCache.state m.retrievalPolicyMu.Unlock() switch state.Mode { @@ -325,20 +404,28 @@ func (m *MemoryStore) postProcessMergedResults( halfLife = m.cfg.DefaultHalfLifeHours } - // Build a createdAt lookup from recall items - createdAtMap := make(map[ids.UUID]time.Time) + // Batch-fetch all recall items referenced by merged results (single query). + recallIDs := make([]ids.UUID, 0, len(merged)) for _, r := range merged { - item, err := m.delegate.GetRecallItem(ctx, m.agentID, r.ID) - if err == nil && item != nil { - createdAtMap[r.ID] = item.CreatedAt + if !r.ID.IsZero() && !strings.HasPrefix(r.Source, "working-context:") && !strings.HasPrefix(r.Source, "dag:") { + recallIDs = append(recallIDs, r.ID) } } + recallBatch, batchErr := m.delegate.GetRecallItemsByIDs(ctx, m.agentID, recallIDs) + if batchErr != nil { + recallBatch = make(map[ids.UUID]*memory.RecallItem) + } + + createdAtMap := make(map[ids.UUID]time.Time, len(recallBatch)) + for id, item := range recallBatch { + createdAtMap[id] = item.CreatedAt + } ApplyRecencyDecay(merged, time.Now(), halfLife, func(id ids.UUID) time.Time { return createdAtMap[id] }) // 2. Metadata pre-filtering (sectors, session_key, date range) - merged = m.applyMetadataFilters(ctx, merged, opts) + merged = m.applyMetadataFiltersBatch(ctx, merged, opts, recallBatch) // 3. Filter by min score if opts.MinScore > 0 { @@ -433,10 +520,10 @@ func (m *MemoryStore) hybridProjectionSearch(ctx context.Context, query string, return results, nil } -// applyMetadataFilters removes results that don't match the requested sector, -// session_key, or date range constraints. It fetches recall item metadata -// from the delegate as needed. -func (m *MemoryStore) applyMetadataFilters(ctx context.Context, results []memory.SearchResult, opts memory.SearchOptions) []memory.SearchResult { +// applyMetadataFiltersBatch removes results that don't match the requested sector, +// session_key, or date range constraints. Uses a pre-fetched batch of recall items +// to avoid N+1 queries. +func (m *MemoryStore) applyMetadataFiltersBatch(ctx context.Context, results []memory.SearchResult, opts memory.SearchOptions, batch map[ids.UUID]*memory.RecallItem) []memory.SearchResult { needSectorFilter := len(opts.Sectors) > 0 needSessionFilter := opts.SessionKey != "" needDateFilter := opts.DateAfter != nil || opts.DateBefore != nil @@ -445,7 +532,6 @@ func (m *MemoryStore) applyMetadataFilters(ctx context.Context, results []memory return results } - // Build sector lookup set sectorSet := make(map[memory.Sector]bool, len(opts.Sectors)) for _, s := range opts.Sectors { sectorSet[s] = true @@ -453,8 +539,6 @@ func (m *MemoryStore) applyMetadataFilters(ctx context.Context, results []memory filtered := results[:0] for _, r := range results { - // Synthetic projection entries (working-context + DAG summaries) are - // already scoped by session and do not map to recall item IDs. if strings.HasPrefix(r.Source, "working-context:") || strings.HasPrefix(r.Source, "dag:") { if needSessionFilter && !strings.Contains(r.Source, opts.SessionKey) { continue @@ -462,8 +546,6 @@ func (m *MemoryStore) applyMetadataFilters(ctx context.Context, results []memory if needSectorFilter && !sectorSet[r.Sector] { continue } - // Synthetic entries do not carry durable timestamps, so honor - // explicit date filters conservatively by excluding them. if needDateFilter { continue } @@ -471,23 +553,20 @@ func (m *MemoryStore) applyMetadataFilters(ctx context.Context, results []memory continue } - item, err := m.delegate.GetRecallItem(ctx, m.agentID, r.ID) - if err != nil || item == nil { + item := batch[r.ID] + if item == nil { continue } if needSessionFilter && item.SessionKey != opts.SessionKey { continue } - if needSectorFilter && !sectorSet[item.Sector] { continue } - if opts.DateAfter != nil && item.CreatedAt.Before(*opts.DateAfter) { continue } - if opts.DateBefore != nil && item.CreatedAt.After(*opts.DateBefore) { continue } @@ -533,12 +612,41 @@ func (m *MemoryStore) vectorSearch(ctx context.Context, query string, opts memor return m.vectorSearchGoSide(ctx, queryVec, limit) } -// vectorSearchGoSide performs Go-side brute-force vector search as a fallback. +// vectorSearchGoSide performs Go-side brute-force cosine similarity search. +// Uses an in-memory cache of embeddings (populated lazily, updated on insert) +// to avoid re-scanning the DB on every search call. func (m *MemoryStore) vectorSearchGoSide(ctx context.Context, queryVec memory.Embedding, limit int) ([]memory.SearchResult, error) { + items, err := m.ensureVecCacheLoaded(ctx) + if err != nil { + return nil, err + } + return VectorSearch(queryVec, items, limit), nil +} + +const maxVectorSearchChunks = 50_000 + +// ensureVecCacheLoaded populates the vector cache from DB on first access. +func (m *MemoryStore) ensureVecCacheLoaded(ctx context.Context) ([]VectorSearchInput, error) { + m.vecCache.mu.RLock() + if m.vecCache.loaded { + items := m.vecCache.items + m.vecCache.mu.RUnlock() + return items, nil + } + m.vecCache.mu.RUnlock() + + m.vecCache.mu.Lock() + defer m.vecCache.mu.Unlock() + + // Double-check after acquiring write lock + if m.vecCache.loaded { + return m.vecCache.items, nil + } + var allChunks []*memory.ArchivalChunk offset := 0 batchSize := 5000 - for { + for len(allChunks) < maxVectorSearchChunks { batch, err := m.delegate.ListAllArchivalChunks(ctx, m.agentID, batchSize, offset) if err != nil { return nil, err @@ -549,19 +657,46 @@ func (m *MemoryStore) vectorSearchGoSide(ctx context.Context, queryVec memory.Em } offset += batchSize } + if len(allChunks) > maxVectorSearchChunks { + allChunks = allChunks[:maxVectorSearchChunks] + } - inputs := make([]VectorSearchInput, 0, len(allChunks)) + items := make([]VectorSearchInput, 0, len(allChunks)) for _, chunk := range allChunks { if len(chunk.Embedding) == 0 { continue } - inputs = append(inputs, VectorSearchInput{ + items = append(items, VectorSearchInput{ Chunk: chunk, Embedding: chunk.Embedding, }) } - return VectorSearch(queryVec, inputs, limit), nil + m.vecCache.items = items + m.vecCache.loaded = true + return items, nil +} + +// appendToVecCache adds newly inserted chunks to the in-memory vector cache +// without triggering a full reload. +func (m *MemoryStore) appendToVecCache(chunks []*memory.ArchivalChunk) { + m.vecCache.mu.Lock() + defer m.vecCache.mu.Unlock() + if !m.vecCache.loaded { + return // not yet populated; first search will load everything + } + for _, chunk := range chunks { + if len(chunk.Embedding) == 0 { + continue + } + if len(m.vecCache.items) >= maxVectorSearchChunks { + break + } + m.vecCache.items = append(m.vecCache.items, VectorSearchInput{ + Chunk: chunk, + Embedding: chunk.Embedding, + }) + } } // recallItemsToResults converts delegate recall items into search results. @@ -671,8 +806,9 @@ func (m *MemoryStore) OffloadToolResult(ctx context.Context, toolName, content, // --- Lifecycle --- // Sync flushes pending writes to the remote replica (Turso). -// No-op if the underlying delegate doesn't support replication. +// Also flushes in-memory caches (retrieval policy) to KV. func (m *MemoryStore) Sync() error { + m.FlushPolicyCache(context.Background()) type syncer interface{ Sync() error } if s, ok := m.delegate.(syncer); ok { return s.Sync() @@ -681,6 +817,7 @@ func (m *MemoryStore) Sync() error { } func (m *MemoryStore) Close() error { + m.FlushPolicyCache(context.Background()) return m.delegate.Close() } diff --git a/pkg/memory/store/memory_store_test.go b/pkg/memory/store/memory_store_test.go index aa8fb78fa..28baac798 100644 --- a/pkg/memory/store/memory_store_test.go +++ b/pkg/memory/store/memory_store_test.go @@ -318,6 +318,8 @@ func TestSearch_ShadowModeUsesBaselineAndTracksParity(t *testing.T) { // Shadow mode should keep baseline result ordering for production output. assert.NotContains(t, results[0].Source, "working-context:s1") + store.FlushPolicyCache(ctx) + metricsRaw, err := store.delegate.GetKV(ctx, "agent-1", retrievalPolicyMetricsKey) require.NoError(t, err) require.NotEmpty(t, metricsRaw) @@ -363,6 +365,8 @@ func TestSearch_DoesNotPromoteWithoutAugmentedSignals(t *testing.T) { }) require.NoError(t, err) + store.FlushPolicyCache(ctx) + stateRaw, err := store.delegate.GetKV(ctx, "agent-1", retrievalPolicyStateKey) require.NoError(t, err) var state retrievalPolicyState @@ -407,6 +411,8 @@ func TestSearch_PromoteOnlyOnGateWin(t *testing.T) { }) require.NoError(t, err) + store.FlushPolicyCache(ctx) + stateRaw, err := store.delegate.GetKV(ctx, "agent-1", retrievalPolicyStateKey) require.NoError(t, err) require.NotEmpty(t, stateRaw) @@ -538,7 +544,7 @@ func TestUpdateRetrievalPolicy_PersistFailuresDoNotBlockTransitions(t *testing.T {ID: ids.New(), Content: "baseline"}, } - next := store.updateRetrievalPolicy(ctx, state, gates, metrics, parity, true, baseline) + next, _ := store.updateRetrievalPolicy(ctx, state, gates, metrics, parity, true, baseline) assert.Empty(t, cmp.Diff(retrievalModePromoted, next.Mode)) } @@ -582,6 +588,8 @@ func TestSearch_ConcurrentRetrievalPolicyUpdates(t *testing.T) { require.NoError(t, err) } + store.FlushPolicyCache(ctx) + metricsRaw, err := store.delegate.GetKV(ctx, "agent-1", retrievalPolicyMetricsKey) require.NoError(t, err) var metrics retrievalShadowMetrics diff --git a/pkg/memory/store/retrieval_policy.go b/pkg/memory/store/retrieval_policy.go index b5e79b8da..affadf155 100644 --- a/pkg/memory/store/retrieval_policy.go +++ b/pkg/memory/store/retrieval_policy.go @@ -285,7 +285,7 @@ func (m *MemoryStore) updateRetrievalPolicy( parity retrievalParity, augmentedUsed bool, baseline []memory.SearchResult, -) retrievalPolicyState { +) (retrievalPolicyState, retrievalShadowMetrics) { prevMode := state.Mode metrics.TotalQueries++ @@ -369,15 +369,20 @@ func (m *MemoryStore) updateRetrievalPolicy( }) } - if err := m.persistRetrievalShadowMetrics(ctx, metrics); err != nil { - logger.WarnCF("memory", "failed to persist retrieval shadow metrics", - map[string]interface{}{"mode": state.Mode, "error": err.Error()}) + // State and metrics are written back to the in-memory policy cache by the + // caller; the cache handles periodic flush to KV. On mode transitions we + // flush immediately to guarantee durability. + if prevMode != state.Mode { + if err := m.persistRetrievalShadowMetrics(ctx, metrics); err != nil { + logger.WarnCF("memory", "failed to persist retrieval shadow metrics", + map[string]interface{}{"mode": state.Mode, "error": err.Error()}) + } + if err := m.persistRetrievalPolicyState(ctx, state); err != nil { + logger.WarnCF("memory", "failed to persist retrieval policy state", + map[string]interface{}{"mode": state.Mode, "error": err.Error()}) + } } - if err := m.persistRetrievalPolicyState(ctx, state); err != nil { - logger.WarnCF("memory", "failed to persist retrieval policy state", - map[string]interface{}{"mode": state.Mode, "error": err.Error()}) - } - return state + return state, metrics } func safeRate(numerator, denominator int) float64 {