diff options
| author | Chia <Chia@93.nz> | 2026-08-06 15:58:57 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-06 15:58:57 +1200 |
| commit | 3f702084d20b3c3a3ea916f3110e99b22bda60b3 (patch) | |
| tree | 517f76c51025ce1ee085ea4898c60f799e5c37ea /internal/httpapi/api.go | |
| parent | 41e322c53d7b4b796eb377d0df9c29ecd10ba431 (diff) | |
feat: complete commercial developer workflowspublish-commercial-control-plane
Add tenant-safe usage observability, prepaid billing controls, API key lifecycle management, Embeddings metering, configurable billing alerts, and resilient provider health propagation. Harden Stripe failure handling, migrations, readiness, and the authenticated control-plane UI with end-to-end verification evidence.
Diffstat (limited to '')
| -rw-r--r-- | internal/httpapi/api.go | 62 |
1 files changed, 59 insertions, 3 deletions
diff --git a/internal/httpapi/api.go b/internal/httpapi/api.go index e66b170..ea4c0c2 100644 --- a/internal/httpapi/api.go +++ b/internal/httpapi/api.go @@ -150,6 +150,8 @@ func (a *API) registerInference(mux *http.ServeMux) { 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("POST /v1/embeddings", a.openAIEmbeddings) + mux.HandleFunc("POST /api/v1/embeddings", a.openAIEmbeddings) mux.HandleFunc("GET /anthropic/v1/models", a.anthropicModels) mux.HandleFunc("GET /api/anthropic/v1/models", a.anthropicModels) @@ -171,6 +173,10 @@ func (a *API) openAIResponses(w http.ResponseWriter, r *http.Request) { a.serveInference(w, r, domain.ProtocolOpenAIResponses) } +func (a *API) openAIEmbeddings(w http.ResponseWriter, r *http.Request) { + a.serveInference(w, r, domain.ProtocolOpenAIEmbeddings) +} + func (a *API) anthropicMessages(w http.ResponseWriter, r *http.Request) { a.serveInference(w, r, domain.ProtocolAnthropic) } @@ -243,6 +249,16 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do return } publicModel := model.ID + if protocol == domain.ProtocolOpenAIEmbeddings { + if envelope.Stream { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "unsupported_capability", Message: "Embeddings does not support streaming"}, requestID) + return + } + if !modelDeclaresCapability(model, "embeddings") { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "unsupported_capability", Message: "Model does not support embeddings"}, 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) @@ -284,7 +300,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do policy, _ = a.limiter.Policy(principal.ProjectID) } if err := a.billingMeter.Authorize(r.Context(), billing.Authorization{ - RequestID: requestID, Principal: principal, Model: model, Body: body, Policy: policy, + RequestID: requestID, Principal: principal, Model: model, Protocol: protocol, Body: body, Policy: policy, }); err != nil { if errors.Is(err, billing.ErrInsufficientBalance) { apierror.Write(w, apierror.Error{Status: http.StatusPaymentRequired, Type: "insufficient_balance", Message: "Account balance is insufficient"}, requestID) @@ -293,6 +309,13 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do ErrorType: "insufficient_balance", StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds()}) return } + if errors.Is(err, billing.ErrDailyQuotaExceeded) { + apierror.Write(w, apierror.Error{Status: http.StatusTooManyRequests, Type: "daily_quota_exceeded", Message: "API key daily spend quota exceeded"}, requestID) + a.recordUsageOnly(r, domain.UsageEvent{RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, + PublicModel: publicModel, Protocol: protocol, Stream: envelope.Stream, StatusCode: http.StatusTooManyRequests, Success: false, + ErrorType: "daily_quota_exceeded", 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, @@ -341,6 +364,16 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do }) return } + if protocol == domain.ProtocolOpenAIEmbeddings && !strings.HasPrefix(strings.ToLower(result.Response.Header.Get("Content-Type")), "application/json") { + apierror.Write(w, apierror.Error{Status: http.StatusBadGateway, Type: "invalid_provider_response", Message: "Upstream provider returned a non-JSON Embeddings response"}, requestID) + a.finishUsage(r, domain.UsageEvent{ + RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, + PublicModel: publicModel, ProviderID: result.Route.Provider.ID, UpstreamModel: result.Route.UpstreamModel, + Protocol: protocol, StatusCode: http.StatusBadGateway, Success: false, ErrorType: "invalid_provider_response", + Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(), + }) + return + } stream := envelope.Stream || strings.HasPrefix(strings.ToLower(result.Response.Header.Get("Content-Type")), "text/event-stream") copyResponseHeaders(w.Header(), result.Response.Header, stream) @@ -354,6 +387,18 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do copyErr := copyResponse(w, result.Response.Body, observer, stream) usageResult := observer.Usage() usageReported := observer.Reported() + ttftMS := int64(0) + if firstOutputAt := observer.FirstOutputAt(); !firstOutputAt.IsZero() { + ttftMS = firstOutputAt.Sub(startedAt).Milliseconds() + if ttftMS < 1 { + ttftMS = 1 + } + providerStartedAt := result.AttemptStartedAt + if providerStartedAt.IsZero() { + providerStartedAt = startedAt + } + a.forwarder.ObserveTTFT(model.ID, result.Route, firstOutputAt.Sub(providerStartedAt)) + } success = copyErr == nil errorType := "" if copyErr != nil { @@ -363,7 +408,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do RequestID: requestID, KeyID: principal.KeyID, TenantID: principal.TenantID, ProjectID: principal.ProjectID, 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, + ErrorType: errorType, Attempts: result.Attempts, StartedAt: startedAt, DurationMS: time.Since(startedAt).Milliseconds(), TTFTMS: ttftMS, Usage: usageResult, UsageReported: usageReported, }) a.logger.Info("inference_request", @@ -376,6 +421,7 @@ func (a *API) serveInference(w http.ResponseWriter, r *http.Request, protocol do "status", result.Response.StatusCode, "attempts", result.Attempts, "duration_ms", time.Since(startedAt).Milliseconds(), + "ttft_ms", ttftMS, ) } @@ -428,6 +474,7 @@ func modelProviderDescriptors(model domain.Model) []map[string]string { func (a *API) availableOpenAIModels(principal domain.Principal) []domain.Model { combined := append(a.catalog.ModelsFor(domain.ProtocolOpenAI, principal), a.catalog.ModelsFor(domain.ProtocolOpenAIResponses, principal)...) + combined = append(combined, a.catalog.ModelsFor(domain.ProtocolOpenAIEmbeddings, principal)...) seen := make(map[string]struct{}, len(combined)) result := make([]domain.Model, 0, len(combined)) for _, model := range combined { @@ -527,6 +574,15 @@ func modelHasCapability(model domain.Model, wanted string) bool { } return false } + +func modelDeclaresCapability(model domain.Model, wanted string) bool { + 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() @@ -653,7 +709,7 @@ func copyResponse(w http.ResponseWriter, body io.Reader, observer io.Writer, str if stream { destination = &flushingWriter{writer: w, controller: http.NewResponseController(w)} } - _, err := io.CopyBuffer(io.MultiWriter(destination, observer), body, make([]byte, 32<<10)) + _, err := io.CopyBuffer(io.MultiWriter(observer, destination), body, make([]byte, 32<<10)) return err } |
