Files
backend_fiesta/utils/embedding_test.go
2026-09-24 17:20:04 +05:30

162 lines
5.5 KiB
Go

package utils
import (
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"strings"
"testing"
"nearle/config"
)
func TestOpenAIEmbedderSendsTheRequestTheAPIExpects(t *testing.T) {
var got map[string]interface{}
var auth, path string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
auth, path = r.Header.Get("Authorization"), r.URL.Path
json.NewDecoder(r.Body).Decode(&got)
w.Write([]byte(`{"data":[{"embedding":[0.1,0.2,0.3]}]}`))
}))
defer srv.Close()
e, err := NewEmbedder(config.EmbeddingConfig{Provider: "openai", Model: "text-embedding-3-small", APIKey: "sk-test", BaseURL: srv.URL + "/v1/", Dimensions: 3})
if err != nil {
t.Fatal(err)
}
vec, err := e.Embed(context.Background(), "Milk Bikis")
if err != nil {
t.Fatal(err)
}
if len(vec) != 3 || vec[2] != 0.3 {
t.Errorf("vector = %v", vec)
}
if path != "/v1/embeddings" || auth != "Bearer sk-test" {
t.Errorf("path=%s auth=%s", path, auth)
}
if got["model"] != "text-embedding-3-small" || got["input"] != "Milk Bikis" || got["dimensions"] != float64(3) {
t.Errorf("body = %v", got)
}
if e.Model() != "text-embedding-3-small" {
t.Errorf("Model() = %q", e.Model())
}
}
func TestGeminiEmbedderSendsTheRequestTheAPIExpects(t *testing.T) {
var got map[string]interface{}
var key, path string
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
key, path = r.Header.Get("x-goog-api-key"), r.URL.Path
json.NewDecoder(r.Body).Decode(&got)
w.Write([]byte(`{"embedding":{"values":[0.5,0.6]}}`))
}))
defer srv.Close()
e, err := NewEmbedder(config.EmbeddingConfig{Provider: "gemini", Model: "gemini-embedding-001", APIKey: "g-test", BaseURL: srv.URL})
if err != nil {
t.Fatal(err)
}
vec, err := e.Embed(context.Background(), "Milk Bikis")
if err != nil {
t.Fatal(err)
}
if len(vec) != 2 || vec[0] != 0.5 {
t.Errorf("vector = %v", vec)
}
if path != "/models/gemini-embedding-001:embedContent" || key != "g-test" {
t.Errorf("path=%s key=%s", path, key)
}
if got["model"] != "models/gemini-embedding-001" || got["taskType"] != "RETRIEVAL_QUERY" {
t.Errorf("body = %v", got)
}
}
func TestEmbedderSurfacesProviderErrors(t *testing.T) {
// Used to assert this with a 429, which is now the one status that is
// deliberately NOT passed through — see the two tests below. The rule it was
// written for still holds for every other failure: a provider that explains
// itself should not have that explanation swallowed.
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(400)
w.Write([]byte(`{"error":{"message":"input exceeds the maximum token count"}}`))
}))
defer srv.Close()
e, _ := NewEmbedder(config.EmbeddingConfig{Provider: "openai", Model: "m", APIKey: "k", BaseURL: srv.URL})
_, err := e.Embed(context.Background(), "x")
if err == nil || !strings.Contains(err.Error(), "400") || !strings.Contains(err.Error(), "maximum token count") {
t.Fatalf("want a 400 with the provider's message, got %v", err)
}
}
func TestARateLimitIsWaitedOutRatherThanReported(t *testing.T) {
// The common case on a free tier: the quota clears in under a second, so
// the right answer is a slightly slower success rather than an error.
var calls int
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
calls++
if calls == 1 {
w.WriteHeader(429)
w.Write([]byte(`{"error":{"message":"Rate limit reached. Please try again in 20ms."}}`))
return
}
w.Write([]byte(`{"data":[{"embedding":[0.5,0.25]}]}`))
}))
defer srv.Close()
e, _ := NewEmbedder(config.EmbeddingConfig{Provider: "openai", Model: "m", APIKey: "k", BaseURL: srv.URL})
vector, err := e.Embed(context.Background(), "x")
if err != nil {
t.Fatalf("a 20ms quota wait became an error: %v", err)
}
if len(vector) != 2 {
t.Fatalf("the retry did not return the answer: %v", vector)
}
if calls != 2 {
t.Fatalf("expected one retry, saw %d calls", calls)
}
}
func TestAPersistentRateLimitIsReportedWithoutTheProvidersDetails(t *testing.T) {
// When waiting does not help, the person is told to try again — not handed
// our organisation id and the tokens-per-minute arithmetic behind it.
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(429)
w.Write([]byte(`{"error":{"message":"Rate limit reached for model X in organization org_01m38 on tokens per minute (TPM): Limit 8000. Please try again in 20ms."}}`))
}))
defer srv.Close()
e, _ := NewEmbedder(config.EmbeddingConfig{Provider: "openai", Model: "m", APIKey: "k", BaseURL: srv.URL})
_, err := e.Embed(context.Background(), "x")
if !errors.Is(err, ErrBusy) {
t.Fatalf("a persistent quota was not reported as busy: %v", err)
}
for _, leaked := range []string{"org_01m38", "TPM", "8000"} {
if strings.Contains(err.Error(), leaked) {
t.Fatalf("%q reached the caller: %s", leaked, err.Error())
}
}
}
func TestNoProviderMeansNoEmbedder(t *testing.T) {
e, err := NewEmbedder(config.EmbeddingConfig{})
if err != nil || e != nil {
t.Fatalf("got %v / %v", e, err)
}
if _, err := NewEmbedder(config.EmbeddingConfig{Provider: "cohere", Model: "m", APIKey: "k"}); err == nil {
t.Fatal("an unknown provider must be refused")
}
}
func TestVectorLiteral(t *testing.T) {
if got := VectorLiteral([]float32{0.1, -2, 3.5}); got != "[0.1,-2,3.5]" {
t.Errorf("got %q", got)
}
if got := VectorLiteral(nil); got != "[]" {
t.Errorf("got %q", got)
}
}