feat: add openai-compatible embeddings support

This commit is contained in:
Badgerbees 2026-04-22 20:29:28 +07:00
parent 68ceb54b36
commit 9f2b031ed9
5 changed files with 269 additions and 0 deletions

View file

@ -70,6 +70,24 @@ func (p *HTTPProvider) ChatStream(
return p.delegate.ChatStream(ctx, messages, tools, model, options, onChunk)
}
func (p *HTTPProvider) EmbedQuery(
ctx context.Context,
input string,
model string,
dimensions int,
) ([]float64, error) {
return p.delegate.EmbedQuery(ctx, input, model, dimensions)
}
func (p *HTTPProvider) EmbedBatch(
ctx context.Context,
inputs []string,
model string,
dimensions int,
) ([][]float64, error) {
return p.delegate.EmbedBatch(ctx, inputs, model, dimensions)
}
func (p *HTTPProvider) GetDefaultModel() string {
return ""
}

View file

@ -7,6 +7,8 @@ type (
HTTPProvider = httpapi.HTTPProvider
)
var _ EmbeddingProvider = (*HTTPProvider)(nil)
func NewGeminiProvider(
apiKey string,
apiBase string,

View file

@ -0,0 +1,140 @@
package openai_compat
import (
"bytes"
"context"
"encoding/json"
"fmt"
"net/http"
"strings"
"github.com/sipeed/picoclaw/pkg/providers/common"
)
type embeddingResponse struct {
Data []struct {
Embedding []float64 `json:"embedding"`
} `json:"data"`
}
// EmbedQuery returns a single embedding vector for the given text.
func (p *Provider) EmbedQuery(
ctx context.Context,
input string,
model string,
dimensions int,
) ([]float64, error) {
vectors, err := p.EmbedBatch(ctx, []string{input}, model, dimensions)
if err != nil {
return nil, err
}
if len(vectors) == 0 {
return nil, fmt.Errorf("embedding response returned no vectors")
}
return vectors[0], nil
}
// EmbedBatch calls the OpenAI-compatible /embeddings endpoint and truncates
// the returned vectors locally when a smaller output width is requested.
// The request body intentionally omits dimensions so non-Matryoshka models,
// including vLLM-backed embeddings, never see an unsupported field upstream.
func (p *Provider) EmbedBatch(
ctx context.Context,
inputs []string,
model string,
dimensions int,
) ([][]float64, error) {
if p.apiBase == "" {
return nil, fmt.Errorf("API base not configured")
}
if len(inputs) == 0 {
return [][]float64{}, nil
}
if strings.TrimSpace(model) == "" {
return nil, fmt.Errorf("embedding model is required")
}
if dimensions < 0 {
return nil, fmt.Errorf("dimensions must be non-negative")
}
requestBody := map[string]any{
"model": normalizeModel(model, p.apiBase),
"input": inputs,
}
for key, value := range p.extraBody {
if strings.EqualFold(strings.TrimSpace(key), "dimensions") {
continue
}
requestBody[key] = value
}
jsonData, err := json.Marshal(requestBody)
if err != nil {
return nil, fmt.Errorf("failed to marshal embedding request: %w", err)
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, p.apiBase+"/embeddings", bytes.NewReader(jsonData))
if err != nil {
return nil, fmt.Errorf("failed to create embedding request: %w", err)
}
req.Header.Set("Content-Type", "application/json")
if p.userAgent != "" {
req.Header.Set("User-Agent", p.userAgent)
}
if p.apiKey != "" {
req.Header.Set("Authorization", "Bearer "+p.apiKey)
}
p.applyCustomHeaders(req)
resp, err := p.httpClient.Do(req)
if err != nil {
return nil, fmt.Errorf("failed to send embedding request: %w", err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, common.HandleErrorResponse(resp, p.apiBase)
}
var apiResp embeddingResponse
if err := json.NewDecoder(resp.Body).Decode(&apiResp); err != nil {
return nil, fmt.Errorf("failed to decode embedding response: %w", err)
}
if len(apiResp.Data) != len(inputs) {
return nil, fmt.Errorf(
"embedding response returned %d vectors for %d inputs",
len(apiResp.Data),
len(inputs),
)
}
results := make([][]float64, 0, len(apiResp.Data))
for i, item := range apiResp.Data {
vector, err := truncateEmbeddingVector(item.Embedding, dimensions)
if err != nil {
return nil, fmt.Errorf("embedding %d: %w", i, err)
}
results = append(results, vector)
}
return results, nil
}
func truncateEmbeddingVector(vector []float64, dimensions int) ([]float64, error) {
if dimensions < 0 {
return nil, fmt.Errorf("dimensions must be non-negative")
}
if dimensions == 0 {
return append([]float64(nil), vector...), nil
}
if len(vector) < dimensions {
return nil, fmt.Errorf(
"embedding length %d is shorter than requested dimensions %d",
len(vector),
dimensions,
)
}
return append([]float64(nil), vector[:dimensions]...), nil
}

View file

@ -0,0 +1,91 @@
package openai_compat
import (
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
)
func TestProviderEmbedBatch_OmitsDimensionsAndTruncatesLocally(t *testing.T) {
var requestBody map[string]any
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/embeddings" {
t.Fatalf("path = %s, want /embeddings", r.URL.Path)
}
if err := json.NewDecoder(r.Body).Decode(&requestBody); err != nil {
t.Fatalf("decode request body: %v", err)
}
response := map[string]any{
"data": []map[string]any{
{"embedding": []float64{1, 2, 3, 4}},
},
}
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(response)
}))
defer server.Close()
provider := NewProvider(
"test-key",
server.URL,
"",
WithExtraBody(map[string]any{
"dimensions": 8,
"encoding_format": "float",
"user": "picoclaw",
}),
)
vectors, err := provider.EmbedBatch(t.Context(), []string{"hello"}, "gemma-2b-embeddings", 2)
if err != nil {
t.Fatalf("EmbedBatch() error = %v", err)
}
if len(vectors) != 1 {
t.Fatalf("len(vectors) = %d, want 1", len(vectors))
}
if len(vectors[0]) != 2 || vectors[0][0] != 1 || vectors[0][1] != 2 {
t.Fatalf("vectors[0] = %#v, want [1 2]", vectors[0])
}
if _, ok := requestBody["dimensions"]; ok {
t.Fatalf("request body unexpectedly included dimensions: %#v", requestBody)
}
if got := requestBody["encoding_format"]; got != "float" {
t.Fatalf("encoding_format = %#v, want %q", got, "float")
}
if got := requestBody["user"]; got != "picoclaw" {
t.Fatalf("user = %#v, want %q", got, "picoclaw")
}
if got := requestBody["model"]; got != "gemma-2b-embeddings" {
t.Fatalf("model = %#v, want %q", got, "gemma-2b-embeddings")
}
input, ok := requestBody["input"].([]any)
if !ok || len(input) != 1 || input[0] != "hello" {
t.Fatalf("input = %#v, want [hello]", requestBody["input"])
}
}
func TestProviderEmbedQuery_ReturnsErrorWhenVectorIsTooShort(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
response := map[string]any{
"data": []map[string]any{
{"embedding": []float64{1, 2}},
},
}
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(response)
}))
defer server.Close()
provider := NewProvider("test-key", server.URL, "")
_, err := provider.EmbedQuery(t.Context(), "hello", "gemma-2b-embeddings", 3)
if err == nil {
t.Fatal("EmbedQuery() error = nil, want error")
}
if !strings.Contains(err.Error(), "shorter than requested dimensions") {
t.Fatalf("EmbedQuery() error = %q, want short-vector error", err)
}
}

View file

@ -33,6 +33,24 @@ type LLMProvider interface {
GetDefaultModel() string
}
// EmbeddingProvider exposes OpenAI-compatible embeddings helpers.
// Callers can request a local output width, but the provider must still send
// the full native vector upstream and truncate client-side if needed.
type EmbeddingProvider interface {
EmbedQuery(
ctx context.Context,
input string,
model string,
dimensions int,
) ([]float64, error)
EmbedBatch(
ctx context.Context,
inputs []string,
model string,
dimensions int,
) ([][]float64, error)
}
type StatefulProvider interface {
LLMProvider
Close()