package rag_test import ( "context" "fmt" "testing" "github.com/VictorVargas/rony-llm-agent/pkg/rag" "github.com/VictorVargas/rony-llm-agent/pkg/rag/embeddings" ) func TestNew_Memory(t *testing.T) { _, err := rag.New(rag.Config{}) if err == nil { t.Fatal("expected error for missing backend") } _, err = rag.New(rag.Config{ Backend: &mockBackend{}, }) if err == nil { t.Fatal("expected error for missing embedder") } m, err := rag.New(rag.Config{ Backend: &mockBackend{}, Embedder: &embeddings.MockEmbedder{}, }) if err != nil { t.Fatalf("unexpected error: %v", err) } if m == nil { t.Fatal("expected non-nil memory") } } func TestMemory_Add(t *testing.T) { m, err := rag.New(rag.Config{ Backend: &mockBackend{}, Embedder: &embeddings.MockEmbedder{}, }) if err != nil { t.Fatalf("unexpected error: %v", err) } err = m.Add(context.Background(), rag.Fragment{ Content: "test content", }) if err != nil { t.Fatalf("unexpected error: %v", err) } } func TestMemory_Add_PassesContentToBackend(t *testing.T) { var gotContent string m, err := rag.New(rag.Config{ Backend: &mockBackend{ upsertFunc: func(_ context.Context, _ string, _ []float32, content string, _ map[string]string) error { gotContent = content return nil }, }, Embedder: &embeddings.MockEmbedder{}, }) if err != nil { t.Fatalf("unexpected error: %v", err) } err = m.Add(context.Background(), rag.Fragment{Content: "remember this process"}) if err != nil { t.Fatalf("unexpected error: %v", err) } if gotContent != "remember this process" { t.Fatalf("expected backend to receive the fragment content, got %q", gotContent) } } func TestMemory_Add_EmbeddingErrorStillSavesWithNoVector(t *testing.T) { var gotVector []float32 sawCall := false m, err := rag.New(rag.Config{ Backend: &mockBackend{ upsertFunc: func(_ context.Context, _ string, vector []float32, _ string, _ map[string]string) error { sawCall = true gotVector = vector return nil }, }, Embedder: &embeddings.MockEmbedder{EmbedFunc: func(ctx context.Context, text string) ([]float32, error) { return nil, fmt.Errorf("embed error") }}, }) if err != nil { t.Fatalf("unexpected error: %v", err) } err = m.Add(context.Background(), rag.Fragment{ Content: "test content", }) if err != nil { t.Fatalf("expected Add to succeed so the backend can still index the text lexically, got: %v", err) } if !sawCall { t.Fatal("expected the backend to still be called despite the embedding failure") } if len(gotVector) != 0 { t.Fatalf("expected no vector to be passed through, got %v", gotVector) } } func TestMemory_Search(t *testing.T) { m, err := rag.New(rag.Config{ Backend: &mockBackend{ searchFunc: func(ctx context.Context, query string, queryVector []float32, topK int) ([]rag.SearchResult, error) { return []rag.SearchResult{ {ID: "1", Content: "result 1", Score: 0.9}, }, nil }, }, Embedder: &embeddings.MockEmbedder{}, }) if err != nil { t.Fatalf("unexpected error: %v", err) } results, err := m.Search(context.Background(), "test query", 5) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(results) != 1 { t.Fatalf("expected 1 result, got %d", len(results)) } if results[0].Content != "result 1" { t.Errorf("expected 'result 1', got %q", results[0].Content) } } func TestMemory_ForgetAll(t *testing.T) { forgetAllCalled := false m, err := rag.New(rag.Config{ Backend: &mockBackend{ forgetAllFunc: func(ctx context.Context) error { forgetAllCalled = true return nil }, }, Embedder: &embeddings.MockEmbedder{}, }) if err != nil { t.Fatalf("unexpected error: %v", err) } err = m.ForgetAll(context.Background()) if err != nil { t.Fatalf("unexpected error: %v", err) } if !forgetAllCalled { t.Fatal("expected forgetAll to be called") } } func TestMemory_Search_EmbeddingErrorFallsBackToLexicalSearch(t *testing.T) { var gotQuery string var gotVector []float32 m, err := rag.New(rag.Config{ Backend: &mockBackend{ searchFunc: func(_ context.Context, query string, queryVector []float32, _ int) ([]rag.SearchResult, error) { gotQuery = query gotVector = queryVector return []rag.SearchResult{{ID: "1", Content: "matched lexically"}}, nil }, }, Embedder: &embeddings.MockEmbedder{EmbedFunc: func(ctx context.Context, text string) ([]float32, error) { return nil, fmt.Errorf("embed error") }}, }) if err != nil { t.Fatalf("unexpected error: %v", err) } results, err := m.Search(context.Background(), "test query", 5) if err != nil { t.Fatalf("expected Search to fall back to the backend's lexical search, got error: %v", err) } if len(results) != 1 || results[0].Content != "matched lexically" { t.Fatalf("expected the backend's fallback result to come through, got %v", results) } if gotQuery != "test query" { t.Fatalf("expected the raw query text to reach the backend, got %q", gotQuery) } if len(gotVector) != 0 { t.Fatalf("expected no query vector to be passed through, got %v", gotVector) } } // mockBackend implements chroma.Backend for testing. type mockBackend struct { upsertFunc func(ctx context.Context, id string, vector []float32, content string, metadata map[string]string) error searchFunc func(ctx context.Context, query string, queryVector []float32, topK int) ([]rag.SearchResult, error) forgetAllFunc func(ctx context.Context) error } func (m *mockBackend) Upsert(ctx context.Context, id string, vector []float32, content string, metadata map[string]string) error { if m.upsertFunc != nil { return m.upsertFunc(ctx, id, vector, content, metadata) } return nil } func (m *mockBackend) Search(ctx context.Context, query string, queryVector []float32, topK int) ([]rag.SearchResult, error) { if m.searchFunc != nil { return m.searchFunc(ctx, query, queryVector, topK) } return nil, nil } func (m *mockBackend) Forget(ctx context.Context, id string) error { return nil } func (m *mockBackend) ForgetAll(ctx context.Context) error { if m.forgetAllFunc != nil { return m.forgetAllFunc(ctx) } return nil } func TestMemory_Add_MultipleFragments(t *testing.T) { m, err := rag.New(rag.Config{ Backend: &mockBackend{}, Embedder: &embeddings.MockEmbedder{}, }) if err != nil { t.Fatalf("unexpected error: %v", err) } for i := 0; i < 10; i++ { err = m.Add(context.Background(), rag.Fragment{ Content: fmt.Sprintf("test content %d", i), }) if err != nil { t.Fatalf("unexpected error on iteration %d: %v", i, err) } } } func TestMemory_Search_EmptyQuery(t *testing.T) { m, err := rag.New(rag.Config{ Backend: &mockBackend{ searchFunc: func(ctx context.Context, query string, queryVector []float32, topK int) ([]rag.SearchResult, error) { return []rag.SearchResult{}, nil }, }, Embedder: &embeddings.MockEmbedder{}, }) if err != nil { t.Fatalf("unexpected error: %v", err) } results, err := m.Search(context.Background(), "", 5) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(results) != 0 { t.Fatalf("expected 0 results, got %d", len(results)) } } func TestMemory_ForgetAll_Error(t *testing.T) { m, err := rag.New(rag.Config{ Backend: &mockBackend{ forgetAllFunc: func(ctx context.Context) error { return fmt.Errorf("forget all error") }, }, Embedder: &embeddings.MockEmbedder{}, }) if err != nil { t.Fatalf("unexpected error: %v", err) } err = m.ForgetAll(context.Background()) if err == nil { t.Fatal("expected error for forget all failure") } }