diff options
Diffstat (limited to 'internal/httpapi')
| -rw-r--r-- | internal/httpapi/api.go | 202 | ||||
| -rw-r--r-- | internal/httpapi/api_test.go | 2 | ||||
| -rw-r--r-- | internal/httpapi/proxy.go | 44 |
3 files changed, 187 insertions, 61 deletions
diff --git a/internal/httpapi/api.go b/internal/httpapi/api.go index becfbf3..6ab437c 100644 --- a/internal/httpapi/api.go +++ b/internal/httpapi/api.go @@ -28,54 +28,57 @@ import ( type requestIDKey struct{} +type UsageRecorder interface { + RecordUsage(context.Context, domain.UsageEvent) error +} + type API struct { - authenticator auth.Authenticator - catalog *catalog.Catalog - router *routing.Router - forwarder *provider.Forwarder - usageSink telemetry.UsageSink - billingMeter billing.Meter - limiter *limits.Limiter - usageRecorder interface { - RecordUsage(context.Context, domain.UsageEvent) error - } - metrics *telemetry.Metrics - logger *slog.Logger - maxBodyBytes int64 - exposeMetrics bool + authenticator auth.Authenticator + catalog *catalog.Catalog + router *routing.Router + forwarder *provider.Forwarder + usageSink telemetry.UsageSink + billingMeter billing.Meter + limiter *limits.Limiter + usageRecorder UsageRecorder + metrics *telemetry.Metrics + logger *slog.Logger + maxBodyBytes int64 + exposeMetrics bool + deploymentRegion string } type Options struct { - Authenticator auth.Authenticator - Catalog *catalog.Catalog - Router *routing.Router - Forwarder *provider.Forwarder - UsageSink telemetry.UsageSink - BillingMeter billing.Meter - Limiter *limits.Limiter - UsageRecorder interface { - RecordUsage(context.Context, domain.UsageEvent) error - } - Metrics *telemetry.Metrics - Logger *slog.Logger - MaxBodyBytes int64 - ExposeMetrics bool + Authenticator auth.Authenticator + Catalog *catalog.Catalog + Router *routing.Router + Forwarder *provider.Forwarder + UsageSink telemetry.UsageSink + BillingMeter billing.Meter + Limiter *limits.Limiter + UsageRecorder UsageRecorder + Metrics *telemetry.Metrics + Logger *slog.Logger + MaxBodyBytes int64 + ExposeMetrics bool + DeploymentRegion string } func New(options Options) *API { return &API{ - authenticator: options.Authenticator, - catalog: options.Catalog, - router: options.Router, - forwarder: options.Forwarder, - usageSink: options.UsageSink, - billingMeter: options.BillingMeter, - limiter: options.Limiter, - usageRecorder: options.UsageRecorder, - metrics: options.Metrics, - logger: options.Logger, - maxBodyBytes: options.MaxBodyBytes, - exposeMetrics: options.ExposeMetrics, + authenticator: options.Authenticator, + catalog: options.Catalog, + router: options.Router, + forwarder: options.Forwarder, + usageSink: options.UsageSink, + billingMeter: options.BillingMeter, + limiter: options.Limiter, + usageRecorder: options.UsageRecorder, + metrics: options.Metrics, + logger: options.Logger, + maxBodyBytes: options.MaxBodyBytes, + exposeMetrics: options.ExposeMetrics, + deploymentRegion: strings.ToLower(strings.TrimSpace(options.DeploymentRegion)), } } @@ -87,6 +90,17 @@ func (a *API) Handler() http.Handler { mux.Handle("GET /metrics", a.metrics) } + a.registerInference(mux) + return a.withRequestID(a.recoverPanics(mux)) +} + +func (a *API) InferenceHandler() http.Handler { + mux := http.NewServeMux() + a.registerInference(mux) + return a.withRequestID(a.recoverPanics(mux)) +} + +func (a *API) registerInference(mux *http.ServeMux) { mux.HandleFunc("GET /v1/models", a.openAIModels) mux.HandleFunc("GET /api/v1/models", a.openAIModels) mux.HandleFunc("POST /v1/chat/completions", a.openAIChat) @@ -97,7 +111,6 @@ func (a *API) Handler() http.Handler { mux.HandleFunc("POST /anthropic/v1/messages", a.anthropicMessages) mux.HandleFunc("POST /api/anthropic/v1/messages", a.anthropicMessages) - return a.withRequestID(a.recoverPanics(mux)) } func (a *API) health(w http.ResponseWriter, _ *http.Request) { @@ -143,8 +156,13 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do } var envelope struct { - Model string `json:"model"` - Stream bool `json:"stream"` + Model string `json:"model"` + Stream bool `json:"stream"` + MaxTokens int64 `json:"max_tokens"` + MaxCompletionTokens int64 `json:"max_completion_tokens"` + MaxOutputTokens int64 `json:"max_output_tokens"` + Tools json.RawMessage `json:"tools"` + ResponseFormat json.RawMessage `json:"response_format"` } if err := json.Unmarshal(body, &envelope); err != nil || strings.TrimSpace(envelope.Model) == "" { apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "invalid_params", Message: "Parameter model is required and the body must be valid JSON"}, requestID) @@ -170,7 +188,30 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do defer lease.Release() } - routes, err := a.router.Plan(envelope.Model, protocol) + model, modelErr := a.catalog.ModelForPrincipal(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 + } + 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) + return + } + if envelope.Stream && !modelHasCapability(model, "streaming") { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "unsupported_capability", Message: "Model does not support streaming"}, requestID) + return + } + if len(envelope.Tools) > 0 && string(envelope.Tools) != "null" && !modelHasCapability(model, "tools") { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "unsupported_capability", Message: "Model does not support tools"}, requestID) + return + } + if len(envelope.ResponseFormat) > 0 && string(envelope.ResponseFormat) != "null" && !modelHasCapability(model, "json") { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "unsupported_capability", Message: "Model does not support structured JSON output"}, requestID) + return + } + + routes, err := a.router.Plan(model.ID, protocol) if err != nil { 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) @@ -180,11 +221,6 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do return } if a.billingMeter != nil { - model, modelErr := a.catalog.Model(envelope.Model) - if modelErr != nil { - apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "invalid_model", Message: "Model does not exist"}, requestID) - return - } policy := domain.LimitPolicy{} if a.limiter != nil { policy, _ = a.limiter.Policy(principal.ProjectID) @@ -257,6 +293,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do observer := usage.NewObserver(protocol, stream) copyErr := copyResponse(w, result.Response.Body, observer, stream) usageResult := observer.Usage() + usageReported := observer.Reported() success = copyErr == nil errorType := "" if copyErr != nil { @@ -267,6 +304,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do PublicModel: envelope.Model, 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, }) a.logger.Info("inference_request", "request_id", requestID, @@ -281,22 +319,27 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do } func (a *API) openAIModels(w http.ResponseWriter, r *http.Request) { - if !a.authorize(w, r) { + principal, ok := a.authorize(w, r) + if !ok { return } - models := a.catalog.Models(domain.ProtocolOpenAI) + models := a.availableModels(a.catalog.ModelsFor(domain.ProtocolOpenAI, 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": 0, "owned_by": model.OwnedBy}) + 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}) } writeJSON(w, map[string]any{"object": "list", "data": data}) } func (a *API) anthropicModels(w http.ResponseWriter, r *http.Request) { - if !a.authorize(w, r) { + principal, ok := a.authorize(w, r) + if !ok { return } - models := a.catalog.Models(domain.ProtocolAnthropic) + 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"}) @@ -309,13 +352,52 @@ func (a *API) anthropicModels(w http.ResponseWriter, r *http.Request) { writeJSON(w, response) } -func (a *API) authorize(w http.ResponseWriter, r *http.Request) bool { +func (a *API) modelAvailableInRegion(model domain.Model) bool { + if a.deploymentRegion == "" || len(model.Regions) == 0 { + return true + } + for _, region := range model.Regions { + if strings.EqualFold(region, a.deploymentRegion) || region == "*" { + return true + } + } + return false +} + +func (a *API) availableModels(models []domain.Model) []domain.Model { + result := models[:0] + for _, model := range models { + if a.modelAvailableInRegion(model) { + result = append(result, model) + } + } + return result +} +func modelHasCapability(model domain.Model, wanted string) bool { + if len(model.Capabilities) == 0 { + return true + } + for _, value := range model.Capabilities { + if value == wanted || value == "*" { + return true + } + } + return false +} +func modelCreated(model domain.Model) int64 { + if model.ReleasedAt != nil { + return model.ReleasedAt.Unix() + } + return 0 +} + +func (a *API) authorize(w http.ResponseWriter, r *http.Request) (domain.Principal, bool) { principal, err := a.authenticator.Authenticate(r) if err != nil || !hasScope(principal, "inference") { apierror.Write(w, apierror.Error{Status: http.StatusForbidden, Type: "access_denied", Message: "Invalid API key or insufficient permission"}, requestIDFrom(r.Context())) - return false + return domain.Principal{}, false } - return true + return principal, true } func (a *API) publishUsage(event domain.UsageEvent) { @@ -327,11 +409,11 @@ func (a *API) publishUsage(event domain.UsageEvent) { func (a *API) finishUsage(r *http.Request, event domain.UsageEvent) { settled := false if a.billingMeter != nil { - ctx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), 5*time.Second) - err := a.billingMeter.Settle(ctx, event) + ctx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), 500*time.Millisecond) + err := a.billingMeter.EnqueueSettlement(ctx, event) cancel() if err != nil { - a.logger.Error("billing_settlement_failed", "request_id", event.RequestID, "tenant_id", event.TenantID, "error", err) + a.logger.Error("billing_settlement_enqueue_failed", "request_id", event.RequestID, "tenant_id", event.TenantID, "error", err) } else { settled = true } diff --git a/internal/httpapi/api_test.go b/internal/httpapi/api_test.go index 014b99b..9407530 100644 --- a/internal/httpapi/api_test.go +++ b/internal/httpapi/api_test.go @@ -37,7 +37,7 @@ func (m *fakeBillingMeter) Authorize(context.Context, billing.Authorization) err return m.authorizeErr } -func (m *fakeBillingMeter) Settle(_ context.Context, event domain.UsageEvent) error { +func (m *fakeBillingMeter) EnqueueSettlement(_ context.Context, event domain.UsageEvent) error { m.settled <- event return nil } diff --git a/internal/httpapi/proxy.go b/internal/httpapi/proxy.go new file mode 100644 index 0000000..3cd2797 --- /dev/null +++ b/internal/httpapi/proxy.go @@ -0,0 +1,44 @@ +package httpapi + +import ( + "net" + "net/http" + "strings" +) + +// TrustProxyHeaders accepts forwarding metadata only from explicitly trusted +// CIDRs. This prevents a direct client from forging HTTPS or audit IP state. +func TrustProxyHeaders(next http.Handler, trustedCIDRs []string, requireHTTPS bool) (http.Handler, error) { + trusted := make([]*net.IPNet, 0, len(trustedCIDRs)) + for _, value := range trustedCIDRs { + _, network, err := net.ParseCIDR(value) + if err != nil { + return nil, err + } + trusted = append(trusted, network) + } + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + host, _, _ := net.SplitHostPort(r.RemoteAddr) + remote := net.ParseIP(host) + trustedPeer := false + for _, network := range trusted { + if remote != nil && network.Contains(remote) { + trustedPeer = true + break + } + } + if !trustedPeer { + for _, header := range []string{"Forwarded", "X-Forwarded-For", "X-Forwarded-Host", "X-Forwarded-Port", "X-Forwarded-Proto", "X-Real-IP"} { + r.Header.Del(header) + } + } else if forwarded := strings.TrimSpace(strings.Split(r.Header.Get("X-Forwarded-For"), ",")[0]); net.ParseIP(forwarded) != nil { + r.RemoteAddr = net.JoinHostPort(forwarded, "0") + } + secure := r.TLS != nil || (trustedPeer && strings.EqualFold(strings.TrimSpace(r.Header.Get("X-Forwarded-Proto")), "https")) + if requireHTTPS && !secure { + http.Error(w, "HTTPS is required", http.StatusUpgradeRequired) + return + } + next.ServeHTTP(w, r) + }), nil +} |
