diff options
Diffstat (limited to 'internal/httpapi/api.go')
| -rw-r--r-- | internal/httpapi/api.go | 175 |
1 files changed, 159 insertions, 16 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 |
