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:
parent
c8ef234401
commit
159484ca79
4 changed files with 471 additions and 0 deletions
69
pkg/memory/store/embed_factory.go
Normal file
69
pkg/memory/store/embed_factory.go
Normal 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
|
||||
}
|
||||
138
pkg/memory/store/embed_factory_test.go
Normal file
138
pkg/memory/store/embed_factory_test.go
Normal 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'")
|
||||
}
|
||||
}
|
||||
110
pkg/memory/store/embed_ollama_test.go
Normal file
110
pkg/memory/store/embed_ollama_test.go
Normal 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())
|
||||
}
|
||||
}
|
||||
154
pkg/memory/store/embed_openai_test.go
Normal file
154
pkg/memory/store/embed_openai_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue