Formatting only (struct field alignment, import ordering) across the files that didn't comply — no semantic changes. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
285 lines
7.4 KiB
Go
285 lines
7.4 KiB
Go
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")
|
|
}
|
|
}
|