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

230 lines
5.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_EmbeddingError(t *testing.T) {
m, err := rag.New(rag.Config{
Backend: &mockBackend{},
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.Fatal("expected error for embedding failure")
}
}
func TestMemory_Search(t *testing.T) {
m, err := rag.New(rag.Config{
Backend: &mockBackend{
searchFunc: func(ctx context.Context, 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_EmbeddingError(t *testing.T) {
m, err := rag.New(rag.Config{
Backend: &mockBackend{},
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.Search(context.Background(), "test query", 5)
if err == nil {
t.Fatal("expected error for embedding failure")
}
}
// mockBackend implements chroma.Backend for testing.
type mockBackend struct {
upsertFunc func(ctx context.Context, id string, vector []float32, metadata map[string]string) error
searchFunc func(ctx context.Context, queryVector []float32, topK int) ([]rag.SearchResult, error)
forgetAllFunc func(ctx context.Context) error
}
func (m *mockBackend) Upsert(ctx context.Context, id string, vector []float32, metadata map[string]string) error {
if m.upsertFunc != nil {
return m.upsertFunc(ctx, id, vector, metadata)
}
return nil
}
func (m *mockBackend) Search(ctx context.Context, queryVector []float32, topK int) ([]rag.SearchResult, error) {
if m.searchFunc != nil {
return m.searchFunc(ctx, 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, 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")
}
}