feat(memory): add Ollama and OpenAI embedding providers
Concrete Embedder implementations for local Ollama and remote OpenAI APIs. Both implement the memory.Embedder interface with proper error handling, HTTP client configuration, and response parsing. Enables pluggable embedding backends for archival memory vector search.
This commit is contained in:
parent
cf7952f2c5
commit
b9d31a48d2
2 changed files with 286 additions and 0 deletions
134
pkg/memory/store/embed_ollama.go
Normal file
134
pkg/memory/store/embed_ollama.go
Normal file
|
|
@ -0,0 +1,134 @@
|
||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/memory"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
defaultOllamaBase = "http://localhost:11434"
|
||||||
|
defaultOllamaModel = "nomic-embed-text"
|
||||||
|
)
|
||||||
|
|
||||||
|
// OllamaEmbedder implements memory.EmbeddingProvider using Ollama's /api/embed endpoint.
|
||||||
|
type OllamaEmbedder struct {
|
||||||
|
base string
|
||||||
|
model string
|
||||||
|
dims int
|
||||||
|
client *http.Client
|
||||||
|
}
|
||||||
|
|
||||||
|
// OllamaEmbedderConfig configures the Ollama embedding provider.
|
||||||
|
type OllamaEmbedderConfig struct {
|
||||||
|
Base string // API base URL. Empty uses default (http://localhost:11434).
|
||||||
|
Model string // Model name. Empty uses "nomic-embed-text".
|
||||||
|
Dims int // Expected dimensions. Zero auto-detects on first call.
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewOllamaEmbedder creates an Ollama embedding provider.
|
||||||
|
func NewOllamaEmbedder(cfg OllamaEmbedderConfig) *OllamaEmbedder {
|
||||||
|
base := cfg.Base
|
||||||
|
if base == "" {
|
||||||
|
base = defaultOllamaBase
|
||||||
|
}
|
||||||
|
model := cfg.Model
|
||||||
|
if model == "" {
|
||||||
|
model = defaultOllamaModel
|
||||||
|
}
|
||||||
|
return &OllamaEmbedder{
|
||||||
|
base: base,
|
||||||
|
model: model,
|
||||||
|
dims: cfg.Dims,
|
||||||
|
client: &http.Client{
|
||||||
|
Timeout: 60 * time.Second,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type ollamaEmbedRequest struct {
|
||||||
|
Model string `json:"model"`
|
||||||
|
Input string `json:"input"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type ollamaEmbedResponse struct {
|
||||||
|
Embeddings [][]float32 `json:"embeddings"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o *OllamaEmbedder) Embed(ctx context.Context, text string) (memory.Embedding, error) {
|
||||||
|
vecs, err := o.embedBatch(ctx, []string{text})
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if len(vecs) == 0 {
|
||||||
|
return nil, fmt.Errorf("ollama returned no embeddings")
|
||||||
|
}
|
||||||
|
return vecs[0], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o *OllamaEmbedder) EmbedBatch(ctx context.Context, texts []string) ([]memory.Embedding, error) {
|
||||||
|
return o.embedBatch(ctx, texts)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o *OllamaEmbedder) embedBatch(ctx context.Context, texts []string) ([]memory.Embedding, error) {
|
||||||
|
results := make([]memory.Embedding, len(texts))
|
||||||
|
for i, text := range texts {
|
||||||
|
vec, err := o.embedSingle(ctx, text)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("embed text %d: %w", i, err)
|
||||||
|
}
|
||||||
|
results[i] = vec
|
||||||
|
|
||||||
|
if o.dims == 0 && len(vec) > 0 {
|
||||||
|
o.dims = len(vec)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return results, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o *OllamaEmbedder) embedSingle(ctx context.Context, text string) (memory.Embedding, error) {
|
||||||
|
body, err := json.Marshal(ollamaEmbedRequest{
|
||||||
|
Model: o.model,
|
||||||
|
Input: text,
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("marshal request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, o.base+"/api/embed", bytes.NewReader(body))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("create request: %w", err)
|
||||||
|
}
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
|
||||||
|
resp, err := o.client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("ollama request: %w", err)
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
respBody, _ := io.ReadAll(resp.Body)
|
||||||
|
return nil, fmt.Errorf("ollama returned %d: %s", resp.StatusCode, string(respBody))
|
||||||
|
}
|
||||||
|
|
||||||
|
var result ollamaEmbedResponse
|
||||||
|
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
|
||||||
|
return nil, fmt.Errorf("decode response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(result.Embeddings) == 0 {
|
||||||
|
return nil, fmt.Errorf("ollama returned empty embeddings array")
|
||||||
|
}
|
||||||
|
|
||||||
|
return memory.Embedding(result.Embeddings[0]), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o *OllamaEmbedder) Dimensions() int { return o.dims }
|
||||||
|
func (o *OllamaEmbedder) Model() string { return o.model }
|
||||||
152
pkg/memory/store/embed_openai.go
Normal file
152
pkg/memory/store/embed_openai.go
Normal file
|
|
@ -0,0 +1,152 @@
|
||||||
|
package store
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/memory"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
defaultOpenAIBase = "https://api.openai.com/v1"
|
||||||
|
defaultOpenAIModel = "text-embedding-3-small"
|
||||||
|
)
|
||||||
|
|
||||||
|
// OpenAIEmbedder implements memory.EmbeddingProvider using the OpenAI /v1/embeddings API.
|
||||||
|
// Compatible with OpenAI, Azure OpenAI, OpenRouter, and any OpenAI-compatible endpoint.
|
||||||
|
type OpenAIEmbedder struct {
|
||||||
|
base string
|
||||||
|
model string
|
||||||
|
apiKey string
|
||||||
|
dims int
|
||||||
|
client *http.Client
|
||||||
|
}
|
||||||
|
|
||||||
|
// OpenAIEmbedderConfig configures the OpenAI embedding provider.
|
||||||
|
type OpenAIEmbedderConfig struct {
|
||||||
|
Base string // API base URL. Empty uses "https://api.openai.com/v1".
|
||||||
|
Model string // Model name. Empty uses "text-embedding-3-small".
|
||||||
|
APIKey string // Required API key.
|
||||||
|
Dims int // Expected dimensions. Zero auto-detects on first call.
|
||||||
|
}
|
||||||
|
|
||||||
|
// NewOpenAIEmbedder creates an OpenAI-compatible embedding provider.
|
||||||
|
func NewOpenAIEmbedder(cfg OpenAIEmbedderConfig) *OpenAIEmbedder {
|
||||||
|
base := cfg.Base
|
||||||
|
if base == "" {
|
||||||
|
base = defaultOpenAIBase
|
||||||
|
}
|
||||||
|
model := cfg.Model
|
||||||
|
if model == "" {
|
||||||
|
model = defaultOpenAIModel
|
||||||
|
}
|
||||||
|
return &OpenAIEmbedder{
|
||||||
|
base: base,
|
||||||
|
model: model,
|
||||||
|
apiKey: cfg.APIKey,
|
||||||
|
dims: cfg.Dims,
|
||||||
|
client: &http.Client{
|
||||||
|
Timeout: 60 * time.Second,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type openAIEmbedRequest struct {
|
||||||
|
Input interface{} `json:"input"` // string or []string
|
||||||
|
Model string `json:"model"`
|
||||||
|
EncodingFormat string `json:"encoding_format,omitempty"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type openAIEmbedResponse struct {
|
||||||
|
Data []openAIEmbedData `json:"data"`
|
||||||
|
Usage openAIEmbedUsage `json:"usage"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type openAIEmbedData struct {
|
||||||
|
Index int `json:"index"`
|
||||||
|
Embedding []float32 `json:"embedding"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type openAIEmbedUsage struct {
|
||||||
|
PromptTokens int `json:"prompt_tokens"`
|
||||||
|
TotalTokens int `json:"total_tokens"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o *OpenAIEmbedder) Embed(ctx context.Context, text string) (memory.Embedding, error) {
|
||||||
|
vecs, err := o.call(ctx, []string{text})
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if len(vecs) == 0 {
|
||||||
|
return nil, fmt.Errorf("openai returned no embeddings")
|
||||||
|
}
|
||||||
|
return vecs[0], nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o *OpenAIEmbedder) EmbedBatch(ctx context.Context, texts []string) ([]memory.Embedding, error) {
|
||||||
|
return o.call(ctx, texts)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o *OpenAIEmbedder) call(ctx context.Context, texts []string) ([]memory.Embedding, error) {
|
||||||
|
var input interface{}
|
||||||
|
if len(texts) == 1 {
|
||||||
|
input = texts[0]
|
||||||
|
} else {
|
||||||
|
input = texts
|
||||||
|
}
|
||||||
|
|
||||||
|
body, err := json.Marshal(openAIEmbedRequest{
|
||||||
|
Input: input,
|
||||||
|
Model: o.model,
|
||||||
|
EncodingFormat: "float",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("marshal request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(ctx, http.MethodPost, o.base+"/embeddings", bytes.NewReader(body))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("create request: %w", err)
|
||||||
|
}
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
if o.apiKey != "" {
|
||||||
|
req.Header.Set("Authorization", "Bearer "+o.apiKey)
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := o.client.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("openai request: %w", err)
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
respBody, _ := io.ReadAll(resp.Body)
|
||||||
|
return nil, fmt.Errorf("openai returned %d: %s", resp.StatusCode, string(respBody))
|
||||||
|
}
|
||||||
|
|
||||||
|
var result openAIEmbedResponse
|
||||||
|
if err := json.NewDecoder(resp.Body).Decode(&result); err != nil {
|
||||||
|
return nil, fmt.Errorf("decode response: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
embeddings := make([]memory.Embedding, len(result.Data))
|
||||||
|
for _, d := range result.Data {
|
||||||
|
if d.Index < len(embeddings) {
|
||||||
|
embeddings[d.Index] = memory.Embedding(d.Embedding)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if o.dims == 0 && len(embeddings) > 0 && len(embeddings[0]) > 0 {
|
||||||
|
o.dims = len(embeddings[0])
|
||||||
|
}
|
||||||
|
|
||||||
|
return embeddings, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (o *OpenAIEmbedder) Dimensions() int { return o.dims }
|
||||||
|
func (o *OpenAIEmbedder) Model() string { return o.model }
|
||||||
Loading…
Add table
Reference in a new issue