summaryrefslogtreecommitdiff
path: root/internal/httpapi/api.go
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/httpapi/api.go175
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