From ae3bd9864df8a6c030b44a8b5a1d800e418404b9 Mon Sep 17 00:00:00 2001 From: ZanzyTHEbar Date: Wed, 4 Mar 2026 18:42:08 +0000 Subject: [PATCH] feat(memory): wire delegate, store, dag action mining, RL types and tests - delegate: sqlite integration, rl_types, rl_store_test - store: memory_store updates - dag: action_mining - migrate_sessions_test --- pkg/memory/dag/action_mining.go | 473 ++++++++++++++++++ pkg/memory/delegate/rl_store_test.go | 709 +++++++++++++++++++++++++++ pkg/memory/delegate/rl_types.go | 76 +++ pkg/memory/delegate/sqlite.go | 608 +++++++++++++++++++++-- pkg/memory/memory.go | 125 ++++- pkg/memory/migrate_sessions_test.go | 22 + pkg/memory/store/memory_store.go | 26 + 7 files changed, 1990 insertions(+), 49 deletions(-) create mode 100644 pkg/memory/dag/action_mining.go create mode 100644 pkg/memory/delegate/rl_store_test.go create mode 100644 pkg/memory/delegate/rl_types.go diff --git a/pkg/memory/dag/action_mining.go b/pkg/memory/dag/action_mining.go new file mode 100644 index 000000000..9c86ff58e --- /dev/null +++ b/pkg/memory/dag/action_mining.go @@ -0,0 +1,473 @@ +package dag + +import ( + "context" + "fmt" + "math" + "sort" + "strings" + "time" + + "github.com/ZanzyTHEbar/dragonscale/pkg/ids" +) + +// AuditEntry represents a single audit log entry. +type AuditEntry struct { + ID ids.UUID `json:"id"` + SessionID string `json:"session_id"` + AgentID string `json:"agent_id"` + ToolName string `json:"tool_name"` + Input string `json:"input"` + Output string `json:"output"` + Success bool `json:"success"` + DurationMs int `json:"duration_ms"` + CreatedAt time.Time `json:"created_at"` +} + +// ActionChain represents a sequence of related tool calls. +type ActionChain struct { + ID ids.UUID `json:"id"` + SessionID string `json:"session_id"` + AgentID string `json:"agent_id"` + Steps []ChainStep `json:"steps"` + Score float64 `json:"score"` // Quality score (0.0-1.0) + ToolDiversity int `json:"tool_diversity"` // Number of unique tools + Success bool `json:"success"` + CreatedAt time.Time `json:"created_at"` + CompletedAt time.Time `json:"completed_at"` + + // Classification + Intent string `json:"intent"` // High-level goal + Category string `json:"category"` // Task category +} + +// ChainStep represents a single step in an action chain. +type ChainStep struct { + StepNumber int `json:"step_number"` + ToolName string `json:"tool_name"` + Input string `json:"input"` + Output string `json:"output"` + Success bool `json:"success"` + DurationMs int `json:"duration_ms"` + Description string `json:"description"` // Human-readable +} + +// ActionMiner mines successful tool call sequences from audit logs. +type ActionMiner struct { + minChainLength int + maxChainLength int + lookbackWindow time.Duration +} + +// ActionMinerConfig configures the action mining behavior. +type ActionMinerConfig struct { + MinChainLength int // Minimum steps to consider (default 2) + MaxChainLength int // Maximum steps to consider (default 10) + LookbackWindow time.Duration // How far back to mine (default 7 days) +} + +// DefaultActionMinerConfig returns sensible defaults. +func DefaultActionMinerConfig() ActionMinerConfig { + return ActionMinerConfig{ + MinChainLength: 2, + MaxChainLength: 10, + LookbackWindow: 7 * 24 * time.Hour, + } +} + +// NewActionMiner creates a new action miner. +func NewActionMiner(cfg ActionMinerConfig) *ActionMiner { + return &ActionMiner{ + minChainLength: cfg.MinChainLength, + maxChainLength: cfg.MaxChainLength, + lookbackWindow: cfg.LookbackWindow, + } +} + +// AuditStore provides access to audit log entries. +type AuditStore interface { + // GetEntries retrieves audit entries within a time window + GetEntries(ctx context.Context, agentID string, since time.Time) ([]*AuditEntry, error) + + // GetSessionEntries retrieves all entries for a specific session + GetSessionEntries(ctx context.Context, sessionID string) ([]*AuditEntry, error) + + // StoreChain saves a mined action chain + StoreChain(ctx context.Context, chain *ActionChain) error + + // GetTopChains retrieves the highest-scoring chains + GetTopChains(ctx context.Context, agentID string, category string, limit int) ([]*ActionChain, error) +} + +// MineChains extracts action chains from audit logs. +func (m *ActionMiner) MineChains(ctx context.Context, store AuditStore, agentID string) ([]*ActionChain, error) { + since := time.Now().Add(-m.lookbackWindow) + + entries, err := store.GetEntries(ctx, agentID, since) + if err != nil { + return nil, fmt.Errorf("get audit entries: %w", err) + } + + // Group entries by session + sessions := groupBySession(entries) + + var chains []*ActionChain + for sessionID, sessionEntries := range sessions { + if len(sessionEntries) < m.minChainLength { + continue + } + + chain := m.buildChain(sessionID, agentID, sessionEntries) + if chain != nil { + chains = append(chains, chain) + } + } + + // Score and rank chains + m.scoreChains(chains) + + // Sort by score descending + sort.Slice(chains, func(i, j int) bool { + return chains[i].Score > chains[j].Score + }) + + return chains, nil +} + +func groupBySession(entries []*AuditEntry) map[string][]*AuditEntry { + sessions := make(map[string][]*AuditEntry) + for _, entry := range entries { + sessions[entry.SessionID] = append(sessions[entry.SessionID], entry) + } + + // Sort each session by time + for _, sessionEntries := range sessions { + sort.Slice(sessionEntries, func(i, j int) bool { + return sessionEntries[i].CreatedAt.Before(sessionEntries[j].CreatedAt) + }) + } + + return sessions +} + +func (m *ActionMiner) buildChain(sessionID, agentID string, entries []*AuditEntry) *ActionChain { + if len(entries) > m.maxChainLength { + entries = entries[:m.maxChainLength] + } + + chain := &ActionChain{ + ID: ids.New(), + SessionID: sessionID, + AgentID: agentID, + CreatedAt: entries[0].CreatedAt, + Steps: make([]ChainStep, len(entries)), + } + + uniqueTools := make(map[string]struct{}) + allSuccessful := true + var totalDuration int + + for i, entry := range entries { + chain.Steps[i] = ChainStep{ + StepNumber: i + 1, + ToolName: entry.ToolName, + Input: entry.Input, + Output: truncate(entry.Output, 200), + Success: entry.Success, + DurationMs: entry.DurationMs, + Description: fmt.Sprintf("Step %d: %s", i+1, describeToolCall(entry)), + } + + uniqueTools[entry.ToolName] = struct{}{} + if !entry.Success { + allSuccessful = false + } + totalDuration += entry.DurationMs + } + + chain.CompletedAt = entries[len(entries)-1].CreatedAt + chain.ToolDiversity = len(uniqueTools) + chain.Success = allSuccessful + + // Classify intent + chain.Intent = m.classifyIntent(chain) + chain.Category = m.classifyCategory(chain) + + return chain +} + +func (m *ActionMiner) scoreChains(chains []*ActionChain) { + for _, chain := range chains { + chain.Score = m.calculateScore(chain) + } +} + +func (m *ActionMiner) calculateScore(chain *ActionChain) float64 { + // Base score from success + successScore := 0.0 + if chain.Success { + successScore = 1.0 + } else { + // Partial credit if most steps succeeded + successCount := 0 + for _, step := range chain.Steps { + if step.Success { + successCount++ + } + } + successScore = float64(successCount) / float64(len(chain.Steps)) + } + + // Tool diversity bonus (more tools = more interesting) + diversityScore := math.Min(float64(chain.ToolDiversity)/5.0, 1.0) + + // Length score (sweet spot around 3-5 steps) + lengthScore := 1.0 + stepCount := len(chain.Steps) + if stepCount < 2 { + lengthScore = 0.5 + } else if stepCount > 8 { + lengthScore = 0.8 + } + + // Recency bonus (more recent = more relevant) + age := time.Since(chain.CreatedAt) + recencyScore := 1.0 - math.Min(age.Hours()/(7*24), 1.0) + + // Weighted combination + score := 0.4*successScore + 0.25*diversityScore + 0.2*lengthScore + 0.15*recencyScore + + return math.Max(0.0, math.Min(1.0, score)) +} + +func (m *ActionMiner) classifyIntent(chain *ActionChain) string { + // Simple classification based on first tool + if len(chain.Steps) == 0 { + return "unknown" + } + + firstTool := strings.ToLower(chain.Steps[0].ToolName) + + // Map tools to intents + switch { + case strings.Contains(firstTool, "search"): + return "research" + case strings.Contains(firstTool, "read") || strings.Contains(firstTool, "file"): + return "read_file" + case strings.Contains(firstTool, "write") || strings.Contains(firstTool, "edit"): + return "write_file" + case strings.Contains(firstTool, "run") || strings.Contains(firstTool, "exec"): + return "execute" + case strings.Contains(firstTool, "test"): + return "test" + case strings.Contains(firstTool, "git"): + return "version_control" + default: + return "general" + } +} + +func (m *ActionMiner) classifyCategory(chain *ActionChain) string { + // Classify based on tool combination patterns + tools := make(map[string]int) + for _, step := range chain.Steps { + tools[strings.ToLower(step.ToolName)]++ + } + + // Check for common patterns + hasSearch := tools["search"] > 0 || tools["grep"] > 0 + hasFileOps := tools["read_file"] > 0 || tools["write_file"] > 0 || tools["edit_file"] > 0 + hasExec := tools["run_command"] > 0 || tools["execute"] > 0 + hasGit := tools["git"] > 0 + + switch { + case hasGit: + return "git_workflow" + case hasSearch && hasFileOps: + return "file_research" + case hasFileOps && hasExec: + return "development" + case hasSearch && !hasFileOps: + return "research" + case hasExec: + return "execution" + default: + return "general" + } +} + +func describeToolCall(entry *AuditEntry) string { + // Create a human-readable description + switch entry.ToolName { + case "search": + return fmt.Sprintf("Searched for '%s'", truncate(entry.Input, 40)) + case "read_file": + return fmt.Sprintf("Read file: %s", truncate(entry.Input, 40)) + case "write_file": + return fmt.Sprintf("Wrote to file: %s", truncate(entry.Input, 40)) + case "run_command": + return fmt.Sprintf("Ran command: %s", truncate(entry.Input, 40)) + default: + return fmt.Sprintf("Used %s", entry.ToolName) + } +} + +func truncate(s string, maxLen int) string { + if len(s) <= maxLen { + return s + } + return s[:maxLen-3] + "..." +} + +// FewShotFormatter formats action chains as few-shot examples. +type FewShotFormatter struct { + maxExamples int + maxStepsPerExample int +} + +// NewFewShotFormatter creates a formatter for few-shot examples. +func NewFewShotFormatter(maxExamples, maxStepsPerExample int) *FewShotFormatter { + return &FewShotFormatter{ + maxExamples: maxExamples, + maxStepsPerExample: maxStepsPerExample, + } +} + +// FormatChain formats a single chain as a few-shot example. +func (f *FewShotFormatter) FormatChain(chain *ActionChain) string { + var b strings.Builder + + b.WriteString(fmt.Sprintf("\n", chain.Intent, chain.Category)) + + for i, step := range chain.Steps { + if i >= f.maxStepsPerExample { + b.WriteString(fmt.Sprintf(" ... (%d more steps) ...\n", len(chain.Steps)-i)) + break + } + + status := "✓" + if !step.Success { + status = "✗" + } + + b.WriteString(fmt.Sprintf(" %s %d. %s: %s\n", + status, step.StepNumber, step.ToolName, step.Description)) + } + + b.WriteString("\n") + return b.String() +} + +// FormatExamples formats multiple chains as few-shot context. +func (f *FewShotFormatter) FormatExamples(chains []*ActionChain) string { + if len(chains) == 0 { + return "" + } + + var b strings.Builder + b.WriteString("## Successful Action Patterns\n\n") + b.WriteString("Here are examples of successful task completion patterns:\n\n") + + for i, chain := range chains { + if i >= f.maxExamples { + break + } + + b.WriteString(f.FormatChain(chain)) + b.WriteString("\n") + } + + return b.String() +} + +// FormatForPrompt formats chains for direct injection into prompts. +func (f *FewShotFormatter) FormatForPrompt(chains []*ActionChain, intent string) string { + // Filter by intent if specified + var filtered []*ActionChain + for _, chain := range chains { + if intent == "" || chain.Intent == intent || chain.Category == intent { + filtered = append(filtered, chain) + } + } + + if len(filtered) == 0 { + return "" + } + + return f.FormatExamples(filtered) +} + +// ActionReplayTask is a Cortex task for mining and caching action patterns. +type ActionReplayTask struct { + miner *ActionMiner + store AuditStore + config ActionReplayConfig +} + +// ActionReplayConfig configures the action replay task. +type ActionReplayConfig struct { + MineInterval time.Duration // How often to mine (default 1 hour) + MinChainLength int // Minimum chain length (default 2) + TopKCache int // Number of chains to cache (default 20) +} + +// DefaultActionReplayConfig returns sensible defaults. +func DefaultActionReplayConfig() ActionReplayConfig { + return ActionReplayConfig{ + MineInterval: 1 * time.Hour, + MinChainLength: 2, + TopKCache: 20, + } +} + +// NewActionReplayTask creates an action replay Cortex task. +func NewActionReplayTask(config ActionReplayConfig, store AuditStore) *ActionReplayTask { + return &ActionReplayTask{ + miner: NewActionMiner(DefaultActionMinerConfig()), + store: store, + config: config, + } +} + +// Name returns the task identifier. +func (t *ActionReplayTask) Name() string { + return "action_replay" +} + +// Interval returns the task run interval. +func (t *ActionReplayTask) Interval() time.Duration { + return t.config.MineInterval +} + +// Execute performs action mining and caching. +func (t *ActionReplayTask) Execute(ctx context.Context) error { + // Mine new chains from recent sessions + chains, err := t.miner.MineChains(ctx, t.store, "default") + if err != nil { + return fmt.Errorf("mine chains: %w", err) + } + + // Store top-K chains + for i, chain := range chains { + if i >= t.config.TopKCache { + break + } + if err := t.store.StoreChain(ctx, chain); err != nil { + // Log but continue + continue + } + } + + return nil +} + +// GetFewShotContext retrieves formatted few-shot examples for prompt injection. +func GetFewShotContext(ctx context.Context, store AuditStore, agentID, intent string, maxExamples int) (string, error) { + chains, err := store.GetTopChains(ctx, agentID, intent, maxExamples) + if err != nil { + return "", fmt.Errorf("get top chains: %w", err) + } + + formatter := NewFewShotFormatter(maxExamples, 5) + return formatter.FormatForPrompt(chains, intent), nil +} diff --git a/pkg/memory/delegate/rl_store_test.go b/pkg/memory/delegate/rl_store_test.go new file mode 100644 index 000000000..1085c8a62 --- /dev/null +++ b/pkg/memory/delegate/rl_store_test.go @@ -0,0 +1,709 @@ +package delegate + +import ( + "context" + "testing" + "time" + + "github.com/ZanzyTHEbar/dragonscale/pkg/ids" + "github.com/ZanzyTHEbar/dragonscale/pkg/memory" + memsqlc "github.com/ZanzyTHEbar/dragonscale/pkg/memory/sqlc" +) + +// setupRLTest creates an in-memory delegate with initialized schema for RL tests. +func setupRLTest(t *testing.T) *LibSQLDelegate { + t.Helper() + d, err := NewLibSQLInMemory() + if err != nil { + t.Fatalf("NewLibSQLInMemory: %v", err) + } + if err := d.Init(t.Context()); err != nil { + t.Fatalf("Init: %v", err) + } + t.Cleanup(func() { d.Close() }) + return d +} + +// insertTestRecallItem creates a recall item for testing RL operations. +func insertTestRecallItem(ctx context.Context, t *testing.T, d *LibSQLDelegate, agentID string) ids.UUID { + t.Helper() + item := &memory.RecallItem{ + ID: ids.New(), + AgentID: agentID, + SessionKey: "test-session", + Role: "assistant", + Sector: memory.SectorEpisodic, + Importance: 0.8, + Salience: 0.6, + DecayRate: 0.01, + Content: "Test content for RL weight updates", + Tags: "test,rl", + } + if err := d.InsertRecallItem(ctx, item); err != nil { + t.Fatalf("InsertRecallItem: %v", err) + } + return item.ID +} + +func TestSQLiteDelegate_GetTaskBaseline_NoBaseline(t *testing.T) { + t.Parallel() + d := setupRLTest(t) + ctx := t.Context() + + // Get baseline for agent without one - should return nil + baseline, err := d.GetTaskBaseline(ctx, "new-agent") + if err != nil { + t.Fatalf("GetTaskBaseline: %v", err) + } + if baseline != nil { + t.Error("expected nil baseline for new agent") + } +} + +func TestSQLiteDelegate_GetTaskBaseline_AfterUpdate(t *testing.T) { + t.Parallel() + d := setupRLTest(t) + ctx := t.Context() + agentID := "test-agent" + + // Initially no baseline + baseline, err := d.GetTaskBaseline(ctx, agentID) + if err != nil { + t.Fatalf("GetTaskBaseline: %v", err) + } + if baseline != nil { + t.Error("expected nil baseline initially") + } + + // Update baseline + newBaseline := &TaskBaseline{ + Count: 10, + MeanTokens: 1000, + MeanErrors: 5, + MeanUserCorrections: 2, + M2Tokens: 5000, + M2Errors: 50, + M2UserCorrections: 20, + } + if err := d.UpdateTaskBaseline(ctx, agentID, newBaseline); err != nil { + t.Fatalf("UpdateTaskBaseline: %v", err) + } + + // Get baseline again + baseline, err = d.GetTaskBaseline(ctx, agentID) + if err != nil { + t.Fatalf("GetTaskBaseline after update: %v", err) + } + if baseline == nil { + t.Fatal("expected non-nil baseline after update") + } + + // Verify values + if baseline.Count != 10 { + t.Errorf("Count = %d, want 10", baseline.Count) + } + if baseline.MeanTokens != 1000 { + t.Errorf("MeanTokens = %f, want 1000", baseline.MeanTokens) + } + if baseline.MeanErrors != 5 { + t.Errorf("MeanErrors = %f, want 5", baseline.MeanErrors) + } + if baseline.MeanUserCorrections != 2 { + t.Errorf("MeanUserCorrections = %f, want 2", baseline.MeanUserCorrections) + } + if baseline.M2Tokens != 5000 { + t.Errorf("M2Tokens = %f, want 5000", baseline.M2Tokens) + } + if baseline.M2Errors != 50 { + t.Errorf("M2Errors = %f, want 50", baseline.M2Errors) + } + if baseline.M2UserCorrections != 20 { + t.Errorf("M2UserCorrections = %f, want 20", baseline.M2UserCorrections) + } +} + +func TestSQLiteDelegate_UpdateTaskBaseline_MultipleUpdates(t *testing.T) { + t.Parallel() + d := setupRLTest(t) + ctx := t.Context() + agentID := "test-agent" + + // First update + baseline1 := &TaskBaseline{ + Count: 5, + MeanTokens: 500, + } + if err := d.UpdateTaskBaseline(ctx, agentID, baseline1); err != nil { + t.Fatalf("UpdateTaskBaseline (1): %v", err) + } + + // Second update (should overwrite) + baseline2 := &TaskBaseline{ + Count: 15, + MeanTokens: 1500, + MeanErrors: 10, + MeanUserCorrections: 3, + M2Tokens: 10000, + M2Errors: 100, + M2UserCorrections: 30, + } + if err := d.UpdateTaskBaseline(ctx, agentID, baseline2); err != nil { + t.Fatalf("UpdateTaskBaseline (2): %v", err) + } + + // Verify second values + baseline, err := d.GetTaskBaseline(ctx, agentID) + if err != nil { + t.Fatalf("GetTaskBaseline: %v", err) + } + if baseline == nil { + t.Fatal("expected non-nil baseline") + } + + if baseline.Count != 15 { + t.Errorf("Count = %d, want 15", baseline.Count) + } + if baseline.MeanTokens != 1500 { + t.Errorf("MeanTokens = %f, want 1500", baseline.MeanTokens) + } + if baseline.MeanErrors != 10 { + t.Errorf("MeanErrors = %f, want 10", baseline.MeanErrors) + } +} + +func TestSQLiteDelegate_UpdateMemoryWeight(t *testing.T) { + t.Parallel() + d := setupRLTest(t) + ctx := t.Context() + agentID := "test-agent" + + // Insert a recall item first + memoryID := insertTestRecallItem(ctx, t, d, agentID) + + // Update the memory weight via direct query + newWeight := 2.5 + credit := 3.0 + paramsRLWeight := newWeight + paramsRLCredit := credit + err := d.Queries().UpdateMemoryWeight(ctx, memsqlc.UpdateMemoryWeightParams{ + RlWeight: ¶msRLWeight, + RlCredit: ¶msRLCredit, + ID: memoryID, + AgentID: agentID, + }) + if err != nil { + t.Fatalf("UpdateMemoryWeight: %v", err) + } + + // Verify the item still exists (GetRecallItem doesn't return RL fields) + item, err := d.GetRecallItem(ctx, agentID, memoryID) + if err != nil { + t.Fatalf("GetRecallItem: %v", err) + } + if item == nil { + t.Fatal("expected non-nil recall item") + } +} + +func TestSQLiteDelegate_UpdateMemoryWeight_NonExistent(t *testing.T) { + t.Parallel() + d := setupRLTest(t) + ctx := t.Context() + agentID := "test-agent" + + // Try to update weight for non-existent memory via direct query + nonExistentID := ids.New() + paramsWeight := 2.0 + paramsCredit := 1.0 + err := d.Queries().UpdateMemoryWeight(ctx, memsqlc.UpdateMemoryWeightParams{ + RlWeight: ¶msWeight, + RlCredit: ¶msCredit, + ID: nonExistentID, + AgentID: agentID, + }) + // Query succeeds but doesn't update anything (no error for non-existent) + if err != nil { + t.Errorf("UpdateMemoryWeight should not error for non-existent: %v", err) + } +} + +func TestSQLiteDelegate_UpdateMemorySelfReport(t *testing.T) { + t.Parallel() + d := setupRLTest(t) + ctx := t.Context() + agentID := "test-agent" + + // Insert a recall item first + memoryID := insertTestRecallItem(ctx, t, d, agentID) + + // Update self-report score via direct query + score := int64(2) + err := d.Queries().UpdateMemorySelfReportScore(ctx, memsqlc.UpdateMemorySelfReportScoreParams{ + SelfReportScore: &score, + ID: memoryID, + AgentID: agentID, + }) + if err != nil { + t.Fatalf("UpdateMemorySelfReportScore: %v", err) + } + + // Verify the item still exists (GetRecallItem doesn't return self_report_score) + item, err := d.GetRecallItem(ctx, agentID, memoryID) + if err != nil { + t.Fatalf("GetRecallItem: %v", err) + } + if item == nil { + t.Fatal("expected non-nil recall item") + } +} + +func TestSQLiteDelegate_UpdateMemorySelfReport_NonExistent(t *testing.T) { + t.Parallel() + d := setupRLTest(t) + ctx := t.Context() + agentID := "test-agent" + + // Try to update self-report for non-existent memory via direct query + nonExistentID := ids.New() + score := int64(3) + err := d.Queries().UpdateMemorySelfReportScore(ctx, memsqlc.UpdateMemorySelfReportScoreParams{ + SelfReportScore: &score, + ID: nonExistentID, + AgentID: agentID, + }) + // Query succeeds but doesn't update anything (no error for non-existent) + if err != nil { + t.Errorf("UpdateMemorySelfReportScore should not error for non-existent: %v", err) + } +} + +func TestSQLiteDelegate_StoreDetectedPattern(t *testing.T) { + t.Parallel() + d := setupRLTest(t) + ctx := t.Context() + + // Store a detected pattern + pattern := DetectedPattern{ + Type: "correction", + Description: "Tool read_file was corrected from wrong path to correct path", + Weight: 1.0, + Category: "correction", + SessionID: "session-123", + AgentID: "agent-456", + } + + if err := d.StoreDetectedPattern(ctx, pattern); err != nil { + t.Fatalf("StoreDetectedPattern: %v", err) + } + + // Verify the pattern was stored as a recall item + items, err := d.ListRecallItems(ctx, pattern.AgentID, pattern.SessionID, 10, 0) + if err != nil { + t.Fatalf("ListRecallItems: %v", err) + } + if len(items) != 1 { + t.Fatalf("expected 1 recall item, got %d", len(items)) + } + + item := items[0] + if item.Content != pattern.Description { + t.Errorf("Content = %q, want %q", item.Content, pattern.Description) + } + if item.Importance != pattern.Weight { + t.Errorf("Importance = %f, want %f", item.Importance, pattern.Weight) + } + if item.Salience != pattern.Weight { + t.Errorf("Salience = %f, want %f", item.Salience, pattern.Weight) + } + if item.Sector != memory.SectorReflective { + t.Errorf("Sector = %v, want %v", item.Sector, memory.SectorReflective) + } + if item.Role != "system" { + t.Errorf("Role = %q, want %q", item.Role, "system") + } +} + +func TestSQLiteDelegate_StoreDetectedPattern_Multiple(t *testing.T) { + t.Parallel() + d := setupRLTest(t) + ctx := t.Context() + + patterns := []DetectedPattern{ + {Type: "correction", Description: "Pattern 1", Weight: 1.0, Category: "correction", SessionID: "s1", AgentID: "a1"}, + {Type: "discovery", Description: "Pattern 2", Weight: 1.2, Category: "discovery", SessionID: "s1", AgentID: "a1"}, + {Type: "failure_pattern", Description: "Pattern 3", Weight: 1.5, Category: "correction", SessionID: "s2", AgentID: "a2"}, + } + + for _, pattern := range patterns { + if err := d.StoreDetectedPattern(ctx, pattern); err != nil { + t.Fatalf("StoreDetectedPattern: %v", err) + } + } + + // Check items in first session + items1, err := d.ListRecallItems(ctx, "a1", "s1", 10, 0) + if err != nil { + t.Fatalf("ListRecallItems (a1/s1): %v", err) + } + if len(items1) != 2 { + t.Errorf("expected 2 items in a1/s1, got %d", len(items1)) + } + + // Check items in second session + items2, err := d.ListRecallItems(ctx, "a2", "s2", 10, 0) + if err != nil { + t.Fatalf("ListRecallItems (a2/s2): %v", err) + } + if len(items2) != 1 { + t.Errorf("expected 1 item in a2/s2, got %d", len(items2)) + } +} + +func TestSQLiteDelegate_GetCompletedTasks(t *testing.T) { + t.Parallel() + d := setupRLTest(t) + ctx := t.Context() + + // This is a placeholder implementation that returns empty list + tasks, err := d.GetCompletedTasks(ctx, time.Time{}) + if err != nil { + t.Fatalf("GetCompletedTasks: %v", err) + } + if len(tasks) != 0 { + t.Errorf("expected 0 tasks (placeholder), got %d", len(tasks)) + } +} + +func TestSQLiteDelegate_GetRetrievedMemories(t *testing.T) { + t.Parallel() + d := setupRLTest(t) + ctx := t.Context() + + // This is a placeholder implementation that returns empty list + memories, err := d.GetRetrievedMemories(ctx, "task-123") + if err != nil { + t.Fatalf("GetRetrievedMemories: %v", err) + } + if len(memories) != 0 { + t.Errorf("expected 0 memories (placeholder), got %d", len(memories)) + } +} + +func TestSQLiteDelegate_GetRecentAuditEntries(t *testing.T) { + t.Parallel() + d := setupRLTest(t) + ctx := t.Context() + + // Insert some audit entries + entries := []*memory.AuditEntry{ + { + ID: ids.New(), + AgentID: "audit-agent", + SessionKey: "session-1", + Action: "read_file", + Target: "/path/to/file", + Input: `{"path": "/test"}`, + Output: "content", + }, + { + ID: ids.New(), + AgentID: "audit-agent", + SessionKey: "session-1", + Action: "write_file", + Target: "/path/to/output", + Input: `{"path": "/output"}`, + Output: "success", + }, + } + + for _, entry := range entries { + if err := d.InsertAuditEntry(ctx, entry); err != nil { + t.Fatalf("InsertAuditEntry: %v", err) + } + } + + // Get recent audit entries (all of them, since time is in the past) + auditEntries, err := d.GetRecentAuditEntries(ctx, time.Time{}) + if err != nil { + t.Fatalf("GetRecentAuditEntries: %v", err) + } + + // The implementation uses ListAuditEntries with empty agent_id which may filter results + // Just verify the query executes without error + t.Logf("Got %d audit entries", len(auditEntries)) +} + +func TestSQLiteDelegate_GetHighTokenSessions(t *testing.T) { + t.Parallel() + d := setupRLTest(t) + ctx := t.Context() + + // This is a placeholder implementation that returns empty list + sessions, err := d.GetHighTokenSessions(ctx, 1000) + if err != nil { + t.Fatalf("GetHighTokenSessions: %v", err) + } + if len(sessions) != 0 { + t.Errorf("expected 0 sessions (placeholder), got %d", len(sessions)) + } +} + +func TestSQLiteDelegate_RLStore_Integration(t *testing.T) { + t.Parallel() + d := setupRLTest(t) + ctx := t.Context() + agentID := "integration-agent" + + t.Run("BaselineFlow", func(t *testing.T) { + // Initially no baseline + baseline, err := d.GetTaskBaseline(ctx, agentID) + if err != nil { + t.Fatalf("GetTaskBaseline: %v", err) + } + if baseline != nil { + t.Error("expected nil baseline initially") + } + + // Update baseline multiple times (simulating task processing) + for i := 1; i <= 5; i++ { + baseline := &TaskBaseline{ + Count: i, + MeanTokens: float64(1000 + i*100), + MeanErrors: float64(i), + MeanUserCorrections: float64(i % 2), + M2Tokens: float64(i * 1000), + M2Errors: float64(i * 10), + M2UserCorrections: float64(i * 5), + } + if err := d.UpdateTaskBaseline(ctx, agentID, baseline); err != nil { + t.Fatalf("UpdateTaskBaseline iteration %d: %v", i, err) + } + } + + // Verify final baseline + baseline, err = d.GetTaskBaseline(ctx, agentID) + if err != nil { + t.Fatalf("GetTaskBaseline final: %v", err) + } + if baseline == nil { + t.Fatal("expected non-nil baseline") + } + if baseline.Count != 5 { + t.Errorf("Count = %d, want 5", baseline.Count) + } + }) + + t.Run("MemoryWeightUpdates", func(t *testing.T) { + // Create multiple recall items + memoryIDs := make([]ids.UUID, 3) + for i := 0; i < 3; i++ { + memoryIDs[i] = insertTestRecallItem(ctx, t, d, agentID) + } + + // Update weights for each memory via direct query + // Note: GetRecallItem doesn't return RL fields, so we just verify no errors + weights := []float64{1.5, 2.0, 2.5} + credits := []float64{1.0, 2.0, 3.0} + for i, memoryID := range memoryIDs { + paramsRLWeight := weights[i] + paramsRLCredit := credits[i] + err := d.Queries().UpdateMemoryWeight(ctx, memsqlc.UpdateMemoryWeightParams{ + RlWeight: ¶msRLWeight, + RlCredit: ¶msRLCredit, + ID: memoryID, + AgentID: agentID, + }) + if err != nil { + t.Fatalf("UpdateMemoryWeight %d: %v", i, err) + } + } + + // Verify items still exist after update + for i, memoryID := range memoryIDs { + item, err := d.GetRecallItem(ctx, agentID, memoryID) + if err != nil { + t.Fatalf("GetRecallItem %d: %v", i, err) + } + if item == nil { + t.Fatalf("item %d is nil after weight update", i) + } + } + }) + + t.Run("SelfReportUpdates", func(t *testing.T) { + memoryID := insertTestRecallItem(ctx, t, d, agentID) + + // Update self-report scores via direct query + // Note: GetRecallItem doesn't return self_report_score, so we just verify no errors + scores := []int64{0, 1, 2, 3} + for _, score := range scores { + err := d.Queries().UpdateMemorySelfReportScore(ctx, memsqlc.UpdateMemorySelfReportScoreParams{ + SelfReportScore: &score, + ID: memoryID, + AgentID: agentID, + }) + if err != nil { + t.Fatalf("UpdateMemorySelfReportScore %d: %v", score, err) + } + } + + // Verify item still exists after updates + item, err := d.GetRecallItem(ctx, agentID, memoryID) + if err != nil { + t.Fatalf("GetRecallItem: %v", err) + } + if item == nil { + t.Fatal("item is nil after self-report updates") + } + }) + + t.Run("PatternStorage", func(t *testing.T) { + patterns := []DetectedPattern{ + {Type: "correction", Description: "Correction pattern", Weight: 1.0, Category: "correction", SessionID: "sess-1", AgentID: agentID}, + {Type: "discovery", Description: "Discovery pattern", Weight: 1.2, Category: "discovery", SessionID: "sess-2", AgentID: agentID}, + {Type: "failure_pattern", Description: "Failure pattern", Weight: 1.5, Category: "correction", SessionID: "sess-3", AgentID: agentID}, + } + + for _, pattern := range patterns { + if err := d.StoreDetectedPattern(ctx, pattern); err != nil { + t.Fatalf("StoreDetectedPattern: %v", err) + } + } + + // Count all patterns stored for this agent + count, err := d.CountRecallItems(ctx, agentID, "") + if err != nil { + t.Fatalf("CountRecallItems: %v", err) + } + // Should have 3 patterns + previous test items + if count < 3 { + t.Errorf("expected at least 3 recall items for patterns, got %d", count) + } + }) +} + +func TestSQLiteDelegate_TaskRecordTypes(t *testing.T) { + t.Parallel() + + // Test that TaskRecord type is properly defined + record := TaskRecord{ + ID: "task-1", + Description: "Test task", + TokensUsed: 100, + ToolCalls: 5, + Errors: 1, + UserCorrections: 0, + Completed: true, + } + + if record.ID != "task-1" { + t.Error("TaskRecord ID mismatch") + } + if record.TokensUsed != 100 { + t.Error("TaskRecord TokensUsed mismatch") + } + if !record.Completed { + t.Error("TaskRecord Completed should be true") + } +} + +func TestSQLiteDelegate_RetrievedMemoryRecordTypes(t *testing.T) { + t.Parallel() + + score := 2 + record := RetrievedMemoryRecord{ + MemoryID: ids.New(), + Similarity: 0.85, + SelfReportScore: &score, + } + + if record.Similarity != 0.85 { + t.Error("RetrievedMemoryRecord Similarity mismatch") + } + if record.SelfReportScore == nil || *record.SelfReportScore != 2 { + t.Error("RetrievedMemoryRecord SelfReportScore mismatch") + } + + // Test with nil score + record2 := RetrievedMemoryRecord{ + MemoryID: ids.New(), + Similarity: 0.75, + SelfReportScore: nil, + } + if record2.SelfReportScore != nil { + t.Error("RetrievedMemoryRecord SelfReportScore should be nil") + } +} + +func TestSQLiteDelegate_DetectedPatternTypes(t *testing.T) { + t.Parallel() + + pattern := DetectedPattern{ + Type: "correction", + Description: "Tool corrected", + Weight: 1.0, + Category: "correction", + SessionID: "session-123", + AgentID: "agent-456", + } + + if pattern.Type != "correction" { + t.Error("DetectedPattern Type mismatch") + } + if pattern.Weight != 1.0 { + t.Error("DetectedPattern Weight mismatch") + } + if pattern.SessionID != "session-123" { + t.Error("DetectedPattern SessionID mismatch") + } +} + +func TestSQLiteDelegate_AuditEntryTypes(t *testing.T) { + t.Parallel() + + entry := AuditEntry{ + ID: "entry-1", + Timestamp: time.Now(), + ToolName: "read_file", + ToolInput: `{"path": "/test"}`, + Success: true, + ErrorMsg: "", + SessionID: "session-1", + AgentID: "agent-1", + } + + if entry.ID != "entry-1" { + t.Error("AuditEntry ID mismatch") + } + if entry.ToolName != "read_file" { + t.Error("AuditEntry ToolName mismatch") + } + if !entry.Success { + t.Error("AuditEntry Success should be true") + } +} + +func TestSQLiteDelegate_SessionSummaryTypes(t *testing.T) { + t.Parallel() + + summary := SessionSummary{ + SessionID: "session-1", + AgentID: "agent-1", + TotalTokens: 10000, + ToolCounts: map[string]int{ + "read": 10, + "write": 5, + "search": 15, + }, + } + + if summary.SessionID != "session-1" { + t.Error("SessionSummary SessionID mismatch") + } + if summary.TotalTokens != 10000 { + t.Error("SessionSummary TotalTokens mismatch") + } + if summary.ToolCounts["read"] != 10 { + t.Error("SessionSummary ToolCounts[read] mismatch") + } +} diff --git a/pkg/memory/delegate/rl_types.go b/pkg/memory/delegate/rl_types.go new file mode 100644 index 000000000..279288f40 --- /dev/null +++ b/pkg/memory/delegate/rl_types.go @@ -0,0 +1,76 @@ +// Package delegate provides RL (Reinforcement Learning) type definitions +// that mirror the cortex package types to avoid circular imports. +package delegate + +import ( + "time" + + "github.com/ZanzyTHEbar/dragonscale/pkg/ids" +) + +// TaskBaseline tracks running statistics for task performance using Welford's online algorithm. +// This enables incremental calculation of mean and variance without storing all historical data. +// Mirrors cortex.TaskBaseline. +type TaskBaseline struct { + Count int + MeanTokens float64 + MeanErrors float64 + MeanUserCorrections float64 + M2Tokens float64 // sum of squares of differences from mean (for variance) + M2Errors float64 + M2UserCorrections float64 +} + +// TaskRecord represents a completed task with performance metrics. +// Mirrors cortex.TaskRecord. +type TaskRecord struct { + ID string + Description string + TokensUsed int + ToolCalls int + Errors int + UserCorrections int + Completed bool + CreatedAt time.Time +} + +// RetrievedMemoryRecord represents a memory retrieved during task execution. +// Mirrors cortex.RetrievedMemoryRecord. +type RetrievedMemoryRecord struct { + MemoryID ids.UUID + Similarity float64 + SelfReportScore *int // nullable 0-3 scale +} + +// AuditEntry represents a single audit log entry for analysis. +// Mirrors cortex.AuditEntry. +type AuditEntry struct { + ID string + Timestamp time.Time + ToolName string + ToolInput string + Success bool + ErrorMsg string + SessionID string + AgentID string +} + +// DetectedPattern represents a pattern detected from audit analysis. +// Mirrors cortex.DetectedPattern. +type DetectedPattern struct { + Type string // "correction", "discovery", "failure_pattern" + Description string + Weight float64 + Category string + SessionID string + AgentID string +} + +// SessionSummary represents token usage summary for a session. +// Mirrors cortex.SessionSummary. +type SessionSummary struct { + SessionID string + AgentID string + TotalTokens int64 + ToolCounts map[string]int +} diff --git a/pkg/memory/delegate/sqlite.go b/pkg/memory/delegate/sqlite.go index 52633288e..1b3428a61 100644 --- a/pkg/memory/delegate/sqlite.go +++ b/pkg/memory/delegate/sqlite.go @@ -8,6 +8,7 @@ import ( "time" "github.com/ZanzyTHEbar/dragonscale/pkg/ids" + "github.com/ZanzyTHEbar/dragonscale/pkg/logger" "github.com/ZanzyTHEbar/dragonscale/pkg/memory" "github.com/ZanzyTHEbar/dragonscale/pkg/memory/dag" "github.com/ZanzyTHEbar/dragonscale/pkg/memory/migrations" @@ -242,6 +243,10 @@ func (d *LibSQLDelegate) InsertRecallItem(ctx context.Context, item *memory.Reca } func recallItemToParams(item *memory.RecallItem) memsqlc.InsertRecallItemParams { + rlWeight := item.RLWeight + if rlWeight == 0 { + rlWeight = 1.0 // Default weight + } return memsqlc.InsertRecallItemParams{ ID: item.ID, AgentID: item.AgentID, @@ -253,6 +258,7 @@ func recallItemToParams(item *memory.RecallItem) memsqlc.InsertRecallItemParams DecayRate: item.DecayRate, Content: item.Content, Tags: item.Tags, + RlWeight: &rlWeight, } } @@ -264,7 +270,22 @@ func (d *LibSQLDelegate) GetRecallItem(ctx context.Context, agentID string, id i if err != nil { return nil, err } - return sqlcRecallToMemory(row), nil + // Inline conversion from GetRecallItemRow (no SuppressedAt in row) + return &memory.RecallItem{ + ID: row.ID, + AgentID: row.AgentID, + SessionKey: row.SessionKey, + Role: row.Role, + Sector: row.Sector, + Importance: row.Importance, + Salience: row.Salience, + DecayRate: row.DecayRate, + Content: row.Content, + Tags: row.Tags, + CreatedAt: row.CreatedAt, + UpdatedAt: row.UpdatedAt, + SuppressedAt: nil, // GetRecallItem query filters out suppressed items + }, nil } func (d *LibSQLDelegate) GetRecallItemsByIDs(ctx context.Context, agentID string, itemIDs []ids.UUID) (map[ids.UUID]*memory.RecallItem, error) { @@ -280,7 +301,22 @@ func (d *LibSQLDelegate) GetRecallItemsByIDs(ctx context.Context, agentID string } result := make(map[ids.UUID]*memory.RecallItem, len(rows)) for _, row := range rows { - result[row.ID] = sqlcRecallToMemory(row) + // Inline conversion from GetRecallItemsByIDsRow (no SuppressedAt in row) + result[row.ID] = &memory.RecallItem{ + ID: row.ID, + AgentID: row.AgentID, + SessionKey: row.SessionKey, + Role: row.Role, + Sector: row.Sector, + Importance: row.Importance, + Salience: row.Salience, + DecayRate: row.DecayRate, + Content: row.Content, + Tags: row.Tags, + CreatedAt: row.CreatedAt, + UpdatedAt: row.UpdatedAt, + SuppressedAt: nil, // GetRecallItemsByIDs query filters out suppressed items + } } return result, nil } @@ -303,6 +339,62 @@ func (d *LibSQLDelegate) DeleteRecallItem(ctx context.Context, agentID string, i return d.queries.DeleteRecallItem(ctx, memsqlc.DeleteRecallItemParams{ID: id, AgentID: agentID}) } +// SoftDeleteRecallItem sets suppressed_at instead of permanently deleting. +func (d *LibSQLDelegate) SoftDeleteRecallItem(ctx context.Context, agentID string, id ids.UUID) error { + if err := d.queries.SoftDeleteRecallItem(ctx, memsqlc.SoftDeleteRecallItemParams{ID: id, AgentID: agentID}); err != nil { + return err + } + // Also soft-delete associated archival chunks + return d.queries.SoftDeleteArchivalChunks(ctx, memsqlc.SoftDeleteArchivalChunksParams{RecallID: id}) +} + +// ListQuarantinedRecallItems returns recall items ready for permanent deletion. +func (d *LibSQLDelegate) ListQuarantinedRecallItems(ctx context.Context, agentID string, cutoff time.Time, limit int) ([]*memory.RecallItem, error) { + rows, err := d.queries.ListQuarantinedRecallItems(ctx, memsqlc.ListQuarantinedRecallItemsParams{ + BeforeDate: &cutoff, + Lim: int64(limit), + }) + if err != nil { + return nil, err + } + items := make([]*memory.RecallItem, 0, len(rows)) + for _, r := range rows { + items = append(items, sqlcRecallToMemory(r)) + } + return items, nil +} + +// ListQuarantinedArchivalChunks returns chunks ready for permanent deletion. +func (d *LibSQLDelegate) ListQuarantinedArchivalChunks(ctx context.Context, cutoff time.Time, limit int) ([]*memory.ArchivalChunk, error) { + rows, err := d.queries.ListQuarantinedArchivalChunks(ctx, memsqlc.ListQuarantinedArchivalChunksParams{ + BeforeDate: &cutoff, + Lim: int64(limit), + }) + if err != nil { + return nil, err + } + chunks := make([]*memory.ArchivalChunk, 0, len(rows)) + for _, r := range rows { + chunks = append(chunks, sqlcChunkToMemory(r)) + } + return chunks, nil +} + +// HardDeleteRecallItem permanently deletes a recall item. +func (d *LibSQLDelegate) HardDeleteRecallItem(ctx context.Context, agentID string, id ids.UUID) error { + return d.queries.HardDeleteRecallItem(ctx, memsqlc.HardDeleteRecallItemParams{ID: id, AgentID: agentID}) +} + +// HardDeleteArchivalChunks permanently deletes chunks for a recall item. +func (d *LibSQLDelegate) HardDeleteArchivalChunks(ctx context.Context, recallID ids.UUID) error { + return d.queries.HardDeleteArchivalChunks(ctx, memsqlc.HardDeleteArchivalChunksParams{RecallID: recallID}) +} + +// HardDeleteChunk permanently deletes a single archival chunk by ID. +func (d *LibSQLDelegate) HardDeleteChunk(ctx context.Context, id ids.UUID) error { + return d.queries.HardDeleteChunk(ctx, memsqlc.HardDeleteChunkParams{ID: id}) +} + func (d *LibSQLDelegate) ListRecallItems(ctx context.Context, agentID, sessionKey string, limit, offset int) ([]*memory.RecallItem, error) { rows, err := d.queries.ListRecallItems(ctx, memsqlc.ListRecallItemsParams{ AgentID: agentID, @@ -315,7 +407,22 @@ func (d *LibSQLDelegate) ListRecallItems(ctx context.Context, agentID, sessionKe } items := make([]*memory.RecallItem, len(rows)) for i, row := range rows { - items[i] = sqlcRecallToMemory(row) + // Inline conversion from ListRecallItemsRow (no SuppressedAt in row) + items[i] = &memory.RecallItem{ + ID: row.ID, + AgentID: row.AgentID, + SessionKey: row.SessionKey, + Role: row.Role, + Sector: row.Sector, + Importance: row.Importance, + Salience: row.Salience, + DecayRate: row.DecayRate, + Content: row.Content, + Tags: row.Tags, + CreatedAt: row.CreatedAt, + UpdatedAt: row.UpdatedAt, + SuppressedAt: nil, // ListRecallItems query filters out suppressed items + } } return items, nil } @@ -331,7 +438,22 @@ func (d *LibSQLDelegate) SearchRecallByKeyword(ctx context.Context, query, agent } items := make([]*memory.RecallItem, len(rows)) for i, row := range rows { - items[i] = sqlcRecallToMemory(row) + // Inline conversion from SearchRecallByKeywordRow (no SuppressedAt in row) + items[i] = &memory.RecallItem{ + ID: row.ID, + AgentID: row.AgentID, + SessionKey: row.SessionKey, + Role: row.Role, + Sector: row.Sector, + Importance: row.Importance, + Salience: row.Salience, + DecayRate: row.DecayRate, + Content: row.Content, + Tags: row.Tags, + CreatedAt: row.CreatedAt, + UpdatedAt: row.UpdatedAt, + SuppressedAt: nil, // SearchRecallByKeyword query filters out suppressed items + } } return items, nil } @@ -394,7 +516,18 @@ func (d *LibSQLDelegate) GetArchivalChunk(ctx context.Context, agentID string, i if err != nil { return nil, err } - return sqlcChunkToMemory(row), nil + // Inline conversion from GetArchivalChunkRow (no SuppressedAt in row) + return &memory.ArchivalChunk{ + ID: row.ID, + RecallID: row.RecallID, + ChunkIndex: int(row.ChunkIndex), + Content: row.Content, + Embedding: row.Embedding, + Source: row.Source, + Hash: row.Hash, + CreatedAt: row.CreatedAt, + SuppressedAt: nil, // GetArchivalChunk query doesn't return suppressed items + }, nil } func (d *LibSQLDelegate) ListArchivalChunks(ctx context.Context, agentID string, recallID ids.UUID) ([]*memory.ArchivalChunk, error) { @@ -404,7 +537,18 @@ func (d *LibSQLDelegate) ListArchivalChunks(ctx context.Context, agentID string, } chunks := make([]*memory.ArchivalChunk, len(rows)) for i, row := range rows { - chunks[i] = sqlcChunkToMemory(row) + // Inline conversion from ListArchivalChunksRow (no SuppressedAt in row) + chunks[i] = &memory.ArchivalChunk{ + ID: row.ID, + RecallID: row.RecallID, + ChunkIndex: int(row.ChunkIndex), + Content: row.Content, + Embedding: row.Embedding, + Source: row.Source, + Hash: row.Hash, + CreatedAt: row.CreatedAt, + SuppressedAt: nil, // ListArchivalChunks query filters out suppressed items + } } return chunks, nil } @@ -420,7 +564,18 @@ func (d *LibSQLDelegate) ListAllArchivalChunks(ctx context.Context, agentID stri } chunks := make([]*memory.ArchivalChunk, len(rows)) for i, row := range rows { - chunks[i] = sqlcChunkToMemory(row) + // Inline conversion from ListAllArchivalChunksRow (no SuppressedAt in row) + chunks[i] = &memory.ArchivalChunk{ + ID: row.ID, + RecallID: row.RecallID, + ChunkIndex: int(row.ChunkIndex), + Content: row.Content, + Embedding: row.Embedding, + Source: row.Source, + Hash: row.Hash, + CreatedAt: row.CreatedAt, + SuppressedAt: nil, // ListAllArchivalChunks includes all chunks + } } return chunks, nil } @@ -631,7 +786,20 @@ func (d *LibSQLDelegate) ListSessionMessages(ctx context.Context, agentID, sessi } items := make([]*memory.RecallItem, len(rows)) for i, row := range rows { - items[i] = sqlcRecallToMemory(row) + items[i] = &memory.RecallItem{ + ID: row.ID, + AgentID: row.AgentID, + SessionKey: row.SessionKey, + Role: row.Role, + Sector: row.Sector, + Importance: row.Importance, + Salience: row.Salience, + DecayRate: row.DecayRate, + Content: row.Content, + Tags: row.Tags, + CreatedAt: row.CreatedAt, + UpdatedAt: row.UpdatedAt, + } } return items, nil } @@ -771,35 +939,417 @@ func (d *LibSQLDelegate) PersistDAG(ctx context.Context, agentID, sessionKey str return dag.PersistDAG(ctx, d.db, d.queries, agentID, sessionKey, snap) } +// --- Memory Edges (via sqlc) --- + +func (d *LibSQLDelegate) InsertMemoryEdge(ctx context.Context, edge *memory.MemoryEdge) error { + row, err := d.queries.InsertMemoryEdge(ctx, memsqlc.InsertMemoryEdgeParams{ + FromID: edge.FromID, + ToID: edge.ToID, + EdgeType: string(edge.EdgeType), + Weight: edge.Weight, + }) + if err != nil { + return err + } + edge.ID = row.ID + edge.CreatedAt = row.CreatedAt + return nil +} + +func (d *LibSQLDelegate) ListMemoryEdges(ctx context.Context, memoryID ids.UUID) ([]*memory.MemoryEdge, error) { + rows, err := d.queries.ListMemoryEdgesForItem(ctx, memsqlc.ListMemoryEdgesForItemParams{ + MemoryID: memoryID, + Lim: 1000, + }) + if err != nil { + return nil, err + } + edges := make([]*memory.MemoryEdge, 0, len(rows)) + for _, r := range rows { + edges = append(edges, sqlcEdgeToMemory(r)) + } + return edges, nil +} + +func (d *LibSQLDelegate) CountMemoryEdgesForItem(ctx context.Context, memoryID ids.UUID) (int, error) { + count, err := d.queries.CountMemoryEdgesForItem(ctx, memsqlc.CountMemoryEdgesForItemParams{ + MemoryID: memoryID, + }) + return int(count), err +} + +// ListRecallItemsForConsolidation returns recall items with embeddings for similarity comparison. +// Used by the Cortex consolidation task to build the memory graph. +func (d *LibSQLDelegate) ListRecallItemsForConsolidation(ctx context.Context, agentID string, cutoff time.Time, limit int) ([]*memory.RecallItem, error) { + rows, err := d.queries.ListRecallItemsForConsolidation(ctx, memsqlc.ListRecallItemsForConsolidationParams{ + Cutoff: cutoff, + AgentID: agentID, + Lim: int64(limit), + }) + if err != nil { + return nil, err + } + + items := make([]*memory.RecallItem, 0, len(rows)) + for _, r := range rows { + item := &memory.RecallItem{ + ID: r.ID, + AgentID: r.AgentID, + SessionKey: r.SessionKey, + Role: r.Role, + Sector: r.Sector, + Importance: r.Importance, + Salience: r.Salience, + DecayRate: r.DecayRate, + Content: r.Content, + Tags: r.Tags, + CreatedAt: r.CreatedAt, + UpdatedAt: r.UpdatedAt, + Embedding: r.Embedding, + } + items = append(items, item) + } + return items, nil +} + +// --- RL (Reinforcement Learning) Store Methods --- + +// GetTaskBaseline retrieves the baseline statistics for an agent. +// Implements cortex.RLStore interface. +func (d *LibSQLDelegate) GetTaskBaseline(ctx context.Context, agentID string) (*TaskBaseline, error) { + row, err := d.queries.GetTaskBaseline(ctx, memsqlc.GetTaskBaselineParams{AgentID: agentID}) + if err == sql.ErrNoRows { + // Return nil baseline for new agents - cold start handling + return nil, nil + } + if err != nil { + return nil, err + } + + baseline := &TaskBaseline{} + if row.Count != nil { + baseline.Count = int(*row.Count) + } + if row.MeanTokens != nil { + baseline.MeanTokens = float64(*row.MeanTokens) + } + if row.MeanErrors != nil { + baseline.MeanErrors = *row.MeanErrors + } + if row.MeanUserCorrections != nil { + baseline.MeanUserCorrections = *row.MeanUserCorrections + } + if row.M2Tokens != nil { + baseline.M2Tokens = *row.M2Tokens + } + if row.M2Errors != nil { + baseline.M2Errors = *row.M2Errors + } + if row.M2UserCorrections != nil { + baseline.M2UserCorrections = *row.M2UserCorrections + } + return baseline, nil +} + +// UpdateTaskBaseline saves the baseline statistics for an agent. +// Implements cortex.RLStore interface. +func (d *LibSQLDelegate) UpdateTaskBaseline(ctx context.Context, agentID string, baseline *TaskBaseline) error { + count := int64(baseline.Count) + meanTokens := int64(baseline.MeanTokens) + meanErrors := baseline.MeanErrors + meanUserCorrections := baseline.MeanUserCorrections + m2Tokens := baseline.M2Tokens + m2Errors := baseline.M2Errors + m2UserCorrections := baseline.M2UserCorrections + + return d.queries.UpdateTaskBaseline(ctx, memsqlc.UpdateTaskBaselineParams{ + AgentID: agentID, + Count: &count, + MeanTokens: &meanTokens, + MeanErrors: &meanErrors, + MeanUserCorrections: &meanUserCorrections, + M2Tokens: &m2Tokens, + M2Errors: &m2Errors, + M2UserCorrections: &m2UserCorrections, + }) +} + +// UpdateMemoryWeight updates the RL weight and credit for a specific memory. +// Implements cortex.RLStore interface. +func (d *LibSQLDelegate) UpdateMemoryWeight(ctx context.Context, memoryID ids.UUID, weight, credit float64) error { + rlWeight := weight + rlCredit := credit + // Get the agent_id from the memory item first + item, err := d.GetRecallItem(ctx, "", memoryID) + if err != nil { + return err + } + if item == nil { + return fmt.Errorf("memory item not found: %s", memoryID) + } + return d.queries.UpdateMemoryWeight(ctx, memsqlc.UpdateMemoryWeightParams{ + RlWeight: &rlWeight, + RlCredit: &rlCredit, + ID: memoryID, + AgentID: item.AgentID, + }) +} + +// UpdateMemorySelfReport updates the self-reported score for a memory. +// Implements cortex.RLStore interface. +func (d *LibSQLDelegate) UpdateMemorySelfReport(ctx context.Context, memoryID ids.UUID, score int) error { + selfReportScore := int64(score) + // Get the agent_id from the memory item first + item, err := d.GetRecallItem(ctx, "", memoryID) + if err != nil { + return err + } + if item == nil { + return fmt.Errorf("memory item not found: %s", memoryID) + } + return d.queries.UpdateMemorySelfReportScore(ctx, memsqlc.UpdateMemorySelfReportScoreParams{ + SelfReportScore: &selfReportScore, + ID: memoryID, + AgentID: item.AgentID, + }) +} + +// GetCompletedTasks returns tasks completed since the given time. +// Implements cortex.RLStore interface. +// Note: This is a placeholder implementation - actual task storage needs to be defined. +func (d *LibSQLDelegate) GetCompletedTasks(ctx context.Context, since time.Time) ([]TaskRecord, error) { + // TODO: Implement actual task retrieval from jobs or runs tables + // For now, return empty list + return []TaskRecord{}, nil +} + +// GetRetrievedMemories returns memories retrieved during a task. +// Implements cortex.RLStore interface. +// Note: This is a placeholder implementation - actual retrieval tracking needs to be defined. +func (d *LibSQLDelegate) GetRetrievedMemories(ctx context.Context, taskID string) ([]RetrievedMemoryRecord, error) { + // TODO: Implement actual retrieved memory tracking + // For now, return empty list + return []RetrievedMemoryRecord{}, nil +} + +// --- Audit Analysis Store Methods --- + +// GetRecentAuditEntries returns audit entries since the given time. +// Implements cortex.AuditAnalysisStore interface. +func (d *LibSQLDelegate) GetRecentAuditEntries(ctx context.Context, since time.Time) ([]AuditEntry, error) { + // Get all audit entries and filter by time + rows, err := d.queries.ListAuditEntries(ctx, memsqlc.ListAuditEntriesParams{ + AgentID: "", // Get all agents + Lim: 10000, + }) + if err != nil { + return nil, err + } + + var entries []AuditEntry + for _, row := range rows { + if row.CreatedAt.After(since) { + entry := AuditEntry{ + ID: row.ID.String(), + Timestamp: row.CreatedAt, + ToolName: row.Action, // Using action as tool name proxy + ToolInput: "", + Success: true, // Default to success + SessionID: row.SessionKey, + AgentID: row.AgentID, + } + if row.Input != nil { + entry.ToolInput = *row.Input + } + entries = append(entries, entry) + } + } + return entries, nil +} + +// StoreDetectedPattern stores a detected pattern as a recall item. +// Implements cortex.AuditAnalysisStore interface. +func (d *LibSQLDelegate) StoreDetectedPattern(ctx context.Context, pattern DetectedPattern) error { + item := &memory.RecallItem{ + ID: ids.New(), + AgentID: pattern.AgentID, + SessionKey: pattern.SessionID, + Role: "system", + Sector: memory.SectorReflective, + Importance: pattern.Weight, + Salience: pattern.Weight, + Content: pattern.Description, + Tags: fmt.Sprintf("audit,%s,%s", pattern.Type, pattern.Category), + CreatedAt: time.Now(), + UpdatedAt: time.Now(), + } + return d.InsertRecallItem(ctx, item) +} + +// GetHighTokenSessions returns sessions with token usage above threshold. +// Implements cortex.AuditAnalysisStore interface. +// Note: This is a placeholder - actual token tracking needs to be implemented. +func (d *LibSQLDelegate) GetHighTokenSessions(ctx context.Context, minTokens int64) ([]SessionSummary, error) { + // TODO: Implement token-based session filtering when token tracking is available + return []SessionSummary{}, nil +} + +// --- Batch Operations for Cortex Tasks (via sqlc) --- + +// DecayRecallImportance applies multiplicative decay to the oldest recall items +// whose importance exceeds the floor. Uses sqlc-generated query string with +// raw ExecContext to preserve RowsAffected for observability. +func (d *LibSQLDelegate) DecayRecallImportance(ctx context.Context, factor, floor float64, batchSize int) (int64, error) { + result, err := d.db.ExecContext(ctx, memsqlc.DecayRecallImportanceBatch, factor, floor, int64(batchSize)) + if err != nil { + return 0, err + } + return result.RowsAffected() +} + +// CountArchivalChunksWithoutEmbedding returns the count of chunks with NULL embeddings. +func (d *LibSQLDelegate) CountArchivalChunksWithoutEmbedding(ctx context.Context) (int, error) { + count, err := d.queries.CountArchivalChunksWithoutEmbedding(ctx) + return int(count), err +} + +// BackfillArchivalEmbeddings finds chunks without embeddings, calls embedFn for each, +// and writes the embedding back via sqlc UpdateArchivalChunkEmbedding. +func (d *LibSQLDelegate) BackfillArchivalEmbeddings(ctx context.Context, batchSize int, embedFn func(ctx context.Context, text string) ([]float32, error)) (int, error) { + chunks, err := d.queries.ListArchivalChunksWithoutEmbedding(ctx, memsqlc.ListArchivalChunksWithoutEmbeddingParams{ + Lim: int64(batchSize), + }) + if err != nil { + return 0, err + } + + processed := 0 + for _, c := range chunks { + vec, err := embedFn(ctx, c.Content) + if err != nil { + logger.WarnCF("cortex", "Embedding failed for chunk", map[string]interface{}{ + "chunk_id": c.ID.String(), + "error": err.Error(), + }) + continue + } + emb := memory.Embedding(vec) + if err := d.queries.UpdateArchivalChunkEmbedding(ctx, memsqlc.UpdateArchivalChunkEmbeddingParams{ + Embedding: emb, + ID: c.ID, + }); err != nil { + logger.WarnCF("cortex", "Failed to update chunk embedding", map[string]interface{}{ + "chunk_id": c.ID.String(), + "error": err.Error(), + }) + continue + } + processed++ + } + return processed, nil +} + +// --- Immutable Messages (via sqlc) --- + +func (d *LibSQLDelegate) InsertImmutableMessage(ctx context.Context, msg *memory.ImmutableMessage) error { + row, err := d.queries.InsertImmutableMessage(ctx, memsqlc.InsertImmutableMessageParams{ + ID: msg.ID, + SessionKey: msg.SessionKey, + Role: msg.Role, + Content: msg.Content, + ToolCallID: msg.ToolCallID, + ToolCalls: msg.ToolCalls, + TokenEstimate: int64(msg.TokenEstimate), + }) + if err != nil { + return err + } + msg.CreatedAt = row.CreatedAt + return nil +} + +func (d *LibSQLDelegate) GetImmutableMessage(ctx context.Context, id ids.UUID) (*memory.ImmutableMessage, error) { + row, err := d.queries.GetImmutableMessage(ctx, memsqlc.GetImmutableMessageParams{ID: id}) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, err + } + return sqlcImmutableToMemory(row), nil +} + +func (d *LibSQLDelegate) ListImmutableMessages(ctx context.Context, sessionKey string, limit, offset int) ([]*memory.ImmutableMessage, error) { + rows, err := d.queries.ListImmutableMessages(ctx, memsqlc.ListImmutableMessagesParams{ + SessionKey: sessionKey, + Lim: int64(limit), + Off: int64(offset), + }) + if err != nil { + return nil, err + } + msgs := make([]*memory.ImmutableMessage, 0, len(rows)) + for _, r := range rows { + msgs = append(msgs, sqlcImmutableToMemory(r)) + } + return msgs, nil +} + // --- Conversion helpers --- -func sqlcRecallToMemory(row memsqlc.RecallItem) *memory.RecallItem { +func sqlcRecallToMemory(row memsqlc.ListQuarantinedRecallItemsRow) *memory.RecallItem { return &memory.RecallItem{ - ID: row.ID, - AgentID: row.AgentID, - SessionKey: row.SessionKey, - Role: row.Role, - Sector: row.Sector, // already memory.Sector via sqlc override - Importance: row.Importance, - Salience: row.Salience, - DecayRate: row.DecayRate, - Content: row.Content, - Tags: row.Tags, - CreatedAt: row.CreatedAt, - UpdatedAt: row.UpdatedAt, + ID: row.ID, + AgentID: row.AgentID, + SessionKey: row.SessionKey, + Role: row.Role, + Sector: row.Sector, // already memory.Sector via sqlc override + Importance: row.Importance, + Salience: row.Salience, + DecayRate: row.DecayRate, + Content: row.Content, + Tags: row.Tags, + CreatedAt: row.CreatedAt, + UpdatedAt: row.UpdatedAt, + SuppressedAt: row.SuppressedAt, } } func sqlcChunkToMemory(row memsqlc.ArchivalChunk) *memory.ArchivalChunk { return &memory.ArchivalChunk{ - ID: row.ID, - RecallID: row.RecallID, - ChunkIndex: int(row.ChunkIndex), - Content: row.Content, - Embedding: row.Embedding, // memory.Embedding with auto-deserialization via Scanner - Source: row.Source, - Hash: row.Hash, - CreatedAt: row.CreatedAt, + ID: row.ID, + RecallID: row.RecallID, + ChunkIndex: int(row.ChunkIndex), + Content: row.Content, + Embedding: row.Embedding, // memory.Embedding with auto-deserialization via Scanner + Source: row.Source, + Hash: row.Hash, + CreatedAt: row.CreatedAt, + SuppressedAt: row.SuppressedAt, + } +} + +func sqlcImmutableToMemory(row memsqlc.ImmutableMessage) *memory.ImmutableMessage { + return &memory.ImmutableMessage{ + ID: row.ID, + SessionKey: row.SessionKey, + Role: row.Role, + Content: row.Content, + ToolCallID: row.ToolCallID, + ToolCalls: row.ToolCalls, + TokenEstimate: int(row.TokenEstimate), + CreatedAt: row.CreatedAt, + } +} + +func sqlcEdgeToMemory(row memsqlc.MemoryEdge) *memory.MemoryEdge { + return &memory.MemoryEdge{ + ID: row.ID, + FromID: row.FromID, + ToID: row.ToID, + EdgeType: memory.EdgeType(row.EdgeType), + Weight: row.Weight, + CreatedAt: row.CreatedAt, } } diff --git a/pkg/memory/memory.go b/pkg/memory/memory.go index 4d8ffc3f7..08a64a86d 100644 --- a/pkg/memory/memory.go +++ b/pkg/memory/memory.go @@ -31,6 +31,18 @@ const ( SectorReflective Sector = "reflective" // Meta-observations, self-assessments ) +// Category classifies memories by their semantic role for RL weight assignment. +type Category string + +const ( + CategoryUnknown Category = "unknown" + CategoryFact Category = "fact" + CategoryInsight Category = "insight" + CategoryCorrection Category = "correction" + CategoryDiscovery Category = "discovery" + CategoryUserInput Category = "user_input" +) + // --- Embedding type (F32_BLOB wire format) --- // Embedding is a float32 vector that transparently serializes to/from @@ -84,30 +96,41 @@ func (e *Embedding) Scan(src interface{}) error { // RecallItem is a memory entry in the warm tier. type RecallItem struct { - ID ids.UUID - AgentID string - SessionKey string - Role string // "system", "user", "assistant", "tool" - Sector Sector - Importance float64 // [0, 1] - Salience float64 // [0, 1] - DecayRate float64 // exponential decay constant - Content string - Tags string // comma-separated - CreatedAt time.Time - UpdatedAt time.Time + ID ids.UUID + AgentID string + SessionKey string + Role string // "system", "user", "assistant", "tool" + Sector Sector + Importance float64 // [0, 1] + Salience float64 // [0, 1] + DecayRate float64 // exponential decay constant + Content string + Tags string // comma-separated + CreatedAt time.Time + UpdatedAt time.Time + SuppressedAt *time.Time // nil if not soft-deleted + // Embedding is the vector representation of this item (populated during + // consolidation queries; nil for standard CRUD operations). + Embedding Embedding + // RL fields for Memelord reinforcement learning + RLWeight float64 // current weight for credit assignment (default 1.0) + RLCredit float64 // accumulated credit for this memory + SelfReportScore *int // self-reported usefulness score (0-3 scale) + TaskRetrievalCount int // how many times retrieved for tasks + Category Category // semantic category for initial weight assignment } // ArchivalChunk is an embedded chunk in the cold tier. type ArchivalChunk struct { - ID ids.UUID - RecallID ids.UUID // FK to RecallItem or standalone - ChunkIndex int - Content string - Embedding Embedding // F32_BLOB with auto-serialization via Valuer/Scanner - Source string - Hash string - CreatedAt time.Time + ID ids.UUID + RecallID ids.UUID // FK to RecallItem or standalone + ChunkIndex int + Content string + Embedding Embedding // F32_BLOB with auto-serialization via Valuer/Scanner + Source string + Hash string + CreatedAt time.Time + SuppressedAt *time.Time // nil if not soft-deleted } // WorkingContext is the hot-tier mutable buffer. @@ -129,6 +152,43 @@ type MemorySummary struct { CreatedAt time.Time } +// ImmutableMessage is a verbatim, append-only record of every message +// exchanged during a session. Unlike recall_items (which can be truncated +// during compaction), immutable messages are never modified or deleted. +// They serve as the source of truth for lossless context recovery. +type ImmutableMessage struct { + ID ids.UUID + SessionKey string + Role string // "user", "assistant", "tool", "system" + Content string + ToolCallID string // non-empty for tool-result messages + ToolCalls string // JSON-encoded tool calls for assistant messages + TokenEstimate int + CreatedAt time.Time +} + +// EdgeType classifies the relationship between two memory items. +type EdgeType string + +const ( + EdgeRelatedTo EdgeType = "related_to" + EdgeUpdates EdgeType = "updates" + EdgeContradicts EdgeType = "contradicts" + EdgeCausedBy EdgeType = "caused_by" + EdgeResultOf EdgeType = "result_of" + EdgePartOf EdgeType = "part_of" +) + +// MemoryEdge represents a typed, weighted relationship between two memory items. +type MemoryEdge struct { + ID int64 + FromID ids.UUID + ToID ids.UUID + EdgeType EdgeType + Weight float64 + CreatedAt time.Time +} + // SearchResult represents a result from hybrid retrieval. type SearchResult struct { ID ids.UUID @@ -268,6 +328,13 @@ type MemoryReader interface { 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) + ListImmutableMessages(ctx context.Context, sessionKey string, limit, offset int) ([]*ImmutableMessage, error) + GetImmutableMessage(ctx context.Context, id ids.UUID) (*ImmutableMessage, error) + ListMemoryEdges(ctx context.Context, memoryID ids.UUID) ([]*MemoryEdge, error) + CountMemoryEdgesForItem(ctx context.Context, memoryID ids.UUID) (int, error) + // ListRecallItemsForConsolidation returns recent recall items with embeddings + // for similarity comparison during consolidation. Used by the Cortex scheduler. + ListRecallItemsForConsolidation(ctx context.Context, agentID string, cutoff time.Time, limit int) ([]*RecallItem, error) HasVectorSearch() bool HasFTS() bool } @@ -289,6 +356,24 @@ type MemoryWriter interface { DeleteDocument(ctx context.Context, agentID, name string) error InsertAuditEntry(ctx context.Context, entry *AuditEntry) error InsertAuditEntryBatch(ctx context.Context, entries []*AuditEntry) error + InsertImmutableMessage(ctx context.Context, msg *ImmutableMessage) error + InsertMemoryEdge(ctx context.Context, edge *MemoryEdge) error + // SoftDeleteRecallItem sets suppressed_at instead of permanently deleting. + // The item enters a 30-day quarantine before hard deletion. + SoftDeleteRecallItem(ctx context.Context, agentID string, id ids.UUID) error +} + +// PruneStore is the interface for pruning quarantined items. +// Implemented by LibSQLDelegate for the Cortex prune task. +type PruneStore interface { + // ListQuarantinedRecallItems returns recall items ready for permanent deletion. + ListQuarantinedRecallItems(ctx context.Context, agentID string, cutoff time.Time, limit int) ([]*RecallItem, error) + // ListQuarantinedArchivalChunks returns chunks ready for permanent deletion. + ListQuarantinedArchivalChunks(ctx context.Context, cutoff time.Time, limit int) ([]*ArchivalChunk, error) + // HardDeleteRecallItem permanently deletes a recall item. + HardDeleteRecallItem(ctx context.Context, agentID string, id ids.UUID) error + // HardDeleteArchivalChunks permanently deletes chunks for a recall item. + HardDeleteArchivalChunks(ctx context.Context, recallID ids.UUID) error } // MemoryDelegate is the full-capability interface for memory operations. diff --git a/pkg/memory/migrate_sessions_test.go b/pkg/memory/migrate_sessions_test.go index b6672706d..a4800bef6 100644 --- a/pkg/memory/migrate_sessions_test.go +++ b/pkg/memory/migrate_sessions_test.go @@ -112,6 +112,28 @@ func (m *mockDelegate) ListAuditEntriesByAction(_ context.Context, _, _ string, return nil, nil } func (m *mockDelegate) CountAuditEntries(_ context.Context, _ string) (int, error) { return 0, nil } +func (m *mockDelegate) InsertImmutableMessage(_ context.Context, _ *ImmutableMessage) error { + return nil +} +func (m *mockDelegate) ListImmutableMessages(_ context.Context, _ string, _, _ int) ([]*ImmutableMessage, error) { + return nil, nil +} +func (m *mockDelegate) GetImmutableMessage(_ context.Context, _ ids.UUID) (*ImmutableMessage, error) { + return nil, nil +} +func (m *mockDelegate) InsertMemoryEdge(_ context.Context, _ *MemoryEdge) error { return nil } +func (m *mockDelegate) ListMemoryEdges(_ context.Context, _ ids.UUID) ([]*MemoryEdge, error) { + return nil, nil +} +func (m *mockDelegate) CountMemoryEdgesForItem(_ context.Context, _ ids.UUID) (int, error) { + return 0, nil +} +func (m *mockDelegate) ListRecallItemsForConsolidation(_ context.Context, _ string, _ time.Time, _ int) ([]*RecallItem, error) { + return nil, nil +} +func (m *mockDelegate) SoftDeleteRecallItem(_ context.Context, _ string, _ ids.UUID) error { + return nil +} func writeSessionFile(t *testing.T, dir, name string, sess SessionFile) { t.Helper() diff --git a/pkg/memory/store/memory_store.go b/pkg/memory/store/memory_store.go index d0dc3ddc7..dd2042bc1 100644 --- a/pkg/memory/store/memory_store.go +++ b/pkg/memory/store/memory_store.go @@ -93,6 +93,9 @@ func New(delegate memory.MemoryDelegate, chunker memory.Chunker, embedder memory } } +// Embedder returns the configured EmbeddingProvider, or nil if embeddings are disabled. +func (m *MemoryStore) Embedder() memory.EmbeddingProvider { return m.embedder } + // 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) { @@ -165,9 +168,32 @@ func (m *MemoryStore) StoreRecall(ctx context.Context, item *memory.RecallItem) if item.ID.IsZero() { item.ID = ids.New() } + // Set initial RL weight based on category if not already set + if item.RLWeight == 0 { + item.RLWeight = m.initialWeightByCategory(item.Category) + } return m.delegate.InsertRecallItem(ctx, item) } +// initialWeightByCategory returns the initial RL weight based on memory category. +// Higher weights are assigned to categories that indicate higher value memories. +func (m *MemoryStore) initialWeightByCategory(category memory.Category) float64 { + switch category { + case memory.CategoryCorrection: + return 1.5 // High priority - corrections are valuable + case memory.CategoryDiscovery: + return 1.3 // Good insights - discoveries are useful + case memory.CategoryUserInput: + return 2.5 // User corrections highest priority + case memory.CategoryInsight: + return 1.1 // Slightly above baseline + case memory.CategoryFact: + return 1.0 // Baseline weight + default: + return 1.0 // Unknown category defaults to baseline + } +} + func (m *MemoryStore) GetRecall(ctx context.Context, id ids.UUID) (*memory.RecallItem, error) { return m.delegate.GetRecallItem(ctx, m.agentID, id) }