package rag_test import ( "context" "fmt" "iter" "strings" "testing" "github.com/VictorVargas/rony-llm-agent/pkg/llm" "github.com/VictorVargas/rony-llm-agent/pkg/rag" "github.com/VictorVargas/rony-llm-agent/pkg/rag/embeddings" ) func TestMemory_Add_StoresMemoryTypeInMetadata(t *testing.T) { var gotMeta map[string]string m, err := rag.New(rag.Config{ Backend: &mockBackend{ upsertFunc: func(_ context.Context, _ string, _ []float32, _ string, metadata map[string]string) error { gotMeta = metadata return nil }, }, Embedder: &embeddings.MockEmbedder{}, }) if err != nil { t.Fatalf("unexpected error: %v", err) } if err := m.Add(context.Background(), rag.Fragment{Content: "an event", Type: rag.MemoryEpisodic}); err != nil { t.Fatalf("unexpected error: %v", err) } if gotMeta["memory_type"] != "episodic" { t.Fatalf("expected memory_type=episodic in metadata, got %q", gotMeta["memory_type"]) } } func TestMemory_Add_DefaultsToProcedural(t *testing.T) { var gotMeta map[string]string m, err := rag.New(rag.Config{ Backend: &mockBackend{ upsertFunc: func(_ context.Context, _ string, _ []float32, _ string, metadata map[string]string) error { gotMeta = metadata return nil }, }, Embedder: &embeddings.MockEmbedder{}, }) if err != nil { t.Fatalf("unexpected error: %v", err) } if err := m.Add(context.Background(), rag.Fragment{Content: "a process"}); err != nil { t.Fatalf("unexpected error: %v", err) } if gotMeta["memory_type"] != "procedural" { t.Fatalf("expected untyped fragments to default to procedural, got %q", gotMeta["memory_type"]) } } func TestMemory_SearchByType_FiltersAndTreatsLegacyAsProcedural(t *testing.T) { backendResults := []rag.SearchResult{ {ID: "1", Content: "episode", Metadata: map[string]string{"memory_type": "episodic"}}, {ID: "2", Content: "fact", Metadata: map[string]string{"memory_type": "semantic"}}, {ID: "3", Content: "legacy process", Metadata: map[string]string{}}, // pre-taxonomy fragment {ID: "4", Content: "typed process", Metadata: map[string]string{"memory_type": "procedural"}}, } m, err := rag.New(rag.Config{ Backend: &mockBackend{ searchFunc: func(_ context.Context, _ string, _ []float32, _ int) ([]rag.SearchResult, error) { return backendResults, nil }, }, Embedder: &embeddings.MockEmbedder{}, }) if err != nil { t.Fatalf("unexpected error: %v", err) } got, err := m.SearchByType(context.Background(), "q", 10, rag.MemoryProcedural) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(got) != 2 || got[0].ID != "3" || got[1].ID != "4" { t.Fatalf("expected legacy + typed procedural fragments, got %+v", got) } if got[0].Type != rag.MemoryProcedural { t.Fatalf("expected legacy fragment to surface as procedural, got %q", got[0].Type) } episodes, err := m.SearchByType(context.Background(), "q", 10, rag.MemoryEpisodic) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(episodes) != 1 || episodes[0].ID != "1" { t.Fatalf("expected only the episodic fragment, got %+v", episodes) } all, err := m.SearchByType(context.Background(), "q", 10) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(all) != 4 { t.Fatalf("expected no type restriction to return everything, got %d", len(all)) } } func TestMemory_SearchByType_OverfetchesWhenFiltering(t *testing.T) { var gotTopK int m, err := rag.New(rag.Config{ Backend: &mockBackend{ searchFunc: func(_ context.Context, _ string, _ []float32, topK int) ([]rag.SearchResult, error) { gotTopK = topK return nil, nil }, }, Embedder: &embeddings.MockEmbedder{}, }) if err != nil { t.Fatalf("unexpected error: %v", err) } if _, err := m.SearchByType(context.Background(), "q", 5, rag.MemoryEpisodic); err != nil { t.Fatalf("unexpected error: %v", err) } if gotTopK <= 5 { t.Fatalf("expected the backend to be asked for more than topK candidates when filtering, got %d", gotTopK) } if _, err := m.Search(context.Background(), "q", 5); err != nil { t.Fatalf("unexpected error: %v", err) } if gotTopK != 5 { t.Fatalf("expected unfiltered search to request exactly topK, got %d", gotTopK) } } // captureLLM is a minimal llm.LLMClient stub for capture tests. type captureLLM struct { response string err error gotUser string } func (c *captureLLM) Generate(_ context.Context, req llm.CompletionRequest) (llm.CompletionResponse, error) { for _, m := range req.Messages { if m.Role == llm.RoleUser { c.gotUser = m.Content } } if c.err != nil { return llm.CompletionResponse{}, c.err } return llm.CompletionResponse{Content: c.response}, nil } func (c *captureLLM) Stream(_ context.Context, _ llm.CompletionRequest) iter.Seq2[llm.StreamChunk, error] { return func(func(llm.StreamChunk, error) bool) {} } func (c *captureLLM) Name() string { return "capture-stub" } func (c *captureLLM) Capabilities() llm.ProviderCapabilities { return llm.ProviderCapabilities{} } func TestEpisodeCapture_SavesEpisodicSummary(t *testing.T) { var saved rag.Fragment backend := &mockBackend{ upsertFunc: func(_ context.Context, id string, _ []float32, content string, metadata map[string]string) error { saved = rag.Fragment{ID: id, Content: content, Metadata: metadata} return nil }, } mem, err := rag.New(rag.Config{Backend: backend, Embedder: &embeddings.MockEmbedder{}}) if err != nil { t.Fatalf("unexpected error: %v", err) } stub := &captureLLM{response: " The user asked how to deploy and the assistant explained the release steps. "} cap := &rag.EpisodeCapture{Memory: mem, LLM: stub, ProjectID: "proj1"} err = cap.Capture(context.Background(), "how do I deploy?", "You run make release...", "read", "bash") if err != nil { t.Fatalf("unexpected error: %v", err) } if saved.Content != "The user asked how to deploy and the assistant explained the release steps." { t.Fatalf("expected trimmed summary as content, got %q", saved.Content) } if saved.Metadata["memory_type"] != "episodic" { t.Fatalf("expected episodic type, got %q", saved.Metadata["memory_type"]) } if saved.Metadata["tools"] != "read,bash" { t.Fatalf("expected tools metadata, got %q", saved.Metadata["tools"]) } if saved.Metadata["project_id"] != "proj1" { t.Fatalf("expected project_id metadata, got %q", saved.Metadata["project_id"]) } if !strings.Contains(stub.gotUser, "how do I deploy?") { t.Fatalf("expected the turn transcript to reach the LLM, got %q", stub.gotUser) } } func TestEpisodeCapture_SkipsEmptyTurnsAndFailures(t *testing.T) { upserts := 0 backend := &mockBackend{ upsertFunc: func(_ context.Context, _ string, _ []float32, _ string, _ map[string]string) error { upserts++ return nil }, } mem, err := rag.New(rag.Config{Backend: backend, Embedder: &embeddings.MockEmbedder{}}) if err != nil { t.Fatalf("unexpected error: %v", err) } cap := &rag.EpisodeCapture{Memory: mem, LLM: &captureLLM{response: "summary"}, ProjectID: "p"} if err := cap.Capture(context.Background(), "", "reply"); err == nil { t.Fatal("expected error for empty user input") } if err := cap.Capture(context.Background(), "input", " "); err == nil { t.Fatal("expected error for empty assistant reply") } failing := &rag.EpisodeCapture{Memory: mem, LLM: &captureLLM{err: fmt.Errorf("llm down")}, ProjectID: "p"} if err := failing.Capture(context.Background(), "input", "reply"); err == nil { t.Fatal("expected error when the LLM fails") } empty := &rag.EpisodeCapture{Memory: mem, LLM: &captureLLM{response: " "}, ProjectID: "p"} if err := empty.Capture(context.Background(), "input", "reply"); err == nil { t.Fatal("expected error for an empty summary") } if upserts != 0 { t.Fatalf("expected nothing to be saved on failures, got %d upserts", upserts) } }