refactor(memory): delegate sqlite, store, retrieval policy updates

This commit is contained in:
ZanzyTHEbar 2026-02-22 22:48:36 +00:00
parent 7f1400f6c0
commit 2dd54a8269
6 changed files with 318 additions and 90 deletions

View file

@ -13,7 +13,6 @@ import (
"github.com/ZanzyTHEbar/dragonscale/pkg/memory/migrations" "github.com/ZanzyTHEbar/dragonscale/pkg/memory/migrations"
memsqlc "github.com/ZanzyTHEbar/dragonscale/pkg/memory/sqlc" memsqlc "github.com/ZanzyTHEbar/dragonscale/pkg/memory/sqlc"
"github.com/pressly/goose/v3" "github.com/pressly/goose/v3"
libsql "github.com/tursodatabase/go-libsql" 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 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 { func (d *LibSQLDelegate) UpdateRecallItem(ctx context.Context, item *memory.RecallItem) error {
return d.queries.UpdateRecallItem(ctx, memsqlc.UpdateRecallItemParams{ return d.queries.UpdateRecallItem(ctx, memsqlc.UpdateRecallItemParams{
ID: item.ID, 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) { 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}) row, err := d.queries.GetArchivalChunk(ctx, memsqlc.GetArchivalChunkParams{ID: id, AgentID: agentID})
if err == sql.ErrNoRows { if err == sql.ErrNoRows {
@ -621,6 +663,40 @@ func (d *LibSQLDelegate) InsertAuditEntry(ctx context.Context, entry *memory.Aud
return nil 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) { func (d *LibSQLDelegate) ListAuditEntries(ctx context.Context, agentID string, limit int) ([]*memory.AuditEntry, error) {
rows, err := d.queries.ListAuditEntries(ctx, memsqlc.ListAuditEntriesParams{ rows, err := d.queries.ListAuditEntries(ctx, memsqlc.ListAuditEntriesParams{
AgentID: agentID, AgentID: agentID,

View file

@ -244,74 +244,67 @@ type AuditEntry struct {
// MemoryDelegate is the pure storage backend for the memory system. // MemoryDelegate is the pure storage backend for the memory system.
// Implementations wrap sqlc-generated queries. All persistence goes through here. // Implementations wrap sqlc-generated queries. All persistence goes through here.
// The Memory logic layer composes a MemoryDelegate for its backend. // The Memory logic layer composes a MemoryDelegate for its backend.
type MemoryDelegate interface { // MemoryReader groups all read-only methods on the memory store. Enables
// Init creates tables and runs migrations. // independent optimization of the read path (caching, read replicas).
Init(ctx context.Context) error type MemoryReader interface {
// Close releases database resources.
Close() error
// --- Working Context ---
GetWorkingContext(ctx context.Context, agentID, sessionKey string) (*WorkingContext, error) 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) GetRecallItem(ctx context.Context, agentID string, id ids.UUID) (*RecallItem, error)
UpdateRecallItem(ctx context.Context, item *RecallItem) error GetRecallItemsByIDs(ctx context.Context, agentID string, itemIDs []ids.UUID) (map[ids.UUID]*RecallItem, error)
DeleteRecallItem(ctx context.Context, agentID string, id ids.UUID) error
ListRecallItems(ctx context.Context, agentID, sessionKey string, limit, offset int) ([]*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) 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) 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) 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) GetArchivalChunk(ctx context.Context, agentID string, id ids.UUID) (*ArchivalChunk, error)
ListArchivalChunks(ctx context.Context, agentID string, recallID 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) 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) ListSummaries(ctx context.Context, agentID, sessionKey string, limit int) ([]*MemorySummary, error)
// --- Stats ---
CountRecallItems(ctx context.Context, agentID, sessionKey string) (int, error) CountRecallItems(ctx context.Context, agentID, sessionKey string) (int, error)
CountArchivalChunks(ctx context.Context, agentID string) (int, error) CountArchivalChunks(ctx context.Context, agentID string) (int, error)
// --- Key-Value Store ---
GetKV(ctx context.Context, agentID, key string) (string, error) 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) ListKVByPrefix(ctx context.Context, agentID, prefix string, limit int) (map[string]string, error)
// --- Documents ---
GetDocument(ctx context.Context, agentID, name string) (*AgentDocument, error) 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) ListDocumentsByCategory(ctx context.Context, agentID, category string) ([]*AgentDocument, error)
ListAllDocuments(ctx context.Context, agentID 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) ListAuditEntries(ctx context.Context, agentID string, limit int) ([]*AuditEntry, error)
ListAuditEntriesByAction(ctx context.Context, agentID, action string, limit int) ([]*AuditEntry, error) ListAuditEntriesByAction(ctx context.Context, agentID, action string, limit int) ([]*AuditEntry, error)
CountAuditEntries(ctx context.Context, agentID string) (int, error) CountAuditEntries(ctx context.Context, agentID string) (int, error)
// --- Capability Detection ---
HasVectorSearch() bool HasVectorSearch() bool
HasFTS() 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 --- // --- Embedding interface ---
// EmbeddingProvider generates vector embeddings from text. // EmbeddingProvider generates vector embeddings from text.

View file

@ -43,6 +43,9 @@ func (m *mockDelegate) InsertRecallItem(_ context.Context, item *RecallItem) err
func (m *mockDelegate) GetRecallItem(_ context.Context, _ string, _ ids.UUID) (*RecallItem, error) { func (m *mockDelegate) GetRecallItem(_ context.Context, _ string, _ ids.UUID) (*RecallItem, error) {
return nil, nil 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) UpdateRecallItem(_ context.Context, _ *RecallItem) error { return nil }
func (m *mockDelegate) DeleteRecallItem(_ context.Context, _ string, _ ids.UUID) 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) { func (m *mockDelegate) ListRecallItems(_ context.Context, _, _ string, _, _ int) ([]*RecallItem, error) {
@ -58,6 +61,9 @@ func (m *mockDelegate) SearchArchivalByVector(_ context.Context, _ Embedding, _,
return nil, nil return nil, nil
} }
func (m *mockDelegate) InsertArchivalChunk(_ context.Context, _ *ArchivalChunk) error { return 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) { func (m *mockDelegate) GetArchivalChunk(_ context.Context, _ string, _ ids.UUID) (*ArchivalChunk, error) {
return nil, nil return nil, nil
} }
@ -96,6 +102,9 @@ func (m *mockDelegate) ListAllDocuments(_ context.Context, _ string) ([]*AgentDo
return nil, nil return nil, nil
} }
func (m *mockDelegate) InsertAuditEntry(_ context.Context, _ *AuditEntry) error { return 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) { func (m *mockDelegate) ListAuditEntries(_ context.Context, _ string, _ int) ([]*AuditEntry, error) {
return nil, nil return nil, nil
} }

View file

@ -44,6 +44,30 @@ type MemoryStore struct {
cfg Config cfg Config
agentID string agentID string
retrievalPolicyMu sync.Mutex 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. // New creates a MemoryStore.
@ -63,12 +87,59 @@ func New(delegate memory.MemoryDelegate, chunker memory.Chunker, embedder memory
embedder: embedder, embedder: embedder,
chunker: chunker, chunker: chunker,
cfg: cfg, cfg: cfg,
policyCache: policyCacheState{
flushEvery: 10,
},
} }
} }
// SetAgentID sets the agent identity used to scope all memory operations. // 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) { func (m *MemoryStore) SetAgentID(agentID string) {
m.agentID = agentID 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) --- // --- 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 { if err := m.delegate.DeleteArchivalChunks(ctx, id); err != nil {
return fmt.Errorf("delete archival chunks: %w", err) return fmt.Errorf("delete archival chunks: %w", err)
} }
m.invalidateVecCache()
return m.delegate.DeleteRecallItem(ctx, m.agentID, id) 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 { for i, chunk := range chunks {
var emb memory.Embedding var emb memory.Embedding
if i < len(embeddings) { if i < len(embeddings) {
emb = embeddings[i] emb = embeddings[i]
} }
archChunk := &memory.ArchivalChunk{ stored[i] = &memory.ArchivalChunk{
ID: ids.New(), ID: ids.New(),
RecallID: recallID, RecallID: recallID,
ChunkIndex: chunk.Index, ChunkIndex: chunk.Index,
@ -174,11 +247,13 @@ func (m *MemoryStore) StoreArchival(ctx context.Context, content, source string,
Source: source, Source: source,
Hash: hashContent(chunk.Text), 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 return recallID, nil
} }
@ -292,11 +367,15 @@ func (m *MemoryStore) Search(ctx context.Context, query string, opts memory.Sear
} }
m.retrievalPolicyMu.Lock() m.retrievalPolicyMu.Lock()
state := m.loadRetrievalPolicyState(ctx) m.ensurePolicyCacheLoaded(ctx)
gates := m.loadRetrievalPromotionGates(ctx)
metrics := m.loadRetrievalShadowMetrics(ctx)
parity := computeRetrievalParity(baselineFinal, augmentedFinal, limit) 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() m.retrievalPolicyMu.Unlock()
switch state.Mode { switch state.Mode {
@ -325,20 +404,28 @@ func (m *MemoryStore) postProcessMergedResults(
halfLife = m.cfg.DefaultHalfLifeHours halfLife = m.cfg.DefaultHalfLifeHours
} }
// Build a createdAt lookup from recall items // Batch-fetch all recall items referenced by merged results (single query).
createdAtMap := make(map[ids.UUID]time.Time) recallIDs := make([]ids.UUID, 0, len(merged))
for _, r := range merged { for _, r := range merged {
item, err := m.delegate.GetRecallItem(ctx, m.agentID, r.ID) if !r.ID.IsZero() && !strings.HasPrefix(r.Source, "working-context:") && !strings.HasPrefix(r.Source, "dag:") {
if err == nil && item != nil { recallIDs = append(recallIDs, r.ID)
createdAtMap[r.ID] = item.CreatedAt
} }
} }
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 { ApplyRecencyDecay(merged, time.Now(), halfLife, func(id ids.UUID) time.Time {
return createdAtMap[id] return createdAtMap[id]
}) })
// 2. Metadata pre-filtering (sectors, session_key, date range) // 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 // 3. Filter by min score
if opts.MinScore > 0 { if opts.MinScore > 0 {
@ -433,10 +520,10 @@ func (m *MemoryStore) hybridProjectionSearch(ctx context.Context, query string,
return results, nil return results, nil
} }
// applyMetadataFilters removes results that don't match the requested sector, // applyMetadataFiltersBatch removes results that don't match the requested sector,
// session_key, or date range constraints. It fetches recall item metadata // session_key, or date range constraints. Uses a pre-fetched batch of recall items
// from the delegate as needed. // to avoid N+1 queries.
func (m *MemoryStore) applyMetadataFilters(ctx context.Context, results []memory.SearchResult, opts memory.SearchOptions) []memory.SearchResult { 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 needSectorFilter := len(opts.Sectors) > 0
needSessionFilter := opts.SessionKey != "" needSessionFilter := opts.SessionKey != ""
needDateFilter := opts.DateAfter != nil || opts.DateBefore != nil needDateFilter := opts.DateAfter != nil || opts.DateBefore != nil
@ -445,7 +532,6 @@ func (m *MemoryStore) applyMetadataFilters(ctx context.Context, results []memory
return results return results
} }
// Build sector lookup set
sectorSet := make(map[memory.Sector]bool, len(opts.Sectors)) sectorSet := make(map[memory.Sector]bool, len(opts.Sectors))
for _, s := range opts.Sectors { for _, s := range opts.Sectors {
sectorSet[s] = true sectorSet[s] = true
@ -453,8 +539,6 @@ func (m *MemoryStore) applyMetadataFilters(ctx context.Context, results []memory
filtered := results[:0] filtered := results[:0]
for _, r := range results { 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 strings.HasPrefix(r.Source, "working-context:") || strings.HasPrefix(r.Source, "dag:") {
if needSessionFilter && !strings.Contains(r.Source, opts.SessionKey) { if needSessionFilter && !strings.Contains(r.Source, opts.SessionKey) {
continue continue
@ -462,8 +546,6 @@ func (m *MemoryStore) applyMetadataFilters(ctx context.Context, results []memory
if needSectorFilter && !sectorSet[r.Sector] { if needSectorFilter && !sectorSet[r.Sector] {
continue continue
} }
// Synthetic entries do not carry durable timestamps, so honor
// explicit date filters conservatively by excluding them.
if needDateFilter { if needDateFilter {
continue continue
} }
@ -471,23 +553,20 @@ func (m *MemoryStore) applyMetadataFilters(ctx context.Context, results []memory
continue continue
} }
item, err := m.delegate.GetRecallItem(ctx, m.agentID, r.ID) item := batch[r.ID]
if err != nil || item == nil { if item == nil {
continue continue
} }
if needSessionFilter && item.SessionKey != opts.SessionKey { if needSessionFilter && item.SessionKey != opts.SessionKey {
continue continue
} }
if needSectorFilter && !sectorSet[item.Sector] { if needSectorFilter && !sectorSet[item.Sector] {
continue continue
} }
if opts.DateAfter != nil && item.CreatedAt.Before(*opts.DateAfter) { if opts.DateAfter != nil && item.CreatedAt.Before(*opts.DateAfter) {
continue continue
} }
if opts.DateBefore != nil && item.CreatedAt.After(*opts.DateBefore) { if opts.DateBefore != nil && item.CreatedAt.After(*opts.DateBefore) {
continue continue
} }
@ -533,12 +612,41 @@ func (m *MemoryStore) vectorSearch(ctx context.Context, query string, opts memor
return m.vectorSearchGoSide(ctx, queryVec, limit) 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) { 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 var allChunks []*memory.ArchivalChunk
offset := 0 offset := 0
batchSize := 5000 batchSize := 5000
for { for len(allChunks) < maxVectorSearchChunks {
batch, err := m.delegate.ListAllArchivalChunks(ctx, m.agentID, batchSize, offset) batch, err := m.delegate.ListAllArchivalChunks(ctx, m.agentID, batchSize, offset)
if err != nil { if err != nil {
return nil, err return nil, err
@ -549,19 +657,46 @@ func (m *MemoryStore) vectorSearchGoSide(ctx context.Context, queryVec memory.Em
} }
offset += batchSize 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 { for _, chunk := range allChunks {
if len(chunk.Embedding) == 0 { if len(chunk.Embedding) == 0 {
continue continue
} }
inputs = append(inputs, VectorSearchInput{ items = append(items, VectorSearchInput{
Chunk: chunk, Chunk: chunk,
Embedding: chunk.Embedding, 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. // recallItemsToResults converts delegate recall items into search results.
@ -671,8 +806,9 @@ func (m *MemoryStore) OffloadToolResult(ctx context.Context, toolName, content,
// --- Lifecycle --- // --- Lifecycle ---
// Sync flushes pending writes to the remote replica (Turso). // 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 { func (m *MemoryStore) Sync() error {
m.FlushPolicyCache(context.Background())
type syncer interface{ Sync() error } type syncer interface{ Sync() error }
if s, ok := m.delegate.(syncer); ok { if s, ok := m.delegate.(syncer); ok {
return s.Sync() return s.Sync()
@ -681,6 +817,7 @@ func (m *MemoryStore) Sync() error {
} }
func (m *MemoryStore) Close() error { func (m *MemoryStore) Close() error {
m.FlushPolicyCache(context.Background())
return m.delegate.Close() return m.delegate.Close()
} }

View file

@ -318,6 +318,8 @@ func TestSearch_ShadowModeUsesBaselineAndTracksParity(t *testing.T) {
// Shadow mode should keep baseline result ordering for production output. // Shadow mode should keep baseline result ordering for production output.
assert.NotContains(t, results[0].Source, "working-context:s1") assert.NotContains(t, results[0].Source, "working-context:s1")
store.FlushPolicyCache(ctx)
metricsRaw, err := store.delegate.GetKV(ctx, "agent-1", retrievalPolicyMetricsKey) metricsRaw, err := store.delegate.GetKV(ctx, "agent-1", retrievalPolicyMetricsKey)
require.NoError(t, err) require.NoError(t, err)
require.NotEmpty(t, metricsRaw) require.NotEmpty(t, metricsRaw)
@ -363,6 +365,8 @@ func TestSearch_DoesNotPromoteWithoutAugmentedSignals(t *testing.T) {
}) })
require.NoError(t, err) require.NoError(t, err)
store.FlushPolicyCache(ctx)
stateRaw, err := store.delegate.GetKV(ctx, "agent-1", retrievalPolicyStateKey) stateRaw, err := store.delegate.GetKV(ctx, "agent-1", retrievalPolicyStateKey)
require.NoError(t, err) require.NoError(t, err)
var state retrievalPolicyState var state retrievalPolicyState
@ -407,6 +411,8 @@ func TestSearch_PromoteOnlyOnGateWin(t *testing.T) {
}) })
require.NoError(t, err) require.NoError(t, err)
store.FlushPolicyCache(ctx)
stateRaw, err := store.delegate.GetKV(ctx, "agent-1", retrievalPolicyStateKey) stateRaw, err := store.delegate.GetKV(ctx, "agent-1", retrievalPolicyStateKey)
require.NoError(t, err) require.NoError(t, err)
require.NotEmpty(t, stateRaw) require.NotEmpty(t, stateRaw)
@ -538,7 +544,7 @@ func TestUpdateRetrievalPolicy_PersistFailuresDoNotBlockTransitions(t *testing.T
{ID: ids.New(), Content: "baseline"}, {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)) assert.Empty(t, cmp.Diff(retrievalModePromoted, next.Mode))
} }
@ -582,6 +588,8 @@ func TestSearch_ConcurrentRetrievalPolicyUpdates(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
} }
store.FlushPolicyCache(ctx)
metricsRaw, err := store.delegate.GetKV(ctx, "agent-1", retrievalPolicyMetricsKey) metricsRaw, err := store.delegate.GetKV(ctx, "agent-1", retrievalPolicyMetricsKey)
require.NoError(t, err) require.NoError(t, err)
var metrics retrievalShadowMetrics var metrics retrievalShadowMetrics

View file

@ -285,7 +285,7 @@ func (m *MemoryStore) updateRetrievalPolicy(
parity retrievalParity, parity retrievalParity,
augmentedUsed bool, augmentedUsed bool,
baseline []memory.SearchResult, baseline []memory.SearchResult,
) retrievalPolicyState { ) (retrievalPolicyState, retrievalShadowMetrics) {
prevMode := state.Mode prevMode := state.Mode
metrics.TotalQueries++ metrics.TotalQueries++
@ -369,15 +369,20 @@ func (m *MemoryStore) updateRetrievalPolicy(
}) })
} }
if err := m.persistRetrievalShadowMetrics(ctx, metrics); err != nil { // State and metrics are written back to the in-memory policy cache by the
logger.WarnCF("memory", "failed to persist retrieval shadow metrics", // caller; the cache handles periodic flush to KV. On mode transitions we
map[string]interface{}{"mode": state.Mode, "error": err.Error()}) // 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 { return state, metrics
logger.WarnCF("memory", "failed to persist retrieval policy state",
map[string]interface{}{"mode": state.Mode, "error": err.Error()})
}
return state
} }
func safeRate(numerator, denominator int) float64 { func safeRate(numerator, denominator int) float64 {