refactor(memory): delegate sqlite, store, retrieval policy updates
This commit is contained in:
parent
7f1400f6c0
commit
2dd54a8269
6 changed files with 318 additions and 90 deletions
|
|
@ -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,
|
||||||
|
|
|
||||||
|
|
@ -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.
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -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()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -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
|
||||||
|
|
|
||||||
|
|
@ -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,6 +369,10 @@ func (m *MemoryStore) updateRetrievalPolicy(
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 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 {
|
if err := m.persistRetrievalShadowMetrics(ctx, metrics); err != nil {
|
||||||
logger.WarnCF("memory", "failed to persist retrieval shadow metrics",
|
logger.WarnCF("memory", "failed to persist retrieval shadow metrics",
|
||||||
map[string]interface{}{"mode": state.Mode, "error": err.Error()})
|
map[string]interface{}{"mode": state.Mode, "error": err.Error()})
|
||||||
|
|
@ -377,7 +381,8 @@ func (m *MemoryStore) updateRetrievalPolicy(
|
||||||
logger.WarnCF("memory", "failed to persist retrieval policy state",
|
logger.WarnCF("memory", "failed to persist retrieval policy state",
|
||||||
map[string]interface{}{"mode": state.Mode, "error": err.Error()})
|
map[string]interface{}{"mode": state.Mode, "error": err.Error()})
|
||||||
}
|
}
|
||||||
return state
|
}
|
||||||
|
return state, metrics
|
||||||
}
|
}
|
||||||
|
|
||||||
func safeRate(numerator, denominator int) float64 {
|
func safeRate(numerator, denominator int) float64 {
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue