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 | |
| 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')
| -rw-r--r-- | internal/httpapi/api.go | 175 | ||||
| -rw-r--r-- | internal/httpapi/api_test.go | 281 |
2 files changed, 438 insertions, 18 deletions
diff --git a/internal/httpapi/api.go b/internal/httpapi/api.go index 6ab437c..e66b170 100644 --- a/internal/httpapi/api.go +++ b/internal/httpapi/api.go @@ -10,6 +10,7 @@ import ( "io" "log/slog" "net/http" + "net/url" "runtime/debug" "strings" "time" @@ -46,6 +47,7 @@ type API struct { maxBodyBytes int64 exposeMetrics bool deploymentRegion string + browserOrigin string } type Options struct { @@ -62,6 +64,7 @@ type Options struct { MaxBodyBytes int64 ExposeMetrics bool DeploymentRegion string + BrowserOrigin string } func New(options Options) *API { @@ -79,6 +82,7 @@ func New(options Options) *API { maxBodyBytes: options.MaxBodyBytes, exposeMetrics: options.ExposeMetrics, deploymentRegion: strings.ToLower(strings.TrimSpace(options.DeploymentRegion)), + browserOrigin: browserOrigin(options.BrowserOrigin), } } @@ -91,13 +95,52 @@ func (a *API) Handler() http.Handler { } a.registerInference(mux) - return a.withRequestID(a.recoverPanics(mux)) + return a.withRequestID(a.withBrowserOrigin(a.recoverPanics(mux))) } func (a *API) InferenceHandler() http.Handler { mux := http.NewServeMux() a.registerInference(mux) - return a.withRequestID(a.recoverPanics(mux)) + return a.withRequestID(a.withBrowserOrigin(a.recoverPanics(mux))) +} + +// withBrowserOrigin permits the authenticated console to call the split +// inference listener. It deliberately does not allow credentials or wildcard +// origins: the API key remains an explicit bearer credential in the request. +func (a *API) withBrowserOrigin(next http.Handler) http.Handler { + if a.browserOrigin == "" { + return next + } + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + origin := strings.TrimSpace(r.Header.Get("Origin")) + if origin == "" { + next.ServeHTTP(w, r) + return + } + if origin != a.browserOrigin { + w.WriteHeader(http.StatusForbidden) + return + } + w.Header().Set("Vary", "Origin") + w.Header().Set("Access-Control-Allow-Origin", a.browserOrigin) + w.Header().Set("Access-Control-Expose-Headers", "X-AIGW-Request-ID") + if r.Method == http.MethodOptions { + w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS") + w.Header().Set("Access-Control-Allow-Headers", "Authorization, Content-Type, X-API-Key, Anthropic-Version, X-Request-ID") + w.Header().Set("Access-Control-Max-Age", "600") + w.WriteHeader(http.StatusNoContent) + return + } + next.ServeHTTP(w, r) + }) +} + +func browserOrigin(raw string) string { + u, err := url.Parse(strings.TrimSpace(raw)) + if err != nil || u.Host == "" || (u.Scheme != "http" && u.Scheme != "https") { + return "" + } + return strings.ToLower(u.Scheme) + "://" + u.Host } func (a *API) registerInference(mux *http.ServeMux) { @@ -105,6 +148,8 @@ func (a *API) registerInference(mux *http.ServeMux) { mux.HandleFunc("GET /api/v1/models", a.openAIModels) mux.HandleFunc("POST /v1/chat/completions", a.openAIChat) mux.HandleFunc("POST /api/v1/chat/completions", a.openAIChat) + mux.HandleFunc("POST /v1/responses", a.openAIResponses) + mux.HandleFunc("POST /api/v1/responses", a.openAIResponses) mux.HandleFunc("GET /anthropic/v1/models", a.anthropicModels) mux.HandleFunc("GET /api/anthropic/v1/models", a.anthropicModels) @@ -122,6 +167,10 @@ func (a *API) openAIChat(w http.ResponseWriter, r *http.Request) { a.serveInference(w, r, domain.ProtocolOpenAI) } +func (a *API) openAIResponses(w http.ResponseWriter, r *http.Request) { + a.serveInference(w, r, domain.ProtocolOpenAIResponses) +} + func (a *API) anthropicMessages(w http.ResponseWriter, r *http.Request) { a.serveInference(w, r, domain.ProtocolAnthropic) } @@ -188,11 +237,12 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do defer lease.Release() } - model, modelErr := a.catalog.ModelForPrincipal(envelope.Model, principal) + model, providerSlug, modelErr := a.resolveModelSelector(envelope.Model, principal) if modelErr != nil || !a.modelAvailableInRegion(model) { apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "invalid_model", Message: "Model does not exist or is not allowed for this API key"}, requestID) return } + publicModel := model.ID requestedOutput := max(envelope.MaxTokens, envelope.MaxCompletionTokens, envelope.MaxOutputTokens) if model.MaxOutputTokens > 0 && requestedOutput > model.MaxOutputTokens { apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "invalid_params", Message: "Requested output exceeds the model maximum"}, requestID) @@ -211,9 +261,17 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do return } - routes, err := a.router.Plan(model.ID, protocol) + routes, err := a.router.PlanProvider(model.ID, protocol, providerSlug) if err != nil { - if errors.Is(err, routing.ErrNoRoute) { + if errors.Is(err, routing.ErrProviderNotFound) { + apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "provider_not_found", Message: "Requested provider is not configured for this model and API protocol"}, requestID) + } else if errors.Is(err, routing.ErrNoHealthyRoute) { + message := "All compatible providers are cooling down after retryable failures" + if providerSlug != "" { + message = "Requested provider is cooling down after retryable failures" + } + apierror.Write(w, apierror.Error{Status: http.StatusServiceUnavailable, Type: "provider_unavailable", Message: message}, requestID) + } else if errors.Is(err, routing.ErrNoRoute) { apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "model_not_supported", Message: "Model does not support this API protocol"}, requestID) } else { apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "invalid_model", Message: "Model does not exist"}, requestID) @@ -231,28 +289,30 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do if errors.Is(err, billing.ErrInsufficientBalance) { apierror.Write(w, apierror.Error{Status: http.StatusPaymentRequired, Type: "insufficient_balance", Message: "Account balance is insufficient"}, requestID) a.recordUsageOnly(r, domain.UsageEvent{RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, - PublicModel: envelope.Model, Protocol: protocol, Stream: envelope.Stream, StatusCode: http.StatusPaymentRequired, Success: false, + PublicModel: publicModel, Protocol: protocol, Stream: envelope.Stream, StatusCode: http.StatusPaymentRequired, Success: false, ErrorType: "insufficient_balance", StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds()}) return } if errors.Is(err, billing.ErrQuotaExceeded) { apierror.Write(w, apierror.Error{Status: http.StatusTooManyRequests, Type: "monthly_quota_exceeded", Message: "Project monthly spend quota exceeded"}, requestID) a.recordUsageOnly(r, domain.UsageEvent{RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, - PublicModel: envelope.Model, Protocol: protocol, Stream: envelope.Stream, StatusCode: http.StatusTooManyRequests, Success: false, + PublicModel: publicModel, Protocol: protocol, Stream: envelope.Stream, StatusCode: http.StatusTooManyRequests, Success: false, ErrorType: "monthly_quota_exceeded", StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds()}) return } a.logger.Error("billing_authorization_failed", "request_id", requestID, "tenant_id", principal.TenantID, "error", err) apierror.Write(w, apierror.Error{Status: http.StatusServiceUnavailable, Type: "billing_unavailable", Message: "Billing service is temporarily unavailable"}, requestID) a.recordUsageOnly(r, domain.UsageEvent{RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, - PublicModel: envelope.Model, Protocol: protocol, Stream: envelope.Stream, StatusCode: http.StatusServiceUnavailable, Success: false, + PublicModel: publicModel, Protocol: protocol, Stream: envelope.Stream, StatusCode: http.StatusServiceUnavailable, Success: false, ErrorType: "billing_unavailable", StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds()}) return } } - result, err := a.forwarder.Forward(r.Context(), protocol, requestID, body, r.Header, routes) + result, err := a.forwarder.Forward(r.Context(), protocol, requestID, model.ID, body, r.Header, routes) if err != nil { + a.logger.Warn("inference_upstream_failed", "request_id", requestID, "model", publicModel, + "protocol", protocol, "attempts", result.Attempts, "error", err) errorType := "no_provider_available" status := http.StatusBadGateway if errors.Is(err, context.Canceled) { @@ -263,7 +323,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do } a.finishUsage(r, domain.UsageEvent{ RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, - PublicModel: envelope.Model, Protocol: protocol, Stream: envelope.Stream, StatusCode: status, + PublicModel: publicModel, Protocol: protocol, Stream: envelope.Stream, StatusCode: status, Success: false, ErrorType: errorType, Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(), }) return @@ -275,7 +335,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do apierror.Write(w, gatewayError, requestID) a.finishUsage(r, domain.UsageEvent{ RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, - PublicModel: envelope.Model, ProviderID: result.Route.Provider.ID, UpstreamModel: result.Route.UpstreamModel, + PublicModel: publicModel, ProviderID: result.Route.Provider.ID, UpstreamModel: result.Route.UpstreamModel, Protocol: protocol, Stream: envelope.Stream, StatusCode: gatewayError.Status, Success: false, ErrorType: gatewayError.Type, Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(), }) @@ -301,7 +361,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do } a.finishUsage(r, domain.UsageEvent{ RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, - PublicModel: envelope.Model, ProviderID: result.Route.Provider.ID, UpstreamModel: result.Route.UpstreamModel, + PublicModel: publicModel, ProviderID: result.Route.Provider.ID, UpstreamModel: result.Route.UpstreamModel, Protocol: protocol, Stream: stream, StatusCode: result.Response.StatusCode, Success: success, ErrorType: errorType, Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(), Usage: usageResult, UsageReported: usageReported, @@ -310,7 +370,8 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do "request_id", requestID, "tenant_id", principal.TenantID, "project_id", principal.ProjectID, - "model", envelope.Model, + "model", publicModel, + "provider_slug", result.Route.Provider.EffectiveSlug(), "provider", result.Route.Provider.ID, "status", result.Response.StatusCode, "attempts", result.Attempts, @@ -323,17 +384,62 @@ func (a *API) openAIModels(w http.ResponseWriter, r *http.Request) { if !ok { return } - models := a.availableModels(a.catalog.ModelsFor(domain.ProtocolOpenAI, principal)) + models := a.availableOpenAIModels(principal) data := make([]map[string]any, 0, len(models)) for _, model := range models { data = append(data, map[string]any{"id": model.ID, "object": "model", "created": modelCreated(model), "owned_by": model.OwnedBy, "display_name": model.DisplayName, "context_window": model.ContextWindow, "max_output_tokens": model.MaxOutputTokens, "input_modalities": model.InputModalities, "output_modalities": model.OutputModalities, "capabilities": model.Capabilities, - "lifecycle": model.Lifecycle, "regions": model.Regions, "replacement_model": model.ReplacementModel}) + "lifecycle": model.Lifecycle, "regions": model.Regions, "replacement_model": model.ReplacementModel, + "supported_wire_apis": modelWireAPIs(model), "providers": modelProviderDescriptors(model)}) } writeJSON(w, map[string]any{"object": "list", "data": data}) } +func modelWireAPIs(model domain.Model) []string { + seen := map[string]struct{}{} + result := make([]string, 0, len(model.Routes)) + for _, route := range model.Routes { + wireAPI := route.Provider.EffectiveWireAPI() + if _, exists := seen[wireAPI]; exists { + continue + } + seen[wireAPI] = struct{}{} + result = append(result, wireAPI) + } + return result +} + +func modelProviderDescriptors(model domain.Model) []map[string]string { + seen := map[string]struct{}{} + result := make([]map[string]string, 0, len(model.Routes)) + for _, route := range model.Routes { + slug := route.Provider.EffectiveSlug() + wireAPI := route.Provider.EffectiveWireAPI() + key := slug + "\x00" + wireAPI + if _, exists := seen[key]; exists { + continue + } + seen[key] = struct{}{} + result = append(result, map[string]string{"slug": slug, "wire_api": wireAPI}) + } + return result +} + +func (a *API) availableOpenAIModels(principal domain.Principal) []domain.Model { + combined := append(a.catalog.ModelsFor(domain.ProtocolOpenAI, principal), a.catalog.ModelsFor(domain.ProtocolOpenAIResponses, principal)...) + seen := make(map[string]struct{}, len(combined)) + result := make([]domain.Model, 0, len(combined)) + for _, model := range combined { + if _, exists := seen[model.ID]; exists || !a.modelAvailableInRegion(model) { + continue + } + seen[model.ID] = struct{}{} + result = append(result, model) + } + return result +} + func (a *API) anthropicModels(w http.ResponseWriter, r *http.Request) { principal, ok := a.authorize(w, r) if !ok { @@ -342,7 +448,7 @@ func (a *API) anthropicModels(w http.ResponseWriter, r *http.Request) { models := a.availableModels(a.catalog.ModelsFor(domain.ProtocolAnthropic, principal)) data := make([]map[string]any, 0, len(models)) for _, model := range models { - data = append(data, map[string]any{"id": model.ID, "display_name": model.ID, "created_at": "1970-01-01T00:00:00Z", "type": "model"}) + data = append(data, map[string]any{"id": model.ID, "display_name": model.ID, "created_at": "1970-01-01T00:00:00Z", "type": "model", "providers": modelProviderDescriptors(model)}) } response := map[string]any{"data": data, "has_more": false, "first_id": nil, "last_id": nil} if len(models) > 0 { @@ -373,6 +479,43 @@ func (a *API) availableModels(models []domain.Model) []domain.Model { } return result } + +func (a *API) resolveModelSelector(selector string, principal domain.Principal) (domain.Model, string, error) { + selector = strings.TrimSpace(selector) + if model, err := a.catalog.Model(selector); err == nil { + if !model.Allows(principal) { + return domain.Model{}, "", fmt.Errorf("model %q not allowed", selector) + } + return model, "", nil + } + separator := strings.LastIndexByte(selector, ':') + if separator <= 0 || separator == len(selector)-1 { + return domain.Model{}, "", fmt.Errorf("model %q not found or not allowed", selector) + } + modelID := strings.TrimSpace(selector[:separator]) + providerSlug := strings.TrimSpace(selector[separator+1:]) + if modelID == "" || !validProviderSlug(providerSlug) { + return domain.Model{}, "", fmt.Errorf("invalid model provider selector %q", selector) + } + model, err := a.catalog.ModelForPrincipal(modelID, principal) + if err != nil { + return domain.Model{}, "", err + } + return model, providerSlug, nil +} + +func validProviderSlug(value string) bool { + if len(value) < 3 || len(value) > 64 || value[0] == '-' || value[len(value)-1] == '-' { + return false + } + for _, character := range value { + if (character < 'a' || character > 'z') && (character < '0' || character > '9') && character != '-' { + return false + } + } + return true +} + func modelHasCapability(model domain.Model, wanted string) bool { if len(model.Capabilities) == 0 { return true 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, |
