package httpapi import ( "bufio" "bytes" "context" "encoding/json" "io" "log/slog" "net/http" "net/http/httptest" "strings" "sync/atomic" "testing" "time" "aigw/internal/auth" "aigw/internal/billing" "aigw/internal/catalog" "aigw/internal/config" "aigw/internal/domain" "aigw/internal/provider" "aigw/internal/providerhealth" "aigw/internal/routing" "aigw/internal/telemetry" ) type captureUsageSink struct { events chan domain.UsageEvent } func TestInferenceBrowserOriginCORS(t *testing.T) { api := New(Options{ BrowserOrigin: "https://console.example.test/admin/", Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), Metrics: &telemetry.Metrics{}, }) handler := api.InferenceHandler() request := httptest.NewRequest(http.MethodOptions, "/v1/responses", nil) request.Header.Set("Origin", "https://console.example.test") request.Header.Set("Access-Control-Request-Method", http.MethodPost) request.Header.Set("Access-Control-Request-Headers", "authorization,content-type") response := httptest.NewRecorder() handler.ServeHTTP(response, request) if response.Code != http.StatusNoContent { t.Fatalf("preflight status = %d, want %d", response.Code, http.StatusNoContent) } if got := response.Header().Get("Access-Control-Allow-Origin"); got != "https://console.example.test" { t.Fatalf("allow origin = %q", got) } if response.Header().Get("Access-Control-Allow-Credentials") != "" { t.Fatal("inference CORS must not allow browser credentials") } if !strings.Contains(response.Header().Get("Access-Control-Expose-Headers"), "X-AIGW-Request-ID") { t.Fatal("request ID is not exposed to the developer console") } blocked := httptest.NewRequest(http.MethodOptions, "/v1/responses", nil) blocked.Header.Set("Origin", "https://attacker.example") blockedResponse := httptest.NewRecorder() handler.ServeHTTP(blockedResponse, blocked) if blockedResponse.Code != http.StatusForbidden { t.Fatalf("untrusted preflight status = %d, want %d", blockedResponse.Code, http.StatusForbidden) } blockedRequest := httptest.NewRequest(http.MethodGet, "/v1/models", nil) blockedRequest.Header.Set("Origin", "https://attacker.example") blockedActual := httptest.NewRecorder() handler.ServeHTTP(blockedActual, blockedRequest) if blockedActual.Code != http.StatusForbidden { t.Fatalf("untrusted actual status = %d, want %d", blockedActual.Code, http.StatusForbidden) } } type fakeBillingMeter struct { authorizeErr error authorized chan billing.Authorization settled chan domain.UsageEvent } func (m *fakeBillingMeter) Authorize(_ context.Context, authorization billing.Authorization) error { if m.authorized != nil { m.authorized <- authorization } return m.authorizeErr } func (m *fakeBillingMeter) EnqueueSettlement(_ context.Context, event domain.UsageEvent) error { m.settled <- event return nil } func (s *captureUsageSink) Publish(event domain.UsageEvent) { s.events <- event } func TestOpenAIProxyRewritesModelAndEmitsUsage(t *testing.T) { upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/v1/chat/completions" { t.Errorf("unexpected path: %s", r.URL.Path) } if r.Header.Get("Authorization") != "Bearer upstream-secret" { t.Errorf("upstream authorization leaked or missing: %q", r.Header.Get("Authorization")) } var request map[string]any if err := json.NewDecoder(r.Body).Decode(&request); err != nil { t.Error(err) } if request["model"] != "upstream-model" { t.Errorf("model was not rewritten: %+v", request) } w.Header().Set("Content-Type", "application/json") _, _ = io.WriteString(w, `{"id":"chat-1","choices":[],"usage":{"prompt_tokens":3,"completion_tokens":2,"total_tokens":5}}`) })) defer upstream.Close() gateway, sink := newTestGateway(t, []config.ProviderConfig{{ ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: upstream.URL + "/v1", APIKey: "upstream-secret", }}, []config.RouteConfig{{Provider: "primary", UpstreamModel: "upstream-model", Weight: 1}}) defer gateway.Close() request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/chat/completions", strings.NewReader(`{"model":"public/model","messages":[{"role":"user","content":"hello"}]}`)) request.Header.Set("Authorization", "Bearer client-secret") response, err := http.DefaultClient.Do(request) if err != nil { t.Fatal(err) } defer response.Body.Close() if response.StatusCode != http.StatusOK { payload, _ := io.ReadAll(response.Body) t.Fatalf("unexpected status %d: %s", response.StatusCode, payload) } if response.Header.Get("X-AIGW-Request-ID") == "" { t.Fatal("missing gateway request id") } select { case event := <-sink.events: if event.PublicModel != "public/model" || event.UpstreamModel != "upstream-model" || event.Usage.TotalTokens != 5 || event.TTFTMS < 1 || !event.Success { t.Fatalf("unexpected usage event: %+v", event) } case <-time.After(time.Second): t.Fatal("usage event was not emitted") } } func TestProxyFailsOverBeforeWritingResponse(t *testing.T) { var primaryCalls atomic.Int64 primary := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { primaryCalls.Add(1) w.WriteHeader(http.StatusServiceUnavailable) })) defer primary.Close() var fallbackCalls atomic.Int64 fallback := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { fallbackCalls.Add(1) w.Header().Set("Content-Type", "application/json") _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) })) defer fallback.Close() gateway, sink := newTestGateway(t, []config.ProviderConfig{ {ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: primary.URL + "/v1", APIKey: "one"}, {ID: "fallback", Protocol: domain.ProtocolOpenAI, BaseURL: fallback.URL + "/v1", APIKey: "two"}, }, []config.RouteConfig{ {Provider: "primary", UpstreamModel: "model", Priority: 0, Weight: 1}, {Provider: "fallback", UpstreamModel: "model", Priority: 10, Weight: 1}, }, ) defer gateway.Close() response := postOpenAI(t, gateway.URL, false) defer response.Body.Close() if response.StatusCode != http.StatusOK || primaryCalls.Load() != 1 || fallbackCalls.Load() != 1 { t.Fatalf("failover did not complete: status=%d primary=%d fallback=%d", response.StatusCode, primaryCalls.Load(), fallbackCalls.Load()) } event := <-sink.events if event.Attempts != 2 || event.ProviderID != "fallback" { t.Fatalf("unexpected failover event: %+v", event) } } func TestGatewayLearnsTTFTAndPrefersFasterProvider(t *testing.T) { var fastCalls atomic.Int64 fast := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { fastCalls.Add(1) w.Header().Set("Content-Type", "application/json") _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) })) defer fast.Close() var slowCalls atomic.Int64 slow := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { slowCalls.Add(1) time.Sleep(25 * time.Millisecond) w.Header().Set("Content-Type", "application/json") _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) })) defer slow.Close() gateway, sink := newTestGateway(t, []config.ProviderConfig{ {ID: "fast", Protocol: domain.ProtocolOpenAI, BaseURL: fast.URL + "/v1", APIKey: "one"}, {ID: "slow", Protocol: domain.ProtocolOpenAI, BaseURL: slow.URL + "/v1", APIKey: "two"}, }, []config.RouteConfig{ {Provider: "fast", UpstreamModel: "model", Weight: 1}, {Provider: "slow", UpstreamModel: "model", Weight: 1}, }, ) defer gateway.Close() for range 20 { response := postOpenAI(t, gateway.URL, false) _, _ = io.Copy(io.Discard, response.Body) _ = response.Body.Close() if response.StatusCode != http.StatusOK { t.Fatalf("unexpected gateway status: %d", response.StatusCode) } select { case <-sink.events: case <-time.After(time.Second): t.Fatal("usage event was not emitted") } } if fastCalls.Load() != 16 || slowCalls.Load() != 4 { t.Fatalf("TTFT feedback was not applied with bounded exploration: fast=%d slow=%d", fastCalls.Load(), slowCalls.Load()) } } func TestProviderSelectorPinsRouteAndKeepsCanonicalUsageModel(t *testing.T) { var primaryCalls atomic.Int64 primary := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { primaryCalls.Add(1) w.WriteHeader(http.StatusServiceUnavailable) })) defer primary.Close() var backupCalls atomic.Int64 backup := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { backupCalls.Add(1) var request map[string]any if err := json.NewDecoder(r.Body).Decode(&request); err != nil { t.Error(err) } if request["model"] != "backup-model" { t.Errorf("upstream model = %v, want backup-model", request["model"]) } w.Header().Set("Content-Type", "application/json") _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) })) defer backup.Close() gateway, sink := newTestGateway(t, []config.ProviderConfig{ {ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: primary.URL + "/v1", APIKey: "one"}, {ID: "backup", Protocol: domain.ProtocolOpenAI, BaseURL: backup.URL + "/v1", APIKey: "two"}, }, []config.RouteConfig{ {Provider: "primary", UpstreamModel: "primary-model", Priority: 0, Weight: 1}, {Provider: "backup", UpstreamModel: "backup-model", Priority: 10, Weight: 1}, }, ) defer gateway.Close() request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/chat/completions", strings.NewReader(`{"model":"public/model:backup","messages":[{"role":"user","content":"hello"}]}`)) request.Header.Set("Authorization", "Bearer client-secret") response, err := http.DefaultClient.Do(request) if err != nil { t.Fatal(err) } defer response.Body.Close() if response.StatusCode != http.StatusOK { payload, _ := io.ReadAll(response.Body) t.Fatalf("status = %d: %s", response.StatusCode, payload) } if primaryCalls.Load() != 0 || backupCalls.Load() != 1 { t.Fatalf("pinned routing calls: primary=%d backup=%d", primaryCalls.Load(), backupCalls.Load()) } event := <-sink.events if event.PublicModel != "public/model" || event.ProviderID != "backup" || event.Attempts != 1 { t.Fatalf("unexpected pinned usage event: %+v", event) } } func TestProviderSelectorRejectsUnknownProviderWithoutCallingUpstream(t *testing.T) { var calls atomic.Int64 upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { calls.Add(1) w.WriteHeader(http.StatusOK) })) defer upstream.Close() gateway, _ := newTestGateway(t, []config.ProviderConfig{{ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: upstream.URL + "/v1", APIKey: "one"}}, []config.RouteConfig{{Provider: "primary", UpstreamModel: "model", Weight: 1}}, ) defer gateway.Close() request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/chat/completions", strings.NewReader(`{"model":"public/model:missing","messages":[{"role":"user","content":"hello"}]}`)) request.Header.Set("Authorization", "Bearer client-secret") response, err := http.DefaultClient.Do(request) if err != nil { t.Fatal(err) } defer response.Body.Close() raw, err := io.ReadAll(response.Body) if err != nil { t.Fatal(err) } var payload struct { Error struct { Type string `json:"type"` } `json:"error"` } if err := json.Unmarshal(raw, &payload); err != nil { t.Fatal(err) } if response.StatusCode != http.StatusNotFound || payload.Error.Type != "provider_not_found" || calls.Load() != 0 { t.Fatalf("status=%d type=%q calls=%d", response.StatusCode, payload.Error.Type, calls.Load()) } } func TestResolveModelSelectorChecksBaseModelAllowlistAndPreservesExactColonID(t *testing.T) { modelCatalog := catalog.NewModels([]domain.Model{ {ID: "public/model"}, {ID: "exact:model", AllowedKeyIDs: map[string]struct{}{"other-key": {}}}, {ID: "exact"}, }) api := &API{catalog: modelCatalog} principal := domain.Principal{KeyID: "key-1", AllowedModels: map[string]struct{}{"public/model": {}, "exact": {}}} model, providerSlug, err := api.resolveModelSelector("public/model:backup", principal) if err != nil || model.ID != "public/model" || providerSlug != "backup" { t.Fatalf("base allowlist selector: model=%+v provider=%q err=%v", model, providerSlug, err) } if _, _, err := api.resolveModelSelector("exact:model", principal); err == nil { t.Fatal("an unauthorized exact colon model ID must not be reinterpreted as a provider selector") } } func TestModelsListPublishesProviderSlugsWithoutUpstreamDetails(t *testing.T) { gateway, _ := newTestGateway(t, []config.ProviderConfig{{ID: "openai-primary", Protocol: domain.ProtocolOpenAI, BaseURL: "https://secret-upstream.example/v1", APIKey: "secret"}}, []config.RouteConfig{{Provider: "openai-primary", UpstreamModel: "secret-upstream-model", Weight: 1}}, ) defer gateway.Close() request, _ := http.NewRequest(http.MethodGet, gateway.URL+"/v1/models", nil) request.Header.Set("Authorization", "Bearer client-secret") response, err := http.DefaultClient.Do(request) if err != nil { t.Fatal(err) } defer response.Body.Close() raw, err := io.ReadAll(response.Body) if err != nil { t.Fatal(err) } var payload struct { Data []struct { ID string `json:"id"` Providers []struct { Slug string `json:"slug"` WireAPI string `json:"wire_api"` } `json:"providers"` } `json:"data"` } if err := json.Unmarshal(raw, &payload); err != nil { t.Fatal(err) } if response.StatusCode != http.StatusOK || len(payload.Data) != 1 || len(payload.Data[0].Providers) != 1 || payload.Data[0].Providers[0].Slug != "openai-primary" { t.Fatalf("unexpected models payload: status=%d payload=%+v", response.StatusCode, payload) } if strings.Contains(string(raw), "secret-upstream") { t.Fatalf("models payload leaked upstream detail: %s", raw) } } func TestCircuitBreakerSkipsFailingProviderOnSubsequentRequests(t *testing.T) { var primaryCalls atomic.Int64 primary := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { primaryCalls.Add(1) w.WriteHeader(http.StatusServiceUnavailable) })) defer primary.Close() var fallbackCalls atomic.Int64 fallback := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { fallbackCalls.Add(1) w.Header().Set("Content-Type", "application/json") _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) })) defer fallback.Close() gateway, sink := newTestGateway(t, []config.ProviderConfig{ {ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: primary.URL + "/v1", APIKey: "one"}, {ID: "fallback", Protocol: domain.ProtocolOpenAI, BaseURL: fallback.URL + "/v1", APIKey: "two"}, }, []config.RouteConfig{ {Provider: "primary", UpstreamModel: "model", Priority: 0, Weight: 1}, {Provider: "fallback", UpstreamModel: "model", Priority: 10, Weight: 1}, }, ) defer gateway.Close() for requestNumber := range 4 { response := postOpenAI(t, gateway.URL, false) _, _ = io.Copy(io.Discard, response.Body) _ = response.Body.Close() if response.StatusCode != http.StatusOK { t.Fatalf("request %d status = %d, want %d", requestNumber+1, response.StatusCode, http.StatusOK) } select { case <-sink.events: case <-time.After(time.Second): t.Fatalf("request %d did not emit usage", requestNumber+1) } } if primaryCalls.Load() != 3 { t.Fatalf("primary calls = %d, want 3 before circuit opens", primaryCalls.Load()) } if fallbackCalls.Load() != 4 { t.Fatalf("fallback calls = %d, want 4", fallbackCalls.Load()) } } func TestSSEIsFlushedBeforeUpstreamCompletes(t *testing.T) { release := make(chan struct{}) upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "text/event-stream") _, _ = io.WriteString(w, "data: {\"id\":\"chunk-1\",\"choices\":[]}\n\n") w.(http.Flusher).Flush() <-release _, _ = io.WriteString(w, "data: [DONE]\n\n") })) defer upstream.Close() gateway, _ := newTestGateway(t, []config.ProviderConfig{{ ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: upstream.URL + "/v1", APIKey: "secret", }}, []config.RouteConfig{{Provider: "primary", UpstreamModel: "model", Weight: 1}}) defer gateway.Close() response := postOpenAI(t, gateway.URL, true) defer response.Body.Close() reader := bufio.NewReader(response.Body) line, err := reader.ReadString('\n') if err != nil { close(release) t.Fatal(err) } if !strings.Contains(line, "chunk-1") { close(release) t.Fatalf("unexpected first SSE line: %q", line) } close(release) } func TestAnthropicHeadersAndPath(t *testing.T) { upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/v1/messages" || r.Header.Get("x-api-key") != "anthropic-upstream" || r.Header.Get("anthropic-version") != "2023-06-01" { t.Errorf("unexpected anthropic request: path=%s key=%q version=%q", r.URL.Path, r.Header.Get("x-api-key"), r.Header.Get("anthropic-version")) } _, _ = io.WriteString(w, `{"type":"message","content":[],"usage":{"input_tokens":4,"output_tokens":6}}`) })) defer upstream.Close() gateway, sink := newTestGateway(t, []config.ProviderConfig{{ ID: "anthropic", Protocol: domain.ProtocolAnthropic, BaseURL: upstream.URL + "/v1", APIKey: "anthropic-upstream", }}, []config.RouteConfig{{Provider: "anthropic", UpstreamModel: "claude-upstream", Weight: 1}}) defer gateway.Close() request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/anthropic/v1/messages", strings.NewReader(`{"model":"public/model","max_tokens":10,"messages":[{"role":"user","content":"hello"}]}`)) request.Header.Set("x-api-key", "client-secret") request.Header.Set("anthropic-version", "2023-06-01") response, err := http.DefaultClient.Do(request) if err != nil { t.Fatal(err) } defer response.Body.Close() if response.StatusCode != http.StatusOK { t.Fatalf("unexpected status: %d", response.StatusCode) } event := <-sink.events if event.Usage.InputTokens != 4 || event.Usage.OutputTokens != 6 { t.Fatalf("unexpected anthropic usage: %+v", event.Usage) } } func TestOpenAIResponsesProxyRewritesModelAndEmitsUsage(t *testing.T) { upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/responses" { t.Errorf("unexpected path: %s", r.URL.Path) } var request map[string]any if err := json.NewDecoder(r.Body).Decode(&request); err != nil { t.Error(err) } if request["model"] != "gpt-upstream" || request["input"] != "hello" { t.Errorf("unexpected Responses request: %+v", request) } w.Header().Set("Content-Type", "application/json") _, _ = io.WriteString(w, `{"id":"resp_1","object":"response","status":"completed","model":"gpt-upstream","output":[],"usage":{"input_tokens":9,"output_tokens":4,"total_tokens":13}}`) })) defer upstream.Close() gateway, sink := newTestGateway(t, []config.ProviderConfig{{ ID: "responses", Protocol: domain.ProtocolOpenAI, WireAPI: "responses", BaseURL: upstream.URL, APIKey: "upstream-secret", }}, []config.RouteConfig{{Provider: "responses", UpstreamModel: "gpt-upstream", Weight: 1}}) defer gateway.Close() request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/responses", strings.NewReader(`{"model":"public/model","input":"hello"}`)) request.Header.Set("Authorization", "Bearer client-secret") response, err := http.DefaultClient.Do(request) if err != nil { t.Fatal(err) } defer response.Body.Close() if response.StatusCode != http.StatusOK { payload, _ := io.ReadAll(response.Body) t.Fatalf("unexpected status %d: %s", response.StatusCode, payload) } event := <-sink.events if event.Protocol != domain.ProtocolOpenAIResponses || event.UpstreamModel != "gpt-upstream" || event.Usage.TotalTokens != 13 || event.TTFTMS < 1 || !event.UsageReported { t.Fatalf("unexpected Responses usage event: %+v", event) } } func TestOpenAIEmbeddingsProxyRewritesModelAndMetersInputTokens(t *testing.T) { upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path != "/v1/embeddings" { t.Errorf("unexpected path: %s", r.URL.Path) } if r.Header.Get("Authorization") != "Bearer upstream-secret" { t.Errorf("unexpected upstream authorization: %q", r.Header.Get("Authorization")) } var request map[string]any if err := json.NewDecoder(r.Body).Decode(&request); err != nil { t.Error(err) } if request["model"] != "embedding-upstream" || request["input"] != "hello vector" { t.Errorf("unexpected Embeddings request: %+v", request) } w.Header().Set("Content-Type", "application/json") _, _ = io.WriteString(w, `{"object":"list","data":[{"object":"embedding","index":0,"embedding":[0.1,0.2]}],"model":"embedding-upstream","usage":{"prompt_tokens":8,"total_tokens":8}}`) })) defer upstream.Close() meter := &fakeBillingMeter{authorized: make(chan billing.Authorization, 1), settled: make(chan domain.UsageEvent, 1)} gateway, _ := newTestGatewayWithBilling(t, []config.ProviderConfig{{ ID: "embeddings", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: upstream.URL + "/v1", APIKey: "upstream-secret", }}, []config.RouteConfig{{Provider: "embeddings", UpstreamModel: "embedding-upstream", Weight: 1}}, meter) defer gateway.Close() request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/embeddings", strings.NewReader(`{"model":"public/model","input":"hello vector"}`)) request.Header.Set("Authorization", "Bearer client-secret") response, err := http.DefaultClient.Do(request) if err != nil { t.Fatal(err) } defer response.Body.Close() if response.StatusCode != http.StatusOK { payload, _ := io.ReadAll(response.Body) t.Fatalf("unexpected status %d: %s", response.StatusCode, payload) } authorization := <-meter.authorized if authorization.Protocol != domain.ProtocolOpenAIEmbeddings { t.Fatalf("authorization protocol = %q", authorization.Protocol) } event := <-meter.settled if event.Protocol != domain.ProtocolOpenAIEmbeddings || event.UpstreamModel != "embedding-upstream" || event.Usage.InputTokens != 8 || event.Usage.OutputTokens != 0 || event.Usage.TotalTokens != 8 || !event.UsageReported { t.Fatalf("unexpected Embeddings usage event: %+v", event) } } func TestOpenAIEmbeddingsRequiresDeclaredCapability(t *testing.T) { providerConfig := config.ProviderConfig{ID: "embeddings", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: "https://example.invalid/v1", APIKey: "secret"} modelCatalog := catalog.New(config.Config{Providers: []config.ProviderConfig{providerConfig}, Models: []config.ModelConfig{{ ID: "public/model", Routes: []config.RouteConfig{{Provider: "embeddings", UpstreamModel: "embedding-upstream", Weight: 1}}, }}}) authenticator, err := auth.NewStatic(`[{"key":"client-secret","key_id":"key-1","tenant_id":"tenant-1","project_id":"project-1","scopes":["inference"]}]`, false) if err != nil { t.Fatal(err) } api := New(Options{Authenticator: authenticator, Catalog: modelCatalog, Router: routing.New(modelCatalog), Metrics: &telemetry.Metrics{}, Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), MaxBodyBytes: 1 << 20}) request := httptest.NewRequest(http.MethodPost, "/v1/embeddings", strings.NewReader(`{"model":"public/model","input":"hello"}`)) request.Header.Set("Authorization", "Bearer client-secret") response := httptest.NewRecorder() api.Handler().ServeHTTP(response, request) if response.Code != http.StatusBadRequest || !strings.Contains(response.Body.String(), "unsupported_capability") { t.Fatalf("unexpected response %d: %s", response.Code, response.Body.String()) } } func TestOpenAIEmbeddingsRejectsNonJSONSuccessBeforeResponseStarts(t *testing.T) { upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { w.Header().Set("Content-Type", "text/html; charset=utf-8") _, _ = io.WriteString(w, "provider console") })) defer upstream.Close() meter := &fakeBillingMeter{settled: make(chan domain.UsageEvent, 1)} gateway, _ := newTestGatewayWithBilling(t, []config.ProviderConfig{{ ID: "embeddings", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: upstream.URL, APIKey: "secret", }}, []config.RouteConfig{{Provider: "embeddings", UpstreamModel: "embedding-model", Weight: 1}}, meter) defer gateway.Close() request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/embeddings", strings.NewReader(`{"model":"public/model","input":"hello"}`)) request.Header.Set("Authorization", "Bearer client-secret") response, err := http.DefaultClient.Do(request) if err != nil { t.Fatal(err) } defer response.Body.Close() body, _ := io.ReadAll(response.Body) if response.StatusCode != http.StatusBadGateway || !strings.Contains(string(body), "invalid_provider_response") || strings.Contains(string(body), "provider console") { t.Fatalf("unexpected response %d: %s", response.StatusCode, body) } event := <-meter.settled if event.Success || event.StatusCode != http.StatusBadGateway || event.ErrorType != "invalid_provider_response" || event.UsageReported { t.Fatalf("unexpected invalid provider usage event: %+v", event) } } func TestInsufficientBalanceRejectsBeforeCallingUpstream(t *testing.T) { var calls atomic.Int64 upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { calls.Add(1) w.WriteHeader(http.StatusOK) })) defer upstream.Close() meter := &fakeBillingMeter{authorizeErr: billing.ErrInsufficientBalance, settled: make(chan domain.UsageEvent, 1)} gateway, sink := newTestGatewayWithBilling(t, []config.ProviderConfig{{ ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: upstream.URL + "/v1", APIKey: "secret", }}, []config.RouteConfig{{Provider: "primary", UpstreamModel: "model", Weight: 1}}, meter) defer gateway.Close() response := postOpenAI(t, gateway.URL, false) defer response.Body.Close() if response.StatusCode != http.StatusPaymentRequired { t.Fatalf("status = %d, want 402", response.StatusCode) } if calls.Load() != 0 { t.Fatalf("upstream calls = %d, want 0", calls.Load()) } select { case <-meter.settled: t.Fatal("rejected request was settled") default: } select { case event := <-sink.events: if event.StatusCode != http.StatusPaymentRequired || event.ErrorType != "insufficient_balance" { t.Fatalf("unexpected rejection usage: %+v", event) } case <-time.After(time.Second): t.Fatal("rejected request did not emit usage") } } func TestSuccessfulRequestIsSettled(t *testing.T) { upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":2,"completion_tokens":3,"total_tokens":5}}`) })) defer upstream.Close() meter := &fakeBillingMeter{settled: make(chan domain.UsageEvent, 1)} gateway, _ := newTestGatewayWithBilling(t, []config.ProviderConfig{{ ID: "primary", Protocol: domain.ProtocolOpenAI, BaseURL: upstream.URL + "/v1", APIKey: "secret", }}, []config.RouteConfig{{Provider: "primary", UpstreamModel: "model", Weight: 1}}, meter) defer gateway.Close() response := postOpenAI(t, gateway.URL, false) defer response.Body.Close() if response.StatusCode != http.StatusOK { t.Fatalf("status = %d, want 200", response.StatusCode) } select { case event := <-meter.settled: if event.Usage.TotalTokens != 5 { t.Fatalf("settled usage = %+v", event.Usage) } case <-time.After(time.Second): t.Fatal("request was not settled") } } func newTestGateway(t *testing.T, providers []config.ProviderConfig, routes []config.RouteConfig) (*httptest.Server, *captureUsageSink) { return newTestGatewayWithBilling(t, providers, routes, nil) } func newTestGatewayWithBilling(t *testing.T, providers []config.ProviderConfig, routes []config.RouteConfig, meter billing.Meter) (*httptest.Server, *captureUsageSink) { t.Helper() authenticator, err := auth.NewStatic(`[{"key":"client-secret","key_id":"key-1","tenant_id":"tenant-1","project_id":"project-1","scopes":["inference"]}]`, false) if err != nil { t.Fatal(err) } modelConfig := config.ModelConfig{ID: "public/model", OwnedBy: "test", Routes: routes} for _, providerConfig := range providers { if providerConfig.WireAPI == "embeddings" { modelConfig.Capabilities = []string{"embeddings"} break } } cfg := config.Config{ Providers: providers, Models: []config.ModelConfig{modelConfig}, UpstreamHTTP: config.UpstreamHTTPConfig{ MaxIdleConnections: 100, MaxIdleConnectionsPerHost: 20, IdleConnectionTimeoutSecs: 10, ResponseHeaderTimeoutSecs: 2, }, } metrics := &telemetry.Metrics{} modelCatalog := catalog.New(cfg) routeHealth := providerhealth.New(providerhealth.Options{}) sink := &captureUsageSink{events: make(chan domain.UsageEvent, 10)} api := New(Options{ Authenticator: authenticator, Catalog: modelCatalog, Router: routing.New(modelCatalog, routeHealth), Forwarder: provider.New(cfg.UpstreamHTTP, metrics, routeHealth), UsageSink: sink, BillingMeter: meter, Metrics: metrics, Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), MaxBodyBytes: 1 << 20, }) return httptest.NewServer(api.Handler()), sink } func postOpenAI(t *testing.T, gatewayURL string, stream bool) *http.Response { t.Helper() payload := []byte(`{"model":"public/model","messages":[{"role":"user","content":"hello"}],"stream":false}`) if stream { payload = []byte(`{"model":"public/model","messages":[{"role":"user","content":"hello"}],"stream":true}`) } request, _ := http.NewRequest(http.MethodPost, gatewayURL+"/v1/chat/completions", bytes.NewReader(payload)) request.Header.Set("Authorization", "Bearer client-secret") response, err := http.DefaultClient.Do(request) if err != nil { t.Fatal(err) } return response }