From 159484ca793a2f99f782b59d2817c8e39f7835e2 Mon Sep 17 00:00:00 2001 From: ZanzyTHEbar Date: Wed, 18 Feb 2026 13:07:56 +0000 Subject: [PATCH] feat(memory): add embedder factory and embedding provider tests NewEmbedder factory selects Ollama or OpenAI based on config provider string. Tests cover both providers' HTTP interaction (mocked), error paths, and factory dispatch logic. --- pkg/memory/store/embed_factory.go | 69 +++++++++++ pkg/memory/store/embed_factory_test.go | 138 ++++++++++++++++++++++ pkg/memory/store/embed_ollama_test.go | 110 ++++++++++++++++++ pkg/memory/store/embed_openai_test.go | 154 +++++++++++++++++++++++++ 4 files changed, 471 insertions(+) create mode 100644 pkg/memory/store/embed_factory.go create mode 100644 pkg/memory/store/embed_factory_test.go create mode 100644 pkg/memory/store/embed_ollama_test.go create mode 100644 pkg/memory/store/embed_openai_test.go diff --git a/pkg/memory/store/embed_factory.go b/pkg/memory/store/embed_factory.go new file mode 100644 index 000000000..b77057fa3 --- /dev/null +++ b/pkg/memory/store/embed_factory.go @@ -0,0 +1,69 @@ +package store + +import ( + "fmt" + "strings" + + "github.com/sipeed/picoclaw/pkg/config" + "github.com/sipeed/picoclaw/pkg/logger" + "github.com/sipeed/picoclaw/pkg/memory" +) + +// NewEmbedderFromConfig creates an EmbeddingProvider from config, wrapped +// with the LRU CachedEmbedder. Returns nil if no provider is configured +// (FTS5-only search mode). +// +// The providersCfg parameter provides fallback API keys from the providers section +// when the embedding-specific config doesn't have its own key. +func NewEmbedderFromConfig(embCfg config.EmbeddingConfig, providersCfg config.ProvidersConfig) (memory.EmbeddingProvider, error) { + provider := strings.ToLower(strings.TrimSpace(embCfg.Provider)) + if provider == "" { + logger.InfoCF("memory", "No embedding provider configured, archival search will use FTS5 only", nil) + return nil, nil + } + + var inner memory.EmbeddingProvider + + switch provider { + case "ollama": + base := embCfg.APIBase + if base == "" && providersCfg.Ollama.APIBase != "" { + base = providersCfg.Ollama.APIBase + } + inner = NewOllamaEmbedder(OllamaEmbedderConfig{ + Base: base, + Model: embCfg.Model, + }) + + case "openai": + apiKey := embCfg.APIKey + if apiKey == "" { + apiKey = providersCfg.OpenAI.APIKey + } + base := embCfg.APIBase + if base == "" && providersCfg.OpenAI.APIBase != "" { + base = providersCfg.OpenAI.APIBase + } + if apiKey == "" { + return nil, fmt.Errorf("openai embedding provider requires an API key (set memory.embedding.api_key or providers.openai.api_key)") + } + inner = NewOpenAIEmbedder(OpenAIEmbedderConfig{ + Base: base, + Model: embCfg.Model, + APIKey: apiKey, + }) + + default: + return nil, fmt.Errorf("unknown embedding provider %q (supported: ollama, openai)", provider) + } + + cached := NewCachedEmbedder(inner, DefaultCachedEmbedderConfig()) + + logger.InfoCF("memory", "Embedding provider initialized", + map[string]interface{}{ + "provider": provider, + "model": inner.Model(), + }) + + return cached, nil +} diff --git a/pkg/memory/store/embed_factory_test.go b/pkg/memory/store/embed_factory_test.go new file mode 100644 index 000000000..91bff0178 --- /dev/null +++ b/pkg/memory/store/embed_factory_test.go @@ -0,0 +1,138 @@ +package store + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/sipeed/picoclaw/pkg/config" +) + +func TestNewEmbedderFromConfig_EmptyProvider(t *testing.T) { + emb, err := NewEmbedderFromConfig(config.EmbeddingConfig{}, config.ProvidersConfig{}) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if emb != nil { + t.Error("expected nil embedder for empty provider") + } +} + +func TestNewEmbedderFromConfig_UnknownProvider(t *testing.T) { + _, err := NewEmbedderFromConfig(config.EmbeddingConfig{Provider: "nonexistent"}, config.ProvidersConfig{}) + if err == nil { + t.Error("expected error for unknown provider") + } +} + +func TestNewEmbedderFromConfig_OpenAI_NoKey(t *testing.T) { + _, err := NewEmbedderFromConfig( + config.EmbeddingConfig{Provider: "openai"}, + config.ProvidersConfig{}, + ) + if err == nil { + t.Error("expected error when OpenAI has no API key") + } +} + +func TestNewEmbedderFromConfig_OpenAI_FallbackKey(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + json.NewEncoder(w).Encode(map[string]interface{}{ + "data": []map[string]interface{}{ + {"index": 0, "embedding": []float32{0.1, 0.2, 0.3}}, + }, + }) + })) + defer srv.Close() + + emb, err := NewEmbedderFromConfig( + config.EmbeddingConfig{ + Provider: "openai", + APIBase: srv.URL, + }, + config.ProvidersConfig{ + OpenAI: config.OpenAIProviderConfig{ + ProviderConfig: config.ProviderConfig{APIKey: "test-key"}, + }, + }, + ) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if emb == nil { + t.Fatal("expected non-nil embedder") + } + if emb.Model() != defaultOpenAIModel { + t.Errorf("expected default model %q, got %q", defaultOpenAIModel, emb.Model()) + } +} + +func TestNewEmbedderFromConfig_Ollama(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + json.NewEncoder(w).Encode(map[string]interface{}{ + "embeddings": [][]float32{{0.1, 0.2, 0.3}}, + }) + })) + defer srv.Close() + + emb, err := NewEmbedderFromConfig( + config.EmbeddingConfig{ + Provider: "ollama", + APIBase: srv.URL, + Model: "test-embed", + }, + config.ProvidersConfig{}, + ) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if emb == nil { + t.Fatal("expected non-nil embedder") + } + if emb.Model() != "test-embed" { + t.Errorf("expected 'test-embed', got %q", emb.Model()) + } +} + +func TestNewEmbedderFromConfig_Ollama_FallbackBase(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + json.NewEncoder(w).Encode(map[string]interface{}{ + "embeddings": [][]float32{{0.5}}, + }) + })) + defer srv.Close() + + emb, err := NewEmbedderFromConfig( + config.EmbeddingConfig{Provider: "ollama"}, + config.ProvidersConfig{ + Ollama: config.ProviderConfig{APIBase: srv.URL}, + }, + ) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if emb == nil { + t.Fatal("expected non-nil embedder") + } +} + +func TestNewEmbedderFromConfig_CaseInsensitive(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + json.NewEncoder(w).Encode(map[string]interface{}{ + "embeddings": [][]float32{{0.1}}, + }) + })) + defer srv.Close() + + emb, err := NewEmbedderFromConfig( + config.EmbeddingConfig{Provider: " Ollama ", APIBase: srv.URL}, + config.ProvidersConfig{}, + ) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if emb == nil { + t.Fatal("expected non-nil embedder for case-insensitive 'Ollama'") + } +} diff --git a/pkg/memory/store/embed_ollama_test.go b/pkg/memory/store/embed_ollama_test.go new file mode 100644 index 000000000..ff9f8267e --- /dev/null +++ b/pkg/memory/store/embed_ollama_test.go @@ -0,0 +1,110 @@ +package store + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" +) + +func TestOllamaEmbedder_Embed(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/api/embed" { + t.Errorf("expected /api/embed, got %s", r.URL.Path) + } + if r.Method != http.MethodPost { + t.Errorf("expected POST, got %s", r.Method) + } + + var req ollamaEmbedRequest + json.NewDecoder(r.Body).Decode(&req) + if req.Model != "test-model" { + t.Errorf("expected model 'test-model', got %q", req.Model) + } + + json.NewEncoder(w).Encode(ollamaEmbedResponse{ + Embeddings: [][]float32{{0.1, 0.2, 0.3, 0.4}}, + }) + })) + defer srv.Close() + + e := NewOllamaEmbedder(OllamaEmbedderConfig{ + Base: srv.URL, + Model: "test-model", + }) + + vec, err := e.Embed(context.Background(), "hello world") + if err != nil { + t.Fatalf("Embed: %v", err) + } + if len(vec) != 4 { + t.Fatalf("expected 4 dims, got %d", len(vec)) + } + if vec[0] != 0.1 { + t.Errorf("expected 0.1, got %f", vec[0]) + } + + // Dimensions should auto-detect + if e.Dimensions() != 4 { + t.Errorf("expected 4 dims, got %d", e.Dimensions()) + } +} + +func TestOllamaEmbedder_EmbedBatch(t *testing.T) { + callCount := 0 + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + callCount++ + json.NewEncoder(w).Encode(ollamaEmbedResponse{ + Embeddings: [][]float32{{float32(callCount) * 0.1}}, + }) + })) + defer srv.Close() + + e := NewOllamaEmbedder(OllamaEmbedderConfig{Base: srv.URL}) + + vecs, err := e.EmbedBatch(context.Background(), []string{"a", "b", "c"}) + if err != nil { + t.Fatalf("EmbedBatch: %v", err) + } + if len(vecs) != 3 { + t.Fatalf("expected 3 vectors, got %d", len(vecs)) + } + if callCount != 3 { + t.Errorf("expected 3 API calls (sequential), got %d", callCount) + } +} + +func TestOllamaEmbedder_ServerError(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusInternalServerError) + w.Write([]byte("model not found")) + })) + defer srv.Close() + + e := NewOllamaEmbedder(OllamaEmbedderConfig{Base: srv.URL}) + _, err := e.Embed(context.Background(), "test") + if err == nil { + t.Error("expected error for server error response") + } +} + +func TestOllamaEmbedder_EmptyResponse(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + json.NewEncoder(w).Encode(ollamaEmbedResponse{Embeddings: [][]float32{}}) + })) + defer srv.Close() + + e := NewOllamaEmbedder(OllamaEmbedderConfig{Base: srv.URL}) + _, err := e.Embed(context.Background(), "test") + if err == nil { + t.Error("expected error for empty embeddings") + } +} + +func TestOllamaEmbedder_DefaultModel(t *testing.T) { + e := NewOllamaEmbedder(OllamaEmbedderConfig{}) + if e.Model() != defaultOllamaModel { + t.Errorf("expected %q, got %q", defaultOllamaModel, e.Model()) + } +} diff --git a/pkg/memory/store/embed_openai_test.go b/pkg/memory/store/embed_openai_test.go new file mode 100644 index 000000000..0a553841b --- /dev/null +++ b/pkg/memory/store/embed_openai_test.go @@ -0,0 +1,154 @@ +package store + +import ( + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" +) + +func TestOpenAIEmbedder_Embed(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/embeddings" { + t.Errorf("expected /embeddings, got %s", r.URL.Path) + } + + auth := r.Header.Get("Authorization") + if auth != "Bearer test-key" { + t.Errorf("expected 'Bearer test-key', got %q", auth) + } + + var req openAIEmbedRequest + json.NewDecoder(r.Body).Decode(&req) + if req.Model != "test-embed" { + t.Errorf("expected model 'test-embed', got %q", req.Model) + } + + json.NewEncoder(w).Encode(openAIEmbedResponse{ + Data: []openAIEmbedData{ + {Index: 0, Embedding: []float32{0.5, 0.6, 0.7}}, + }, + }) + })) + defer srv.Close() + + e := NewOpenAIEmbedder(OpenAIEmbedderConfig{ + Base: srv.URL, + Model: "test-embed", + APIKey: "test-key", + }) + + vec, err := e.Embed(context.Background(), "hello") + if err != nil { + t.Fatalf("Embed: %v", err) + } + if len(vec) != 3 { + t.Fatalf("expected 3 dims, got %d", len(vec)) + } + if vec[0] != 0.5 { + t.Errorf("expected 0.5, got %f", vec[0]) + } + if e.Dimensions() != 3 { + t.Errorf("expected 3 dims auto-detected, got %d", e.Dimensions()) + } +} + +func TestOpenAIEmbedder_EmbedBatch(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req openAIEmbedRequest + json.NewDecoder(r.Body).Decode(&req) + + // Batch request should send array + texts, ok := req.Input.([]interface{}) + if !ok { + t.Fatalf("expected array input for batch, got %T", req.Input) + } + + data := make([]openAIEmbedData, len(texts)) + for i := range texts { + data[i] = openAIEmbedData{ + Index: i, + Embedding: []float32{float32(i) * 0.1}, + } + } + + json.NewEncoder(w).Encode(openAIEmbedResponse{Data: data}) + })) + defer srv.Close() + + e := NewOpenAIEmbedder(OpenAIEmbedderConfig{ + Base: srv.URL, + APIKey: "k", + }) + + vecs, err := e.EmbedBatch(context.Background(), []string{"a", "b", "c"}) + if err != nil { + t.Fatalf("EmbedBatch: %v", err) + } + if len(vecs) != 3 { + t.Fatalf("expected 3 vectors, got %d", len(vecs)) + } +} + +func TestOpenAIEmbedder_SingleTextNotArray(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + var req openAIEmbedRequest + json.NewDecoder(r.Body).Decode(&req) + + // Single text should be sent as string, not array + if _, ok := req.Input.(string); !ok { + t.Errorf("expected string input for single text, got %T", req.Input) + } + + json.NewEncoder(w).Encode(openAIEmbedResponse{ + Data: []openAIEmbedData{{Index: 0, Embedding: []float32{1.0}}}, + }) + })) + defer srv.Close() + + e := NewOpenAIEmbedder(OpenAIEmbedderConfig{Base: srv.URL, APIKey: "k"}) + _, err := e.Embed(context.Background(), "single") + if err != nil { + t.Fatalf("Embed: %v", err) + } +} + +func TestOpenAIEmbedder_ServerError(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusUnauthorized) + w.Write([]byte(`{"error":{"message":"invalid api key"}}`)) + })) + defer srv.Close() + + e := NewOpenAIEmbedder(OpenAIEmbedderConfig{Base: srv.URL, APIKey: "bad"}) + _, err := e.Embed(context.Background(), "test") + if err == nil { + t.Error("expected error for 401 response") + } +} + +func TestOpenAIEmbedder_DefaultModel(t *testing.T) { + e := NewOpenAIEmbedder(OpenAIEmbedderConfig{APIKey: "k"}) + if e.Model() != defaultOpenAIModel { + t.Errorf("expected %q, got %q", defaultOpenAIModel, e.Model()) + } +} + +func TestOpenAIEmbedder_NoAuth(t *testing.T) { + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("Authorization") != "" { + t.Error("expected no auth header when key is empty") + } + json.NewEncoder(w).Encode(openAIEmbedResponse{ + Data: []openAIEmbedData{{Index: 0, Embedding: []float32{1.0}}}, + }) + })) + defer srv.Close() + + e := NewOpenAIEmbedder(OpenAIEmbedderConfig{Base: srv.URL}) + _, err := e.Embed(context.Background(), "test") + if err != nil { + t.Fatalf("Embed: %v", err) + } +}