rony-llm-agent/pkg/rag/taxonomy_test.go

235 lines
7.7 KiB
Go
Raw Permalink Normal View History

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)
}
}