summaryrefslogtreecommitdiff
path: root/internal/httpapi/api.go
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-05 22:01:29 +1200
committerChia <Chia@93.nz>2026-08-05 22:07:50 +1200
commiteadb2ffe85c43cf6fc741c9823cd28eedb4a844c (patch)
tree1aba2536d57360da403aa35c9ced58b615c7064e /internal/httpapi/api.go
parentcd0dd91ab93653631904f2ea0e574ccde6d60339 (diff)
feat: harden prepaid billing and commercial operations
Diffstat (limited to '')
-rw-r--r--internal/httpapi/api.go202
1 files changed, 142 insertions, 60 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
}