summaryrefslogtreecommitdiff
path: root/internal/httpapi/api_test.go
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-06 09:29:41 +1200
committerChia <Chia@93.nz>2026-08-06 09:32:46 +1200
commit41e322c53d7b4b796eb377d0df9c29ecd10ba431 (patch)
treec730526150e55e39b822d5197e4a20318ecaa449 /internal/httpapi/api_test.go
parenteadb2ffe85c43cf6fc741c9823cd28eedb4a844c (diff)
feat: complete commercial control plane, billing, auth, and model catalog
- add PostgreSQL control-plane persistence with Redis-degraded hot reload - implement prepaid balance, usage ledger, Stripe top-up and reconciliation - add registration, email verification, password reset, invitations and RBAC - support TOTP, Passkey MFA, device sessions, quotas and rate limits - add tenant billing profiles, audit logs and operational readiness checks - build authenticated admin console, Quickstart, Playground and usage analytics - add public model catalog with pricing, filtering and cost estimation - support OpenAI Responses providers and provider health failover - validate real upstream usage reporting and balance settlement
Diffstat (limited to 'internal/httpapi/api_test.go')
-rw-r--r--internal/httpapi/api_test.go281
1 files changed, 279 insertions, 2 deletions
diff --git a/internal/httpapi/api_test.go b/internal/httpapi/api_test.go
index 9407530..cbb3967 100644
--- a/internal/httpapi/api_test.go
+++ b/internal/httpapi/api_test.go
@@ -20,6 +20,7 @@ import (
"aigw/internal/config"
"aigw/internal/domain"
"aigw/internal/provider"
+ "aigw/internal/providerhealth"
"aigw/internal/routing"
"aigw/internal/telemetry"
)
@@ -28,6 +29,49 @@ 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
settled chan domain.UsageEvent
@@ -134,6 +178,199 @@ func TestProxyFailsOverBeforeWritingResponse(t *testing.T) {
}
}
+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) {
@@ -194,6 +431,45 @@ func TestAnthropicHeadersAndPath(t *testing.T) {
}
}
+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.UsageReported {
+ t.Fatalf("unexpected Responses usage event: %+v", event)
+ }
+}
+
func TestInsufficientBalanceRejectsBeforeCallingUpstream(t *testing.T) {
var calls atomic.Int64
upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
@@ -276,12 +552,13 @@ func newTestGatewayWithBilling(t *testing.T, providers []config.ProviderConfig,
}
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),
- Forwarder: provider.New(cfg.UpstreamHTTP, metrics),
+ Router: routing.New(modelCatalog, routeHealth),
+ Forwarder: provider.New(cfg.UpstreamHTTP, metrics, routeHealth),
UsageSink: sink,
BillingMeter: meter,
Metrics: metrics,