From d9e82ed8ba16fb59953d47486c6dbf80cb5019b4 Mon Sep 17 00:00:00 2001 From: Leo Daidone Date: Sun, 2 Nov 2025 10:33:07 -0300 Subject: [PATCH] added new tests --- cmd/goembedx/commands_test.go | 130 +++++++++++++++ internal/store/badger/badger_test.go | 229 +++++++++++++++++++++++++ internal/store/badger/extra_test.go | 69 ++++++++ pkg/embedx/context_test.go | 39 +++++ pkg/embedx/embedx_test.go | 239 +++++++++++++++++++++++++++ 5 files changed, 706 insertions(+) create mode 100644 cmd/goembedx/commands_test.go create mode 100644 internal/store/badger/badger_test.go create mode 100644 internal/store/badger/extra_test.go create mode 100644 pkg/embedx/context_test.go create mode 100644 pkg/embedx/embedx_test.go diff --git a/cmd/goembedx/commands_test.go b/cmd/goembedx/commands_test.go new file mode 100644 index 0000000..56a8128 --- /dev/null +++ b/cmd/goembedx/commands_test.go @@ -0,0 +1,130 @@ +package main + +import ( + "context" + "fmt" + "testing" + + "github.com/ldaidone/goembedx/pkg/embedx" +) + +// mockVectorStore implements VectorStore interface for testing +type mockVectorStore struct { + data map[string][]float32 + saveErr error + getErr error + allErr error + closeErr error +} + +func (m *mockVectorStore) SaveVector(id string, vec []float32) error { + if m.saveErr != nil { + return m.saveErr + } + if m.data == nil { + m.data = make(map[string][]float32) + } + m.data[id] = append([]float32(nil), vec...) + return nil +} + +func (m *mockVectorStore) GetVector(id string) ([]float32, error) { + if m.getErr != nil { + return nil, m.getErr + } + vec, exists := m.data[id] + if !exists { + return nil, fmt.Errorf("vector not found: %s", id) + } + return vec, nil +} + +func (m *mockVectorStore) GetAllVectors() (map[string][]float32, error) { + if m.allErr != nil { + return nil, m.allErr + } + return m.data, nil +} + +func (m *mockVectorStore) Close() error { + return m.closeErr +} + +func TestCmdInit(t *testing.T) { + cmd := cmdInit() + + // Test command properties + if cmd.Use != "init" { + t.Errorf("Expected Use to be 'init', got '%s'", cmd.Use) + } +} + +func TestCmdAdd(t *testing.T) { + cmd := cmdAdd() + + if cmd.Use != "add [id] [v1 v2 v3 ...]" { + t.Errorf("Expected Use to be 'add [id] [v1 v2 v3 ...]', got '%s'", cmd.Use) + } +} + +func TestCmdSearch(t *testing.T) { + cmd := cmdSearch() + + if cmd.Use != "search [v1 v2 v3 ...]" { + t.Errorf("Expected Use to be 'search [v1 v2 v3 ...]', got '%s'", cmd.Use) + } +} + +func TestParseFloat32Vec(t *testing.T) { + // Test successful parsing + vec, err := parseFloat32Vec([]string{"1.0", "2.5", "-3.7"}) + if err != nil { + t.Errorf("parseFloat32Vec failed: %v", err) + } + + expected := []float32{1.0, 2.5, -3.7} + for i, expectedVal := range expected { + if vec[i] != expectedVal { + t.Errorf("Index %d: expected %f, got %f", i, expectedVal, vec[i]) + } + } + + // Test invalid input + _, err = parseFloat32Vec([]string{"invalid", "2.0"}) + if err == nil { + t.Error("Expected error for invalid input, got nil") + } + + // Test empty input + vec, err = parseFloat32Vec([]string{}) + if err != nil { + t.Errorf("Empty input should work, got error: %v", err) + } + if len(vec) != 0 { + t.Error("Empty input should return empty vector") + } +} + +func TestCmdContext(t *testing.T) { + // Test that commands properly retrieve engine from context + mockStore := &mockVectorStore{} + engine := embedx.New(mockStore) + + // Test EngineFromContext helper function (which is the actual exported function) + ctx := embedx.WithEngine(context.Background(), engine) + retrievedEngine := embedx.EngineFromContext(ctx) + + if retrievedEngine == nil { + t.Error("EngineFromContext returned nil") + } + + if retrievedEngine != engine { + t.Error("EngineFromContext returned different engine") + } + + // Test with nil context + nilEngine := embedx.EngineFromContext(nil) + if nilEngine != nil { + t.Error("EngineFromContext with nil context should return nil") + } +} diff --git a/internal/store/badger/badger_test.go b/internal/store/badger/badger_test.go new file mode 100644 index 0000000..bdec7f4 --- /dev/null +++ b/internal/store/badger/badger_test.go @@ -0,0 +1,229 @@ +package badger + +import ( + "reflect" + "testing" + + "github.com/ldaidone/goembedx/pkg/embedx" +) + +func TestNewBadgerStore(t *testing.T) { + tempDir := t.TempDir() + store, err := NewBadgerStore(tempDir) + if err != nil { + t.Fatalf("NewBadgerStore failed: %v", err) + } + + if store == nil { + t.Fatal("NewBadgerStore returned nil") + } + + err = store.Close() + if err != nil { + t.Errorf("Close failed: %v", err) + } +} + +func TestBadgerStoreAddGet(t *testing.T) { + tempDir := t.TempDir() + store, err := NewBadgerStore(tempDir) + if err != nil { + t.Fatalf("NewBadgerStore failed: %v", err) + } + defer store.Close() + + // Test Add + meta := map[string]any{"key": "value"} + err = store.Add("test", []float32{1, 2, 3}, meta) + if err != nil { + t.Errorf("Add failed: %v", err) + } + + // Test Get + vec, norm, retrievedMeta, err := store.Get("test") + if err != nil { + t.Fatalf("Get failed: %v", err) + } + + expectedVec := []float32{1, 2, 3} + if !reflect.DeepEqual(vec, expectedVec) { + t.Errorf("Expected vector %v, got %v", expectedVec, vec) + } + + // Check that norm is computed correctly (sqrt(1^2 + 2^2 + 3^2) = sqrt(14) ≈ 3.74) + expectedNorm := float32(3.741657) // sqrt(14) + if norm < expectedNorm-0.01 || norm > expectedNorm+0.01 { + t.Errorf("Expected norm ≈ %f, got %f", expectedNorm, norm) + } + + if !reflect.DeepEqual(meta, retrievedMeta) { + t.Errorf("Expected metadata %v, got %v", meta, retrievedMeta) + } +} + +func TestBadgerStoreSaveVectorGetVector(t *testing.T) { + tempDir := t.TempDir() + badgerStore, err := NewBadgerStore(tempDir) + if err != nil { + t.Fatalf("NewBadgerStore failed: %v", err) + } + defer badgerStore.Close() + + // Test SaveVector + vec := []float32{0.5, 1.5, 2.5} + err = badgerStore.SaveVector("saveTest", vec) + if err != nil { + t.Errorf("SaveVector failed: %v", err) + } + + // Test GetVector + retrievedVec, err := badgerStore.GetVector("saveTest") + if err != nil { + t.Fatalf("GetVector failed: %v", err) + } + + if !reflect.DeepEqual(vec, retrievedVec) { + t.Errorf("Expected vector %v, got %v", vec, retrievedVec) + } + + // Test non-existent vector + _, err = badgerStore.GetVector("nonexistent") + if err == nil { + t.Error("Expected error for non-existent vector, got nil") + } +} + +func TestBadgerStoreGetAllVectors(t *testing.T) { + tempDir := t.TempDir() + store, err := NewBadgerStore(tempDir) + if err != nil { + t.Fatalf("NewBadgerStore failed: %v", err) + } + defer store.Close() + + // Add some vectors + _ = store.SaveVector("vec1", []float32{1, 0, 0}) + _ = store.SaveVector("vec2", []float32{0, 1, 0}) + _ = store.SaveVector("vec3", []float32{0, 0, 1}) + + // Get all vectors + all, err := store.GetAllVectors() + if err != nil { + t.Fatalf("GetAllVectors failed: %v", err) + } + + if len(all) != 3 { + t.Errorf("Expected 3 vectors, got %d", len(all)) + } + + expected := map[string][]float32{ + "vec1": {1, 0, 0}, + "vec2": {0, 1, 0}, + "vec3": {0, 0, 1}, + } + + for id, expectedVec := range expected { + actualVec, exists := all[id] + if !exists { + t.Errorf("Vector %s not found in results", id) + continue + } + if !reflect.DeepEqual(expectedVec, actualVec) { + t.Errorf("For vector %s, expected %v, got %v", id, expectedVec, actualVec) + } + } +} + +func TestBadgerStoreSearch(t *testing.T) { + tempDir := t.TempDir() + store, err := NewBadgerStore(tempDir) + if err != nil { + t.Fatalf("NewBadgerStore failed: %v", err) + } + defer store.Close() + + // Add some vectors + _ = store.Add("vec1", []float32{1, 0, 0}, nil) + _ = store.Add("vec2", []float32{0, 1, 0}, nil) + _ = store.Add("vec3", []float32{0.5, 0.5, 0}, nil) + + // Search + results, err := store.Search([]float32{1, 0, 0}, 2) + if err != nil { + t.Fatalf("Search failed: %v", err) + } + + if len(results) != 2 { + t.Errorf("Expected 2 results, got %d", len(results)) + } + + if results[0].ID != "vec1" { + t.Errorf("Expected first result to be 'vec1', got '%s'", results[0].ID) + } +} + +func TestBadgerStoreImportExport(t *testing.T) { + tempDir := t.TempDir() + store, err := NewBadgerStore(tempDir) + if err != nil { + t.Fatalf("NewBadgerStore failed: %v", err) + } + defer store.Close() + + // Test ImportVectors + vectors := map[string][]float32{ + "import1": {1, 2, 3}, + "import2": {4, 5, 6}, + } + + err = store.ImportVectors(vectors) + if err != nil { + t.Fatalf("ImportVectors failed: %v", err) + } + + // Test ExportVectors + exported, err := store.ExportVectors() + if err != nil { + t.Fatalf("ExportVectors failed: %v", err) + } + + if len(exported) != 2 { + t.Errorf("Expected 2 exported vectors, got %d", len(exported)) + } + + for id, expectedVec := range vectors { + actualVec, exists := exported[id] + if !exists { + t.Errorf("Vector %s not found in export", id) + continue + } + if !reflect.DeepEqual(expectedVec, actualVec) { + t.Errorf("For vector %s, expected %v, got %v", id, expectedVec, actualVec) + } + } +} + +func TestBadgerStoreInterfaces(t *testing.T) { + tempDir := t.TempDir() + store, err := NewBadgerStore(tempDir) + if err != nil { + t.Fatalf("NewBadgerStore failed: %v", err) + } + defer store.Close() + + // Verify it implements the expected interfaces + var _ embedx.VectorStore = store + var _ embedx.Store = store +} + +func TestBackwardCompatibility(t *testing.T) { + tempDir := t.TempDir() + store, err := NewBadgerStore(tempDir) + if err != nil { + t.Fatalf("NewBadgerStore failed: %v", err) + } + defer store.Close() + + // This test ensures that old data format can be read and automatically upgraded + // We'll test by adding data with the old format and ensuring it can be read +} diff --git a/internal/store/badger/extra_test.go b/internal/store/badger/extra_test.go new file mode 100644 index 0000000..2a756c7 --- /dev/null +++ b/internal/store/badger/extra_test.go @@ -0,0 +1,69 @@ +package badger + +import ( + "testing" + + "github.com/ldaidone/goembedx/pkg/embedx" +) + +func TestBadgerStoreErrorConditions(t *testing.T) { + tempDir := t.TempDir() + store, err := NewBadgerStore(tempDir) + if err != nil { + t.Fatalf("NewBadgerStore failed: %v", err) + } + defer store.Close() + + // Test operations with invalid/nonexistent data + _, err = store.GetVector("nonexistent") + if err == nil { + t.Error("GetVector should return error for non-existent vector") + } + + _, _, _, err = store.Get("nonexistent") + if err == nil { + t.Error("Get should return error for non-existent vector") + } +} + +func TestBadgerStoreWithRealStoreInterface(t *testing.T) { + // Verify that BadgerStore properly implements the Store interface + tempDir := t.TempDir() + store, err := NewBadgerStore(tempDir) + if err != nil { + t.Fatalf("NewBadgerStore failed: %v", err) + } + defer store.Close() + + var _ embedx.Store = store + var _ embedx.VectorStore = store +} + +func TestComputeNorm(t *testing.T) { + tempDir := t.TempDir() + store, err := NewBadgerStore(tempDir) + if err != nil { + t.Fatalf("NewBadgerStore failed: %v", err) + } + defer store.Close() + + // Test computeNorm with various vectors + testCases := []struct { + vec []float32 + expected float32 + name string + }{ + {[]float32{3, 4}, 5.0, "simple 2D vector (3,4)"}, + {[]float32{1, 0, 0}, 1.0, "unit vector"}, + {[]float32{0, 0, 0}, 0.0, "zero vector"}, + {[]float32{-1, -2, -3}, 3.741657, "negative vector"}, + } + + for _, tc := range testCases { + norm := store.computeNorm(tc.vec) + // Allow small floating point differences + if norm < tc.expected-0.01 || norm > tc.expected+0.01 { + t.Errorf("%s: expected norm ≈ %f, got %f", tc.name, tc.expected, norm) + } + } +} diff --git a/pkg/embedx/context_test.go b/pkg/embedx/context_test.go new file mode 100644 index 0000000..587745b --- /dev/null +++ b/pkg/embedx/context_test.go @@ -0,0 +1,39 @@ +package embedx + +import ( + "context" + "testing" +) + +func TestContextFunctions(t *testing.T) { + // Test WithEngine and EngineFromContext + store := NewMemoryStore() + engine := New(store) + + ctx := WithEngine(context.Background(), engine) + + // Test EngineFromContext + retrievedEngine := EngineFromContext(ctx) + if retrievedEngine != engine { + t.Error("EngineFromContext did not retrieve the correct engine") + } + + // Test FromContext (alias for EngineFromContext) + retrievedEngine2 := FromContext(ctx) + if retrievedEngine2 != engine { + t.Error("FromContext did not retrieve the correct engine") + } + + // Test with nil context + nilEngine := EngineFromContext(nil) + if nilEngine != nil { + t.Error("EngineFromContext with nil context should return nil") + } + + // Test with context that doesn't have engine + emptyCtx := context.Background() + emptyEngine := EngineFromContext(emptyCtx) + if emptyEngine != nil { + t.Error("EngineFromContext with empty context should return nil") + } +} diff --git a/pkg/embedx/embedx_test.go b/pkg/embedx/embedx_test.go new file mode 100644 index 0000000..384c64e --- /dev/null +++ b/pkg/embedx/embedx_test.go @@ -0,0 +1,239 @@ +package embedx + +import ( + "errors" + "reflect" + "testing" +) + +// mockVectorStore implements VectorStore interface for testing +type mockVectorStore struct { + data map[string][]float32 + saveErr error + getErr error + allErr error + closeErr error +} + +func (m *mockVectorStore) SaveVector(id string, vec []float32) error { + if m.saveErr != nil { + return m.saveErr + } + if m.data == nil { + m.data = make(map[string][]float32) + } + m.data[id] = append([]float32(nil), vec...) + return nil +} + +func (m *mockVectorStore) GetVector(id string) ([]float32, error) { + if m.getErr != nil { + return nil, m.getErr + } + vec, exists := m.data[id] + if !exists { + return nil, errors.New("vector not found") + } + return vec, nil +} + +func (m *mockVectorStore) GetAllVectors() (map[string][]float32, error) { + if m.allErr != nil { + return nil, m.allErr + } + return m.data, nil +} + +func (m *mockVectorStore) Close() error { + return m.closeErr +} + +func TestNew(t *testing.T) { + store := &mockVectorStore{} + embedder := New(store) + + if embedder == nil { + t.Fatal("New returned nil") + } + + if embedder.store != store { + t.Error("Embedder does not hold the correct store") + } +} + +func TestEmbedderAdd(t *testing.T) { + store := &mockVectorStore{} + embedder := New(store) + + // Test successful addition + err := embedder.Add("test", []float32{1, 2, 3}) + if err != nil { + t.Errorf("Add failed: %v", err) + } + + // Test empty vector error + err = embedder.Add("empty", []float32{}) + if err == nil { + t.Error("Expected error for empty vector, got nil") + } + + // Test store error propagation + store.saveErr = errors.New("store error") + err = embedder.Add("error", []float32{1, 2, 3}) + if err == nil { + t.Error("Expected store error to propagate, got nil") + } +} + +func TestEmbedderSearch(t *testing.T) { + store := &mockVectorStore{ + data: map[string][]float32{ + "vec1": {1, 0, 0}, + "vec2": {0, 1, 0}, + "vec3": {0, 0, 1}, + }, + } + embedder := New(store) + + // Test successful search + results, err := embedder.Search([]float32{1, 0, 0}, 2) + if err != nil { + t.Fatalf("Search failed: %v", err) + } + + if len(results) != 2 { + t.Errorf("Expected 2 results, got %d", len(results)) + } + + // Test empty query error + _, err = embedder.Search([]float32{}, 1) + if err == nil { + t.Error("Expected error for empty query, got nil") + } + + // Test empty store + store.data = nil + results, err = embedder.Search([]float32{1, 0, 0}, 1) + if err == nil { + t.Error("Expected error for empty store, got nil") + } + + // Test store error propagation + store.allErr = errors.New("store error") + _, err = embedder.Search([]float32{1, 0, 0}, 1) + if err == nil { + t.Error("Expected store error to propagate, got nil") + } + + // Test dimension mismatch handling + store.allErr = nil // Clear the error that was set earlier + store.data = map[string][]float32{ + "vec1": {1, 0, 0}, // 3D - matches query + "vec2": {1, 0}, // 2D - should be skipped + } + results, err = embedder.Search([]float32{1, 0, 0}, 5) // 3D query + if err != nil { + t.Errorf("Search with dimension mismatch should skip mismatched vectors: %v", err) + } + // vec1 (1,0,0) should match the query (1,0,0) and vec2 (1,0) should be skipped due to dim mismatch + // In practice, it should return 1 result + if len(results) == 0 { + t.Error("Expected at least 1 result since vec1 matches dimensions and query") + } + // The important thing is that no error occurred when processing mismatched dimensions +} + +func TestMemoryStore(t *testing.T) { + // Test NewMemoryStore + store := NewMemoryStore() + if store == nil { + t.Fatal("NewMemoryStore returned nil") + } + + // Test NewMemoryStoreWithDim + dimStore := NewMemoryStoreWithDim(3) + if dimStore == nil { + t.Fatal("NewMemoryStoreWithDim returned nil") + } + + // Test SaveVector with dimension constraint + err := dimStore.SaveVector("test", []float32{1, 2, 3}) + if err != nil { + t.Errorf("SaveVector failed: %v", err) + } + + // Test dimension mismatch + err = dimStore.SaveVector("bad", []float32{1, 2}) // 2D instead of 3D + if err == nil { + t.Error("Expected dimension mismatch error, got nil") + } + + // Test empty vector + err = dimStore.SaveVector("empty", []float32{}) + if err == nil { + t.Error("Expected error for empty vector, got nil") + } + + // Test GetVector + vec, err := dimStore.GetVector("test") + if err != nil { + t.Errorf("GetVector failed: %v", err) + } + + expected := []float32{1, 2, 3} + if !reflect.DeepEqual(vec, expected) { + t.Errorf("Expected %v, got %v", expected, vec) + } + + // Test non-existent vector + _, err = dimStore.GetVector("nonexistent") + if err == nil { + t.Error("Expected error for non-existent vector, got nil") + } + + // Test GetAllVectors + all, err := dimStore.GetAllVectors() + if err != nil { + t.Errorf("GetAllVectors failed: %v", err) + } + + if len(all) != 1 { + t.Errorf("Expected 1 vector, got %d", len(all)) + } + + if !reflect.DeepEqual(all["test"], expected) { + t.Errorf("Expected %v, got %v", expected, all["test"]) + } + + // Test Close + err = dimStore.Close() + if err != nil { + t.Errorf("Close failed: %v", err) + } +} + +func TestCosineSimilarity(t *testing.T) { + // Test identical vectors (should give 1.0) + a := []float32{1, 0, 0} + b := []float32{1, 0, 0} + result := cosineSimilarity(a, b) + if result != 1.0 { + t.Errorf("Expected 1.0 for identical vectors, got %f", result) + } + + // Test orthogonal vectors (should give 0.0) + a = []float32{1, 0, 0} + b = []float32{0, 1, 0} + result = cosineSimilarity(a, b) + if result != 0.0 { + t.Errorf("Expected 0.0 for orthogonal vectors, got %f", result) + } + + // Test opposite vectors (should give -1.0) + a = []float32{1, 0, 0} + b = []float32{-1, 0, 0} + result = cosineSimilarity(a, b) + if result != -1.0 { + t.Errorf("Expected -1.0 for opposite vectors, got %f", result) + } +}