summaryrefslogtreecommitdiff
path: root/internal/httpapi
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
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 '')
-rw-r--r--internal/httpapi/api.go175
-rw-r--r--internal/httpapi/api_test.go281
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,