Files
backend_fiesta/utils/embedding.go
2026-09-24 11:01:16 +05:30

233 lines
7.3 KiB
Go

package utils
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
"strings"
"time"
"nearle/config"
)
// Embedder turns a short piece of text — what Google Lens read off a packet —
// into the vector the catalogue was indexed with.
//
// One method on purpose. The scan pipeline needs exactly one thing from the
// model and nothing about which model it is; the provider is an environment
// decision (config.EmbeddingConfig) and the tests supply a fake.
type Embedder interface {
// Embed returns the vector for text. It must be the same length as the
// catalogue's `embedding` column or pgvector refuses the comparison.
Embed(ctx context.Context, text string) ([]float32, error)
// Model names what produced the vector, so a cache key can include it: a
// vector cached under one model must never be served for another.
Model() string
}
// ErrEmbedderNotConfigured is what the scan search sees when no provider is
// set. It falls back to text matching rather than failing the request.
var ErrEmbedderNotConfigured = errors.New("embedding provider is not configured")
// embedTimeout bounds one call to the provider. The scan endpoint has a
// customer waiting with a phone in their hand; a slow model is worse than a
// text-only answer, and the caller falls back on error.
const embedTimeout = 4 * time.Second
// NewEmbedder builds the provider named in the config, or returns nil when
// none is configured. A nil Embedder is a supported state everywhere it is
// used: the search degrades to text matching and says so in the response.
func NewEmbedder(cfg config.EmbeddingConfig) (Embedder, error) {
if !cfg.Enabled() {
return nil, nil
}
client := &http.Client{Timeout: embedTimeout}
switch cfg.Provider {
case "openai":
base := strings.TrimRight(cfg.BaseURL, "/")
if base == "" {
base = "https://api.openai.com/v1"
}
return &openAIEmbedder{cfg: cfg, base: base, client: client}, nil
case "gemini":
base := strings.TrimRight(cfg.BaseURL, "/")
if base == "" {
base = "https://generativelanguage.googleapis.com/v1beta"
}
return &geminiEmbedder{cfg: cfg, base: base, client: client}, nil
}
return nil, fmt.Errorf("embedding provider %q is not supported", cfg.Provider)
}
// ── OpenAI-compatible ───────────────────────────────────────────────────────
//
// POST {base}/embeddings — the shape OpenAI, Azure OpenAI (with a base URL),
// Ollama, vLLM, LM Studio and most hosted models all accept.
type openAIEmbedder struct {
cfg config.EmbeddingConfig
base string
client *http.Client
}
func (e *openAIEmbedder) Model() string { return e.cfg.Model }
func (e *openAIEmbedder) Embed(ctx context.Context, text string) ([]float32, error) {
body := map[string]interface{}{
"model": e.cfg.Model,
"input": text,
}
if e.cfg.Dimensions > 0 {
body["dimensions"] = e.cfg.Dimensions
}
var out struct {
Data []struct {
Embedding []float32 `json:"embedding"`
} `json:"data"`
Error *struct {
Message string `json:"message"`
} `json:"error"`
}
if err := postJSON(ctx, e.client, "embedding", e.base+"/embeddings", "Bearer "+e.cfg.APIKey, body, &out); err != nil {
return nil, err
}
if out.Error != nil {
return nil, fmt.Errorf("embedding: %s", out.Error.Message)
}
if len(out.Data) == 0 || len(out.Data[0].Embedding) == 0 {
return nil, errors.New("embedding: provider returned no vector")
}
return out.Data[0].Embedding, nil
}
// ── Gemini ──────────────────────────────────────────────────────────────────
//
// POST {base}/models/{model}:embedContent with the key as a header.
type geminiEmbedder struct {
cfg config.EmbeddingConfig
base string
client *http.Client
}
func (e *geminiEmbedder) Model() string { return e.cfg.Model }
func (e *geminiEmbedder) Embed(ctx context.Context, text string) ([]float32, error) {
model := e.cfg.Model
if !strings.HasPrefix(model, "models/") {
model = "models/" + model
}
body := map[string]interface{}{
"model": model,
"content": map[string]interface{}{"parts": []map[string]string{{"text": text}}},
"taskType": "RETRIEVAL_QUERY",
}
if e.cfg.Dimensions > 0 {
body["outputDimensionality"] = e.cfg.Dimensions
}
var out struct {
Embedding struct {
Values []float32 `json:"values"`
} `json:"embedding"`
Error *struct {
Message string `json:"message"`
} `json:"error"`
}
url := fmt.Sprintf("%s/%s:embedContent", e.base, model)
if err := postJSON(ctx, e.client, "embedding", url, "", body, &out, "x-goog-api-key", e.cfg.APIKey); err != nil {
return nil, err
}
if out.Error != nil {
return nil, fmt.Errorf("embedding: %s", out.Error.Message)
}
if len(out.Embedding.Values) == 0 {
return nil, errors.New("embedding: provider returned no vector")
}
return out.Embedding.Values, nil
}
// postJSON is the one HTTP call both providers make. Extra header pairs
// follow the body; `auth` is sent as Authorization when non-empty.
// postJSON is shared by the embedder and the chat gateway.
//
// `what` names the caller, and it is a parameter rather than a constant because
// it reaches a person. This helper used to say "embedding:" on every failure,
// so a rate-limited assistant told a shopkeeper "embedding: HTTP 429" — a
// sentence about a subsystem they have never heard of, describing something
// that was not involved.
func postJSON(ctx context.Context, client *http.Client, what, url, auth string, body, out interface{}, headers ...string) error {
payload, err := json.Marshal(body)
if err != nil {
return err
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, url, bytes.NewReader(payload))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
if auth != "" {
req.Header.Set("Authorization", auth)
}
for i := 0; i+1 < len(headers); i += 2 {
req.Header.Set(headers[i], headers[i+1])
}
resp, err := client.Do(req)
if err != nil {
return fmt.Errorf("%s: %w", what, err)
}
defer resp.Body.Close()
// Bounded: an error page from a misconfigured proxy should not be read to
// the end of the internet.
raw, err := io.ReadAll(io.LimitReader(resp.Body, 4<<20))
if err != nil {
return fmt.Errorf("%s: %w", what, err)
}
if err := json.Unmarshal(raw, out); err != nil {
return fmt.Errorf("%s: HTTP %d, unreadable body: %w", what, resp.StatusCode, err)
}
if resp.StatusCode/100 != 2 {
// The decoded body carries the provider's message where there is one;
// this is the fallback for a bare status.
if msg := extractMessage(raw); msg != "" {
return fmt.Errorf("%s: HTTP %d: %s", what, resp.StatusCode, msg)
}
return fmt.Errorf("%s: HTTP %d", what, resp.StatusCode)
}
return nil
}
func extractMessage(raw []byte) string {
var e struct {
Error struct {
Message string `json:"message"`
} `json:"error"`
}
if json.Unmarshal(raw, &e) == nil {
return e.Error.Message
}
return ""
}
// VectorLiteral renders a vector the way pgvector reads one: `[0.1,0.2,...]`.
func VectorLiteral(v []float32) string {
var b strings.Builder
b.Grow(len(v)*10 + 2)
b.WriteByte('[')
for i, f := range v {
if i > 0 {
b.WriteByte(',')
}
fmt.Fprintf(&b, "%g", f)
}
b.WriteByte(']')
return b.String()
}