package utils import ( "context" "encoding/json" "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) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.WriteHeader(429) w.Write([]byte(`{"error":{"message":"Rate limit reached"}}`)) })) 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(), "429") || !strings.Contains(err.Error(), "Rate limit") { t.Fatalf("want a 429 with the provider's message, got %v", err) } } 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) } }