diff options
| author | Chia <Chia@93.nz> | 2026-08-06 09:29:41 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-06 09:32:46 +1200 |
| commit | 41e322c53d7b4b796eb377d0df9c29ecd10ba431 (patch) | |
| tree | c730526150e55e39b822d5197e4a20318ecaa449 /internal/httpapi/api_test.go | |
| parent | eadb2ffe85c43cf6fc741c9823cd28eedb4a844c (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.go | 281 |
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, |
