386 lines
13 KiB
Go
386 lines
13 KiB
Go
package utils
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
"reflect"
|
|
"strconv"
|
|
"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.
|
|
// postJSON sends the request, and sends it a second time if the provider said
|
|
// it was over its quota and named a wait we are willing to hold for.
|
|
//
|
|
// One retry, not a loop: past that, a queue forms behind a limit that is not
|
|
// going to lift, and the person is better told to try again than left watching
|
|
// a spinner. Both the chat gateway and the embedder go through here, so neither
|
|
// can be the one that forgot.
|
|
func postJSON(ctx context.Context, client *http.Client, what, url, auth string, body, out interface{}, headers ...string) error {
|
|
err := postJSONOnce(ctx, client, what, url, auth, body, out, headers...)
|
|
|
|
var busy *tooManyRequests
|
|
if errors.As(err, &busy) && waitBeforeRetry(ctx, busy.after) {
|
|
// `out` must be emptied first. The first attempt decoded the provider's
|
|
// error body into it, and `encoding/json` leaves fields the second
|
|
// payload does not mention exactly as it found them — so a retry that
|
|
// SUCCEEDED came back carrying the 429's `error` object, and every
|
|
// caller here checks that field before the data. The result was a
|
|
// successful call reported as the failure it had just recovered from.
|
|
resetForRetry(out)
|
|
return postJSONOnce(ctx, client, what, url, auth, body, out, headers...)
|
|
}
|
|
return err
|
|
}
|
|
|
|
// resetForRetry empties a decode target so a second attempt cannot inherit the
|
|
// first one's fields.
|
|
func resetForRetry(out interface{}) {
|
|
value := reflect.ValueOf(out)
|
|
if value.Kind() != reflect.Ptr || value.IsNil() {
|
|
return
|
|
}
|
|
value.Elem().Set(reflect.Zero(value.Elem().Type()))
|
|
}
|
|
|
|
func postJSONOnce(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 {
|
|
// Too many requests is the one status that is not about this request.
|
|
// The provider's own sentence is unusable here — Groq's reads
|
|
//
|
|
// "Rate limit reached for model `openai/gpt-oss-120b` in organization
|
|
// `org_01m38x8s72e759kn6g88ve2dhj` service tier `on_demand` on tokens
|
|
// per minute (TPM): Limit 8000, Used 7320…"
|
|
//
|
|
// which is shown to a shopkeeper as Buddy's answer, names our billing
|
|
// account, and tells them nothing they can act on. It is also usually
|
|
// over within a second, so the honest handling is to wait and try again
|
|
// rather than to report it at all.
|
|
if resp.StatusCode == http.StatusTooManyRequests {
|
|
return &tooManyRequests{what: what, after: retryAfter(resp, raw)}
|
|
}
|
|
// 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()
|
|
}
|
|
|
|
// ── Being rate limited ──────────────────────────────────────────────────────
|
|
//
|
|
// A shared provider quota is not a fault in the request that happened to hit
|
|
// it, and on Groq's free tier it is reached by ordinary use: 8,000 tokens a
|
|
// minute is three or four Buddy questions. The waits are short — the provider
|
|
// states them in milliseconds — so one retry turns almost all of them into a
|
|
// slightly slower answer instead of an error.
|
|
|
|
// ErrBusy is what a caller sees when the wait did not help.
|
|
//
|
|
// Sentinel so the HTTP layer can answer 429 and the console can say "a moment"
|
|
// rather than rendering a provider's billing details as an answer.
|
|
var ErrBusy = errors.New("the assistant is busy right now — try again in a moment")
|
|
|
|
type tooManyRequests struct {
|
|
what string
|
|
after time.Duration
|
|
}
|
|
|
|
func (e *tooManyRequests) Error() string { return e.what + ": " + ErrBusy.Error() }
|
|
func (e *tooManyRequests) Unwrap() error { return ErrBusy }
|
|
|
|
// maxRetryWait bounds how long a request may be held. Beyond this the honest
|
|
// answer is "busy" — a person watching a spinner has already decided something
|
|
// is broken, and the provider's own suggestion can be a minute on a hard quota.
|
|
const maxRetryWait = 3 * time.Second
|
|
|
|
// retryAfter reads how long the provider asked us to wait.
|
|
//
|
|
// `Retry-After` first, because it is the standard and a proxy may add it where
|
|
// the body has nothing. Groq puts the number in prose instead — "Please try
|
|
// again in 840ms" — so that is read next. Zero means "no idea", and the caller
|
|
// uses its own floor rather than hammering immediately.
|
|
func retryAfter(resp *http.Response, raw []byte) time.Duration {
|
|
if header := strings.TrimSpace(resp.Header.Get("Retry-After")); header != "" {
|
|
// Seconds, as an integer, is the only form worth reading: the HTTP-date
|
|
// form is for caches and no model provider sends it.
|
|
if secs, err := strconv.ParseFloat(header, 64); err == nil && secs > 0 {
|
|
return time.Duration(secs * float64(time.Second))
|
|
}
|
|
}
|
|
return waitFromMessage(extractMessage(raw))
|
|
}
|
|
|
|
// waitFromMessage pulls "try again in 840ms" or "try again in 1.5s" out of prose.
|
|
//
|
|
// Its own function because it is the part worth testing: the wording comes from
|
|
// somebody else's error strings and is the first thing that will change.
|
|
func waitFromMessage(message string) time.Duration {
|
|
lower := strings.ToLower(message)
|
|
marker := "try again in "
|
|
at := strings.Index(lower, marker)
|
|
if at < 0 {
|
|
return 0
|
|
}
|
|
|
|
rest := lower[at+len(marker):]
|
|
end := 0
|
|
for end < len(rest) && (rest[end] == '.' || (rest[end] >= '0' && rest[end] <= '9')) {
|
|
end++
|
|
}
|
|
if end == 0 {
|
|
return 0
|
|
}
|
|
amount, err := strconv.ParseFloat(rest[:end], 64)
|
|
if err != nil || amount <= 0 {
|
|
return 0
|
|
}
|
|
|
|
switch {
|
|
case strings.HasPrefix(rest[end:], "ms"):
|
|
return time.Duration(amount * float64(time.Millisecond))
|
|
case strings.HasPrefix(rest[end:], "s"):
|
|
return time.Duration(amount * float64(time.Second))
|
|
}
|
|
return 0
|
|
}
|
|
|
|
// waitBeforeRetry sleeps for what the provider asked, bounded, and reports
|
|
// whether waiting is worth it at all.
|
|
//
|
|
// Returns false when the ask is longer than we are prepared to hold a request
|
|
// for, or when the caller's context is done — a retry after the browser has
|
|
// given up is work nobody will see.
|
|
func waitBeforeRetry(ctx context.Context, after time.Duration) bool {
|
|
if after <= 0 {
|
|
// No stated wait. A short one anyway: retrying instantly on a quota is
|
|
// how a burst becomes two bursts.
|
|
after = 250 * time.Millisecond
|
|
}
|
|
if after > maxRetryWait {
|
|
return false
|
|
}
|
|
timer := time.NewTimer(after)
|
|
defer timer.Stop()
|
|
select {
|
|
case <-ctx.Done():
|
|
return false
|
|
case <-timer.C:
|
|
return true
|
|
}
|
|
}
|