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) + } +}