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.
This commit is contained in:
ZanzyTHEbar 2026-02-18 13:07:56 +00:00
parent c8ef234401
commit 159484ca79
4 changed files with 471 additions and 0 deletions

View file

@ -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
}

View file

@ -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'")
}
}

View file

@ -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())
}
}

View file

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