From 3f702084d20b3c3a3ea916f3110e99b22bda60b3 Mon Sep 17 00:00:00 2001 From: Chia Date: Thu, 6 Aug 2026 15:58:57 +1200 Subject: feat: complete commercial developer workflows 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. --- internal/adminapi/api.go | 216 ++++++++++--- internal/adminapi/bootstrap_test.go | 57 ++++ internal/adminapi/model_page.go | 276 ++++++++++++++++ internal/adminapi/model_page_test.go | 44 +++ internal/adminapi/usage_test.go | 26 ++ internal/adminui/assets/app.js | 146 ++++++--- internal/adminui/assets/index.html | 35 ++- internal/adminui/assets/models.css | 21 +- internal/adminui/assets/models.js | 4 +- internal/adminui/assets/style.css | 7 +- internal/auth/static.go | 6 +- internal/auth/static_test.go | 5 +- internal/billing/auto_topup.go | 56 ++-- internal/billing/auto_topup_test.go | 43 +++ internal/billing/ledger.go | 77 +++++ internal/billing/operations.go | 8 +- internal/billing/operations_test.go | 58 ++++ internal/billing/service.go | 115 +++++-- internal/billing/service_test.go | 174 +++++++++++ internal/billing/stripe.go | 71 +++-- internal/billing/stripe_preflight.go | 124 ++++++++ internal/billing/stripe_preflight_test.go | 90 ++++++ internal/billing/types.go | 39 ++- internal/catalog/catalog.go | 19 ++ internal/config/config.go | 74 ++++- internal/config/config_test.go | 83 +++++ internal/controlplane/api_key_test.go | 26 ++ internal/controlplane/mail_operations.go | 16 +- .../mail_operations_integration_test.go | 144 +++++++++ internal/controlplane/manager.go | 38 ++- internal/controlplane/manager_test.go | 49 +++ internal/controlplane/mutations.go | 130 +++++++- internal/controlplane/preferences.go | 75 +++-- internal/controlplane/preferences_test.go | 18 +- internal/controlplane/queries.go | 35 ++- internal/controlplane/schema.sql | 40 ++- internal/controlplane/snapshot.go | 9 +- internal/controlplane/store.go | 48 ++- internal/controlplane/store_integration_test.go | 84 ++++- internal/controlplane/types.go | 64 ++++ internal/controlplane/usage.go | 82 ++++- internal/controlplane/usage_analytics.go | 61 +++- internal/controlplane/usage_integration_test.go | 43 ++- internal/controlplane/usage_test.go | 26 ++ internal/domain/types.go | 22 +- internal/httpapi/api.go | 62 +++- internal/httpapi/api_test.go | 161 +++++++++- internal/limits/limits.go | 116 ++++--- internal/limits/limits_test.go | 38 +++ internal/provider/forwarder.go | 24 +- internal/provider/forwarder_test.go | 18 ++ internal/providerhealth/prober.go | 171 ++++++++++ internal/providerhealth/prober_test.go | 91 ++++++ internal/providerhealth/redis_history.go | 347 +++++++++++++++++++++ internal/providerhealth/tracker.go | 162 +++++++++- internal/providerhealth/tracker_test.go | 152 +++++++++ internal/routing/router.go | 123 +++++++- internal/routing/router_test.go | 45 +++ internal/telemetry/metrics.go | 78 ++++- internal/usage/observer.go | 88 +++++- internal/usage/observer_test.go | 71 +++++ 61 files changed, 4211 insertions(+), 420 deletions(-) create mode 100644 internal/adminapi/bootstrap_test.go create mode 100644 internal/adminapi/model_page.go create mode 100644 internal/adminapi/model_page_test.go create mode 100644 internal/adminapi/usage_test.go create mode 100644 internal/billing/stripe_preflight.go create mode 100644 internal/billing/stripe_preflight_test.go create mode 100644 internal/controlplane/api_key_test.go create mode 100644 internal/controlplane/mail_operations_integration_test.go create mode 100644 internal/controlplane/usage_test.go create mode 100644 internal/providerhealth/prober.go create mode 100644 internal/providerhealth/prober_test.go create mode 100644 internal/providerhealth/redis_history.go (limited to 'internal') diff --git a/internal/adminapi/api.go b/internal/adminapi/api.go index ff87ac9..739cd87 100644 --- a/internal/adminapi/api.go +++ b/internal/adminapi/api.go @@ -44,7 +44,7 @@ type API struct { webauthn *webauthn.WebAuthn mailEnabled bool inferencePublicURL string - defaultLowBalance int64 + billingPreferences controlplane.BillingPreferenceDefaults catalog *catalog.Catalog health *providerhealth.Tracker } @@ -67,22 +67,22 @@ func (w *auditWriter) Write(body []byte) (int, error) { } type Options struct { - Store *controlplane.Store - Manager *controlplane.Manager - Billing *billing.Service - Token string - Logger *slog.Logger - Prefix string - RegistrationEnabled bool - SessionTTL time.Duration - Currency string - PublicURL string - WebAuthn *webauthn.WebAuthn - MailEnabled bool - InferencePublicURL string - DefaultLowBalanceMicros int64 - Catalog *catalog.Catalog - ProviderHealth *providerhealth.Tracker + Store *controlplane.Store + Manager *controlplane.Manager + Billing *billing.Service + Token string + Logger *slog.Logger + Prefix string + RegistrationEnabled bool + SessionTTL time.Duration + Currency string + PublicURL string + WebAuthn *webauthn.WebAuthn + MailEnabled bool + InferencePublicURL string + BillingPreferenceDefaults controlplane.BillingPreferenceDefaults + Catalog *catalog.Catalog + ProviderHealth *providerhealth.Tracker } func New(options Options) *API { @@ -103,7 +103,7 @@ func New(options Options) *API { logger: options.Logger, prefix: prefix, registrationEnabled: options.RegistrationEnabled, sessionTTL: options.SessionTTL, currency: options.Currency, publicURL: strings.TrimRight(options.PublicURL, "/") + "/", webauthn: options.WebAuthn, mailEnabled: options.MailEnabled, inferencePublicURL: strings.TrimRight(options.InferencePublicURL, "/"), - defaultLowBalance: options.DefaultLowBalanceMicros, catalog: options.Catalog, health: options.ProviderHealth} + billingPreferences: options.BillingPreferenceDefaults, catalog: options.Catalog, health: options.ProviderHealth} } func (a *API) Handler() http.Handler { @@ -116,6 +116,7 @@ func (a *API) Handler() http.Handler { } http.Redirect(w, r, target, http.StatusTemporaryRedirect) }) + mux.HandleFunc("GET "+a.prefix+"/models/{id...}", a.public(a.publicModelPage)) mux.Handle(a.prefix+"/", http.StripPrefix(a.prefix, adminui.Handler())) mux.HandleFunc("GET "+apiPrefix+"/public/models", a.public(a.publicModels)) mux.HandleFunc("GET "+apiPrefix+"/public/models/{id...}", a.public(a.publicModel)) @@ -158,6 +159,9 @@ func (a *API) Handler() http.Handler { mux.HandleFunc("POST "+apiPrefix+"/projects", a.withAuth("projects.write", a.createProject)) mux.HandleFunc("GET "+apiPrefix+"/keys", a.withAuth("keys.read", a.listKeys)) mux.HandleFunc("POST "+apiPrefix+"/keys", a.withAuth("keys.write", a.createKey)) + mux.HandleFunc("POST "+apiPrefix+"/keys/{id}/disable", a.withAuth("keys.write", a.disableKey)) + mux.HandleFunc("POST "+apiPrefix+"/keys/{id}/enable", a.withAuth("keys.write", a.enableKey)) + mux.HandleFunc("POST "+apiPrefix+"/keys/{id}/rotate", a.withAuth("keys.write", a.rotateKey)) mux.HandleFunc("POST "+apiPrefix+"/keys/{id}/revoke", a.withAuth("keys.write", a.revokeKey)) mux.HandleFunc("GET "+apiPrefix+"/providers", a.withAuth("platform.read", a.listProviders)) mux.HandleFunc("POST "+apiPrefix+"/providers", a.withAuth("platform.write", a.createProvider)) @@ -175,6 +179,7 @@ func (a *API) Handler() http.Handler { mux.HandleFunc("GET "+apiPrefix+"/billing/orders", a.withAuth("billing.read", a.listTopUpOrders)) mux.HandleFunc("GET "+apiPrefix+"/billing/orders/{id}", a.withAuth("billing.read", a.getTopUpOrder)) mux.HandleFunc("POST "+apiPrefix+"/billing/adjustments", a.withAuth("billing.adjust", a.adjustBalance)) + mux.HandleFunc("POST "+apiPrefix+"/billing/reservations/{request_id}/release", a.withAuth("billing.adjust", a.releaseUnmeteredReservation)) mux.HandleFunc("POST "+apiPrefix+"/billing/checkout-sessions", a.withAuth("billing.topup", a.createCheckoutSession)) mux.HandleFunc("POST "+apiPrefix+"/billing/portal-sessions", a.withAuth("billing.topup", a.createPortalSession)) mux.HandleFunc("GET "+apiPrefix+"/billing/auto-topup", a.withAuth("billing.read", a.getAutoTopUp)) @@ -287,6 +292,14 @@ func (a *API) actor(r *http.Request) controlplane.ConsoleActor { return actor } +func billingResolutionActor(actor controlplane.ConsoleActor, actorType string) billing.ResolutionActor { + actorID := actor.ID + if actor.Bootstrap { + actorID = "bootstrap" + } + return billing.ResolutionActor{ID: actorID, Type: actorType} +} + func (a *API) writeAudit(r *http.Request, actor controlplane.ConsoleActor, action string, status int) { if a.store == nil { return @@ -816,6 +829,10 @@ func (a *API) listSessions(w http.ResponseWriter, r *http.Request) { func (a *API) revokeSession(w http.ResponseWriter, r *http.Request) { actor := a.actor(r) + if actor.ID == "" { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "bootstrap_account", Message: "Bootstrap access has no device sessions"}, requestID(r)) + return + } current := sessionCookie(r) sessions, err := a.store.ListDeviceSessions(r.Context(), actor.ID, current) if err != nil { @@ -840,7 +857,12 @@ func (a *API) revokeSession(w http.ResponseWriter, r *http.Request) { } func (a *API) revokeOtherSessions(w http.ResponseWriter, r *http.Request) { - if err := a.store.RevokeOtherDeviceSessions(r.Context(), a.actor(r).ID, sessionCookie(r)); err != nil { + actor := a.actor(r) + if actor.ID == "" { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "bootstrap_account", Message: "Bootstrap access has no device sessions"}, requestID(r)) + return + } + if err := a.store.RevokeOtherDeviceSessions(r.Context(), actor.ID, sessionCookie(r)); err != nil { a.databaseError(w, r, err) return } @@ -966,6 +988,7 @@ func (a *API) developerConfig(w http.ResponseWriter, _ *http.Request) { "endpoints": map[string]string{ "chat_completions": "/v1/chat/completions", "responses": "/v1/responses", + "embeddings": "/v1/embeddings", "messages": "/anthropic/v1/messages", "models": "/v1/models", }, @@ -1002,18 +1025,15 @@ func (a *API) publicModels(w http.ResponseWriter, r *http.Request) { func (a *API) publicModel(w http.ResponseWriter, r *http.Request) { wanted := strings.Trim(strings.TrimSpace(r.PathValue("id")), "/") - result, err := a.store.ListPublicModels(r.Context()) + item, found, err := a.findPublicModel(r.Context(), wanted) if err != nil { a.databaseError(w, r, err) return } - a.addPublicModelHealth(result) - for _, item := range result { - if item.PublicID == wanted { - w.Header().Set("Cache-Control", "public, max-age=30, stale-while-revalidate=120") - writeJSON(w, item) - return - } + if found { + w.Header().Set("Cache-Control", "public, max-age=30, stale-while-revalidate=120") + writeJSON(w, item) + return } apierror.Write(w, apierror.Error{Status: http.StatusNotFound, Type: "model_not_found", Message: "Model not found"}, requestID(r)) } @@ -1105,8 +1125,11 @@ func (a *API) addDeveloperModelHealth(ctx context.Context, models []controlplane items = append(items, controlplane.DeveloperProviderHealth{Slug: route.Provider.EffectiveSlug(), Name: provider.Name, Protocol: string(route.Provider.Protocol), WireAPI: key.WireAPI, State: state, Attempts: status.Attempts, RecentSamples: status.RecentSamples, AvailabilityPercent: status.AvailabilityPercent, - HeaderLatencyEWMA: status.HeaderLatencyEWMA, ConsecutiveFailures: status.ConsecutiveFailures, - LastObservedAt: status.LastObservedAt, CircuitOpenUntil: status.CircuitOpenUntil}) + HeaderLatencyEWMA: status.HeaderLatencyEWMA, TTFTSamples: status.TTFTSamples, TTFTEWMA: status.TTFTEWMA, + SharedAttempts: status.SharedAttempts, SharedTTFTSamples: status.SharedTTFTSamples, + ConsecutiveFailures: status.ConsecutiveFailures, + ActiveProbes: status.ActiveProbes, LastObservedAt: status.LastObservedAt, + LastProbeAt: status.LastProbeAt, CircuitOpenUntil: status.CircuitOpenUntil}) } sort.Slice(items, func(i, j int) bool { iOpen := items[i].State == "open" @@ -1135,7 +1158,7 @@ func (a *API) addDeveloperModelHealth(ctx context.Context, models []controlplane } func (a *API) developerPreferences(w http.ResponseWriter, r *http.Request) { - result, err := a.store.GetTenantPreferences(r.Context(), a.preferenceTenantID(r, ""), a.defaultLowBalance) + result, err := a.store.GetTenantPreferences(r.Context(), a.preferenceTenantID(r, ""), a.billingPreferences) if err != nil { a.databaseError(w, r, err) return @@ -1149,7 +1172,7 @@ func (a *API) updateDeveloperPreferences(w http.ResponseWriter, r *http.Request) return } input.TenantID = a.preferenceTenantID(r, input.TenantID) - result, err := a.store.SetDeveloperPreferences(r.Context(), input) + result, err := a.store.SetDeveloperPreferences(r.Context(), input, a.billingPreferences) if err != nil { a.mutationError(w, r, err) return @@ -1163,7 +1186,7 @@ func (a *API) updateBillingPreferences(w http.ResponseWriter, r *http.Request) { return } input.TenantID = a.preferenceTenantID(r, input.TenantID) - result, err := a.store.SetBillingPreferences(r.Context(), input, a.defaultLowBalance) + result, err := a.store.SetBillingPreferences(r.Context(), input, a.billingPreferences) if err != nil { a.mutationError(w, r, err) return @@ -1273,6 +1296,33 @@ func (a *API) adjustBalance(w http.ResponseWriter, r *http.Request) { writeStatusJSON(w, http.StatusCreated, result) } +func (a *API) releaseUnmeteredReservation(w http.ResponseWriter, r *http.Request) { + var input billing.ReleaseReservationInput + if !decodeBody(w, r, &input) { + return + } + tenantID := strings.TrimSpace(r.URL.Query().Get("tenant_id")) + actor := a.actor(r) + if actor.TenantID != "" { + tenantID = actor.TenantID + } + if tenantID == "" { + a.scopeError(w, r) + return + } + actorType := "console_user" + if actor.Bootstrap { + actorType = "bootstrap" + } + result, err := a.billing.ReleaseUnmeteredReservation(r.Context(), tenantID, r.PathValue("request_id"), input, + billingResolutionActor(actor, actorType)) + if err != nil { + a.billingError(w, r, err) + return + } + writeJSON(w, result) +} + func (a *API) createCheckoutSession(w http.ResponseWriter, r *http.Request) { var input billing.CheckoutInput if !decodeBody(w, r, &input) { @@ -1377,7 +1427,7 @@ func (a *API) resolveMissingTopUp(w http.ResponseWriter, r *http.Request) { actorType = "bootstrap" } result, err := a.billing.ResolveMissingTopUp(r.Context(), tenantID, r.PathValue("id"), input, - billing.ResolutionActor{ID: actor.ID, Type: actorType}) + billingResolutionActor(actor, actorType)) if err != nil { a.billingError(w, r, err) return @@ -1404,7 +1454,7 @@ func (a *API) reverseMissingTopUpCredit(w http.ResponseWriter, r *http.Request) actorType = "bootstrap" } result, err := a.billing.ReverseMissingTopUpCredit(r.Context(), tenantID, r.PathValue("id"), input, - billing.ResolutionActor{ID: actor.ID, Type: actorType}) + billingResolutionActor(actor, actorType)) if err != nil { a.billingError(w, r, err) return @@ -1568,10 +1618,62 @@ func (a *API) createKey(w http.ResponseWriter, r *http.Request) { a.mutationError(w, r, err) return } - if !a.changed(w, r, generation, "api_key", result.ID) { + syncStatus := a.afterSecretMutation(r, generation, result.ID) + writeStatusJSON(w, http.StatusCreated, struct { + controlplane.CreatedAPIKey + RuntimeSyncStatus string `json:"runtime_sync_status"` + }{CreatedAPIKey: result, RuntimeSyncStatus: syncStatus}) +} + +func (a *API) disableKey(w http.ResponseWriter, r *http.Request) { + a.setKeyStatus(w, r, "disabled", a.store.DisableAPIKey) +} + +func (a *API) enableKey(w http.ResponseWriter, r *http.Request) { + a.setKeyStatus(w, r, "active", a.store.EnableAPIKey) +} + +func (a *API) setKeyStatus(w http.ResponseWriter, r *http.Request, status string, update func(context.Context, string) (int64, error)) { + id := r.PathValue("id") + if err := a.requireResourceTenant(r, "api_key", id); err != nil { + a.scopeError(w, r) return } - writeStatusJSON(w, http.StatusCreated, result) + generation, err := update(r.Context(), id) + if err != nil { + a.mutationError(w, r, err) + return + } + if !a.changed(w, r, generation, "api_key", id) { + return + } + writeJSON(w, map[string]any{"status": status}) +} + +func (a *API) rotateKey(w http.ResponseWriter, r *http.Request) { + id := r.PathValue("id") + if err := a.requireResourceTenant(r, "api_key", id); err != nil { + a.scopeError(w, r) + return + } + result, generation, err := a.store.RotateAPIKey(r.Context(), id) + if err != nil { + a.mutationError(w, r, err) + return + } + syncStatus := a.afterSecretMutation(r, generation, result.ID) + writeStatusJSON(w, http.StatusCreated, struct { + controlplane.CreatedAPIKey + RuntimeSyncStatus string `json:"runtime_sync_status"` + }{CreatedAPIKey: result, RuntimeSyncStatus: syncStatus}) +} + +func (a *API) afterSecretMutation(r *http.Request, generation int64, id string) string { + if err := a.manager.AfterMutation(r.Context(), generation, "api_key", id); err != nil { + a.logger.Error("admin_control_plane_sync_failed", "resource", "api_key", "id", id, "error", err) + return "pending" + } + return "applied" } func (a *API) revokeKey(w http.ResponseWriter, r *http.Request) { @@ -1712,9 +1814,16 @@ func (a *API) listUsage(w http.ResponseWriter, r *http.Request) { } result, err := a.store.ListUsage(r.Context(), query) if err != nil { + if errors.Is(err, controlplane.ErrInvalidUsageCursor) { + a.mutationError(w, r, err) + return + } a.databaseError(w, r, err) return } + if a.actor(r).TenantID != "" { + redactTenantUsagePage(&result) + } writeJSON(w, result) } @@ -1743,6 +1852,9 @@ func (a *API) usageAnalytics(w http.ResponseWriter, r *http.Request) { a.databaseError(w, r, err) return } + if a.actor(r).TenantID != "" { + redactTenantUsageAnalytics(&result) + } writeJSON(w, result) } @@ -1752,12 +1864,16 @@ func (a *API) usageQuery(r *http.Request) (controlplane.UsageQuery, error) { if tenantID == "" { tenantID = strings.TrimSpace(r.URL.Query().Get("tenant_id")) } + provider := "" + if actor.TenantID == "" { + provider = strings.ToLower(strings.TrimSpace(r.URL.Query().Get("provider"))) + } query := controlplane.UsageQuery{ TenantID: tenantID, ProjectID: strings.TrimSpace(r.URL.Query().Get("project_id")), KeyID: strings.TrimSpace(r.URL.Query().Get("key_id")), Model: strings.TrimSpace(r.URL.Query().Get("model")), - Provider: strings.ToLower(strings.TrimSpace(r.URL.Query().Get("provider"))), Protocol: strings.TrimSpace(r.URL.Query().Get("protocol")), + Provider: provider, Protocol: strings.TrimSpace(r.URL.Query().Get("protocol")), ErrorType: strings.TrimSpace(r.URL.Query().Get("error_type")), RequestID: strings.TrimSpace(r.URL.Query().Get("request_id")), - Status: strings.TrimSpace(r.URL.Query().Get("status")), Limit: 200, + Status: strings.TrimSpace(r.URL.Query().Get("status")), Limit: 200, Cursor: strings.TrimSpace(r.URL.Query().Get("cursor")), } if raw := strings.TrimSpace(r.URL.Query().Get("limit")); raw != "" { limit, err := strconv.Atoi(raw) @@ -1769,12 +1885,15 @@ func (a *API) usageQuery(r *http.Request) (controlplane.UsageQuery, error) { if query.Status != "" && query.Status != "success" && query.Status != "error" { return controlplane.UsageQuery{}, errors.New("usage status must be success or error") } - if query.Protocol != "" && query.Protocol != string(domain.ProtocolOpenAI) && query.Protocol != string(domain.ProtocolOpenAIResponses) && query.Protocol != string(domain.ProtocolAnthropic) { + if query.Protocol != "" && query.Protocol != string(domain.ProtocolOpenAI) && query.Protocol != string(domain.ProtocolOpenAIResponses) && query.Protocol != string(domain.ProtocolOpenAIEmbeddings) && query.Protocol != string(domain.ProtocolAnthropic) { return controlplane.UsageQuery{}, errors.New("usage protocol is invalid") } if len(query.Provider) > 64 || len(query.ErrorType) > 128 { return controlplane.UsageQuery{}, errors.New("usage provider or error type is too long") } + if len(query.Cursor) > 2048 { + return controlplane.UsageQuery{}, errors.New("usage cursor is too long") + } if raw := strings.TrimSpace(r.URL.Query().Get("stream")); raw != "" { value, err := strconv.ParseBool(raw) if err != nil { @@ -1795,6 +1914,21 @@ func (a *API) usageQuery(r *http.Request) (controlplane.UsageQuery, error) { return query, nil } +func redactTenantUsagePage(page *controlplane.UsagePage) { + for index := range page.Data { + page.Data[index].ProviderID = "" + page.Data[index].ProviderName = "" + page.Data[index].UpstreamModel = "" + } +} + +func redactTenantUsageAnalytics(analytics *controlplane.UsageAnalytics) { + analytics.Providers = []controlplane.UsageProviderAnalytics{} + for index := range analytics.Models { + analytics.Models[index].ProviderCount = 0 + } +} + func parseUsageTime(raw string, endOfDay bool) (time.Time, error) { raw = strings.TrimSpace(raw) if raw == "" { @@ -2009,6 +2143,10 @@ func (a *API) billingError(w http.ResponseWriter, r *http.Request, err error) { typeName = "billing_profile_sync_failed" message = "Billing details were saved, but Stripe synchronization failed" a.logger.Error("billing_profile_sync_failed", "error", err) + case errors.Is(err, billing.ErrReservationNotReleasable): + status = http.StatusConflict + typeName = "reservation_not_releasable" + message = err.Error() default: a.logger.Error("admin_billing_error", "error", err) } diff --git a/internal/adminapi/bootstrap_test.go b/internal/adminapi/bootstrap_test.go new file mode 100644 index 0000000..b0f0e90 --- /dev/null +++ b/internal/adminapi/bootstrap_test.go @@ -0,0 +1,57 @@ +package adminapi + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "aigw/internal/controlplane" +) + +func TestBootstrapActorHasNoDatabaseUserID(t *testing.T) { + handler := New(Options{Token: "bootstrap-secret", Prefix: "/admin"}).Handler() + + request := httptest.NewRequest(http.MethodGet, "/admin/api/me", nil) + request.Header.Set("Authorization", "Bearer bootstrap-secret") + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + if response.Code != http.StatusOK { + t.Fatalf("bootstrap me status = %d, body = %s", response.Code, response.Body.String()) + } + var payload struct { + Actor controlplane.ConsoleActor `json:"actor"` + } + if err := json.Unmarshal(response.Body.Bytes(), &payload); err != nil { + t.Fatal(err) + } + if !payload.Actor.Bootstrap || payload.Actor.ID != "" || payload.Actor.Role != controlplane.RolePlatformAdmin { + t.Fatalf("unexpected bootstrap actor: %+v", payload.Actor) + } +} + +func TestBootstrapSecurityEndpointsDoNotQueryUserUUID(t *testing.T) { + handler := New(Options{Token: "bootstrap-secret", Prefix: "/admin"}).Handler() + for _, test := range []struct { + path string + wantStatus int + }{ + {path: "/admin/api/auth/mfa", wantStatus: http.StatusBadRequest}, + {path: "/admin/api/auth/sessions", wantStatus: http.StatusOK}, + } { + request := httptest.NewRequest(http.MethodGet, test.path, nil) + request.Header.Set("Authorization", "Bearer bootstrap-secret") + response := httptest.NewRecorder() + handler.ServeHTTP(response, request) + if response.Code != test.wantStatus { + t.Fatalf("%s status = %d, want %d; body = %s", test.path, response.Code, test.wantStatus, response.Body.String()) + } + } +} + +func TestBootstrapBillingResolutionActorUsesTextEvidenceID(t *testing.T) { + actor := billingResolutionActor(controlplane.ConsoleActor{Bootstrap: true, Role: controlplane.RolePlatformAdmin}, "bootstrap") + if actor.ID != "bootstrap" || actor.Type != "bootstrap" { + t.Fatalf("unexpected resolution actor: %+v", actor) + } +} diff --git a/internal/adminapi/model_page.go b/internal/adminapi/model_page.go new file mode 100644 index 0000000..5a94564 --- /dev/null +++ b/internal/adminapi/model_page.go @@ -0,0 +1,276 @@ +package adminapi + +import ( + "context" + "encoding/json" + "fmt" + "html/template" + "net/http" + "net/url" + "strconv" + "strings" + "time" + + "aigw/internal/controlplane" +) + +type publicModelExample struct { + Name string + Endpoint string + Code string +} + +type publicModelPageData struct { + Model controlplane.PublicModel + Title string + Description string + CanonicalURL string + CatalogURL string + SignInURL string + RegistrationURL string + RegistrationEnabled bool + CSSURL string + HealthLabel string + HealthClass string + InputPrice string + OutputPrice string + CacheReadPrice string + ContextWindow string + MaxOutputTokens string + ReleasedAt string + Tags []string + Examples []publicModelExample +} + +var publicModelTemplate = template.Must(template.New("public-model").Parse(` + + + + + + + + + + + {{.Title}} + + + + +
+ AAIGWMODEL CATALOG + +
+
+ Back to models +
+
{{if .Model.OwnedBy}}{{.Model.OwnedBy}}{{else}}INDEPENDENT{{end}}

{{.Model.DisplayName}}

{{.Model.PublicID}}
+ {{.HealthLabel}} +
+

{{.Description}}

+
{{range .Tags}}{{.}}{{end}}
+ +
+
+

Model details

+
+
Input price
{{.InputPrice}} / 1M tokens
+
Output price
{{.OutputPrice}} / 1M tokens
+
Cached input
{{.CacheReadPrice}} / 1M tokens
+
Context window
{{.ContextWindow}} tokens
+
Max output
{{.MaxOutputTokens}} tokens
+
Regions
{{if .Model.Regions}}{{range $index,$region := .Model.Regions}}{{if $index}}, {{end}}{{$region}}{{end}}{{else}}Global{{end}}
+
Released
{{.ReleasedAt}}
+
Lifecycle
{{.Model.Lifecycle}}
+
+
+ +
+ +
+
SUPPORTED APIS

Code examples

+
{{range .Examples}}
{{.Name}}{{.Endpoint}}
{{.Code}}
{{end}}
+
+
+ +`)) + +func (a *API) findPublicModel(ctx context.Context, wanted string) (controlplane.PublicModel, bool, error) { + models, err := a.store.ListPublicModels(ctx) + if err != nil { + return controlplane.PublicModel{}, false, err + } + a.addPublicModelHealth(models) + for _, item := range models { + if item.PublicID == wanted { + return item, true, nil + } + } + return controlplane.PublicModel{}, false, nil +} + +func (a *API) publicModelPage(w http.ResponseWriter, r *http.Request) { + wanted := strings.Trim(strings.TrimSpace(r.PathValue("id")), "/") + model, found, err := a.findPublicModel(r.Context(), wanted) + if err != nil { + a.logger.Error("public_model_page_failed", "error", err) + http.Error(w, "Model catalog unavailable", http.StatusServiceUnavailable) + return + } + if !found { + http.NotFound(w, r) + return + } + data := buildPublicModelPageData(model, a.publicURL, a.inferencePublicURL, a.prefix, a.registrationEnabled) + w.Header().Set("Content-Type", "text/html; charset=utf-8") + w.Header().Set("Cache-Control", "public, max-age=30, stale-while-revalidate=120") + if err := publicModelTemplate.Execute(w, data); err != nil { + a.logger.Error("render_public_model_page_failed", "model", model.PublicID, "error", err) + } +} + +func buildPublicModelPageData(model controlplane.PublicModel, publicURL, inferenceBase, prefix string, registrationEnabled bool) publicModelPageData { + canonicalBase := strings.TrimRight(publicURL, "/") + pathBase := strings.TrimRight(prefix, "/") + if pathBase == "" { + pathBase = "/admin" + } + if canonicalBase == "" { + canonicalBase = pathBase + } + pathID := escapeModelPath(model.PublicID) + canonical := canonicalBase + "/models/" + pathID + description := strings.TrimSpace(model.Description) + if description == "" { + description = fmt.Sprintf("Use %s through the AIGW unified API.", model.PublicID) + } + tags := append([]string(nil), model.Capabilities...) + tags = append(tags, model.InputModalities...) + for _, wireAPI := range model.SupportedWireAPIs { + tags = append(tags, protocolLabel(wireAPI)) + } + healthLabel, healthClass := publicHealthLabel(model) + return publicModelPageData{ + Model: model, Title: model.DisplayName + " API, pricing, and context | AIGW", Description: description, + CanonicalURL: canonical, CatalogURL: pathBase + "/models", SignInURL: pathBase + "/", + RegistrationURL: pathBase + "/?auth=register&model=" + url.QueryEscape(model.PublicID), RegistrationEnabled: registrationEnabled, + CSSURL: pathBase + "/models.css", HealthLabel: healthLabel, HealthClass: healthClass, + InputPrice: formatMicros(model.InputPriceMicrosPerMillion, model.PriceCurrency), + OutputPrice: formatMicros(model.OutputPriceMicrosPerMillion, model.PriceCurrency), + CacheReadPrice: formatMicros(model.CacheReadPriceMicrosPerMillion, model.PriceCurrency), + ContextWindow: formatInteger(model.ContextWindow), MaxOutputTokens: formatInteger(model.MaxOutputTokens), + ReleasedAt: formatPublicDate(model.ReleasedAt), Tags: uniquePageStrings(tags), + Examples: publicModelExamples(model, inferenceBase), + } +} + +func publicModelExamples(model controlplane.PublicModel, inferenceBase string) []publicModelExample { + base := strings.TrimRight(inferenceBase, "/") + if base == "" { + base = "https://api.example.com" + } + modelJSON, _ := json.Marshal(model.PublicID) + result := make([]publicModelExample, 0, len(model.SupportedWireAPIs)) + for _, wireAPI := range model.SupportedWireAPIs { + var endpoint, code string + switch wireAPI { + case "chat_completions": + endpoint = "/v1/chat/completions" + code = fmt.Sprintf("curl %s%s \\\n -H 'Authorization: Bearer $AIGW_API_KEY' \\\n -H 'Content-Type: application/json' \\\n -d '{\"model\":%s,\"messages\":[{\"role\":\"user\",\"content\":\"Hello\"}]}'", base, endpoint, modelJSON) + case "responses": + endpoint = "/v1/responses" + code = fmt.Sprintf("curl %s%s \\\n -H 'Authorization: Bearer $AIGW_API_KEY' \\\n -H 'Content-Type: application/json' \\\n -d '{\"model\":%s,\"input\":\"Hello\"}'", base, endpoint, modelJSON) + case "embeddings": + endpoint = "/v1/embeddings" + code = fmt.Sprintf("curl %s%s \\\n -H 'Authorization: Bearer $AIGW_API_KEY' \\\n -H 'Content-Type: application/json' \\\n -d '{\"model\":%s,\"input\":[\"Text to embed\"]}'", base, endpoint, modelJSON) + case "messages": + endpoint = "/anthropic/v1/messages" + code = fmt.Sprintf("curl %s%s \\\n -H 'x-api-key: $AIGW_API_KEY' \\\n -H 'anthropic-version: 2023-06-01' \\\n -H 'Content-Type: application/json' \\\n -d '{\"model\":%s,\"max_tokens\":256,\"messages\":[{\"role\":\"user\",\"content\":\"Hello\"}]}'", base, endpoint, modelJSON) + default: + continue + } + result = append(result, publicModelExample{Name: protocolLabel(wireAPI), Endpoint: endpoint, Code: code}) + } + return result +} + +func escapeModelPath(publicID string) string { + parts := strings.Split(strings.Trim(publicID, "/"), "/") + for index := range parts { + parts[index] = url.PathEscape(parts[index]) + } + return strings.Join(parts, "/") +} + +func publicHealthLabel(model controlplane.PublicModel) (string, string) { + switch model.HealthStatus { + case "unavailable": + return "Unavailable", "unavailable" + case "degraded": + return fmt.Sprintf("%d/%d routes", model.AvailableProviderCount, model.ProviderCount), "degraded" + default: + return "Available", "available" + } +} + +func protocolLabel(value string) string { + switch value { + case "chat_completions": + return "Chat Completions" + case "responses": + return "Responses" + case "embeddings": + return "Embeddings" + case "messages": + return "Anthropic Messages" + default: + return value + } +} + +func formatMicros(value int64, currency string) string { + amount := strconv.FormatFloat(float64(value)/1_000_000, 'f', 6, 64) + amount = strings.TrimRight(strings.TrimRight(amount, "0"), ".") + if amount == "" { + amount = "0" + } + return strings.ToUpper(currency) + " " + amount +} + +func formatInteger(value int64) string { + raw := strconv.FormatInt(value, 10) + for index := len(raw) - 3; index > 0; index -= 3 { + raw = raw[:index] + "," + raw[index:] + } + return raw +} + +func formatPublicDate(value *time.Time) string { + if value == nil { + return "Not published" + } + return value.UTC().Format("2 Jan 2006") +} + +func uniquePageStrings(values []string) []string { + seen := make(map[string]struct{}, len(values)) + result := make([]string, 0, len(values)) + for _, value := range values { + value = strings.TrimSpace(value) + if value == "" { + continue + } + if _, exists := seen[value]; exists { + continue + } + seen[value] = struct{}{} + result = append(result, value) + } + return result +} diff --git a/internal/adminapi/model_page_test.go b/internal/adminapi/model_page_test.go new file mode 100644 index 0000000..eac6d7d --- /dev/null +++ b/internal/adminapi/model_page_test.go @@ -0,0 +1,44 @@ +package adminapi + +import ( + "strings" + "testing" + "time" + + "aigw/internal/controlplane" +) + +func TestBuildPublicModelPageDataUsesSupportedProtocolsWithoutInternalRouting(t *testing.T) { + released := time.Date(2026, time.August, 1, 0, 0, 0, 0, time.UTC) + model := controlplane.PublicModel{ + PublicID: "openai/text-embedding-3-small", DisplayName: "Text Embedding 3 Small", + Description: "Embeddings for search.", OwnedBy: "OpenAI", InputModalities: []string{"text"}, + OutputModalities: []string{"embedding"}, Capabilities: []string{"embeddings"}, Regions: []string{"global"}, + Lifecycle: "active", ReleasedAt: &released, PriceCurrency: "usd", InputPriceMicrosPerMillion: 20_000, + ContextWindow: 8192, SupportedWireAPIs: []string{"embeddings"}, ProviderCount: 2, + AvailableProviderCount: 1, HealthStatus: "degraded", + } + data := buildPublicModelPageData(model, "https://console.example.test/admin/", "https://api.example.test/", "/admin", true) + if data.CanonicalURL != "https://console.example.test/admin/models/openai/text-embedding-3-small" { + t.Fatalf("canonical=%q", data.CanonicalURL) + } + if data.CSSURL != "/admin/models.css" || data.CatalogURL != "/admin/models" { + t.Fatalf("cross-origin static URL: css=%q catalog=%q", data.CSSURL, data.CatalogURL) + } + if len(data.Examples) != 1 || data.Examples[0].Endpoint != "/v1/embeddings" { + t.Fatalf("examples=%+v", data.Examples) + } + if !strings.Contains(data.Examples[0].Code, `"model":"openai/text-embedding-3-small"`) || + strings.Contains(data.Examples[0].Code, "provider") || strings.Contains(data.Examples[0].Code, "upstream") { + t.Fatalf("unexpected public example: %s", data.Examples[0].Code) + } + if data.HealthLabel != "1/2 routes" || data.InputPrice != "USD 0.02" || data.ContextWindow != "8,192" { + t.Fatalf("unexpected page projection: %+v", data) + } +} + +func TestEscapeModelPathPreservesHierarchyAndEscapesSegments(t *testing.T) { + if got := escapeModelPath("owner/model name"); got != "owner/model%20name" { + t.Fatalf("escapeModelPath=%q", got) + } +} diff --git a/internal/adminapi/usage_test.go b/internal/adminapi/usage_test.go new file mode 100644 index 0000000..e6192d6 --- /dev/null +++ b/internal/adminapi/usage_test.go @@ -0,0 +1,26 @@ +package adminapi + +import ( + "testing" + + "aigw/internal/controlplane" +) + +func TestTenantUsageRedactionRemovesRouteInternals(t *testing.T) { + page := controlplane.UsagePage{Data: []controlplane.UsageRecord{{ + ProviderID: "provider-uuid", ProviderName: "Internal Provider", UpstreamModel: "vendor/model-v2", + }}} + redactTenantUsagePage(&page) + if page.Data[0].ProviderID != "" || page.Data[0].ProviderName != "" || page.Data[0].UpstreamModel != "" { + t.Fatalf("tenant usage leaked route internals: %+v", page.Data[0]) + } + + analytics := controlplane.UsageAnalytics{ + Models: []controlplane.UsageModelAnalytics{{PublicModel: "public/model", ProviderCount: 3}}, + Providers: []controlplane.UsageProviderAnalytics{{ProviderID: "provider-uuid", ProviderName: "Internal Provider"}}, + } + redactTenantUsageAnalytics(&analytics) + if len(analytics.Providers) != 0 || analytics.Models[0].ProviderCount != 0 { + t.Fatalf("tenant analytics leaked provider topology: %+v", analytics) + } +} diff --git a/internal/adminui/assets/app.js b/internal/adminui/assets/app.js index e30a0c8..f2ed886 100644 --- a/internal/adminui/assets/app.js +++ b/internal/adminui/assets/app.js @@ -1,7 +1,7 @@ const state = { token: '', csrf: '', actor: {}, permissions: new Set(), overview: {}, tenants: [], projects: [], keys: [], providers: [], models: [], billingAccounts: [], ledger: [], - usage: [], usageSummary: [], usageDaily: [], usageAnalytics: {models:[],providers:[]}, limits: [], users: [], audit: [], orders: [], refunds: [], disputes: [], invoices: [], sessions: [], + usage: [], usagePaging: {cursor:'',nextCursor:'',history:[]}, usageSummary: [], usageDaily: [], usageAnalytics: {models:[],providers:[],keys:[]}, limits: [], users: [], audit: [], orders: [], refunds: [], disputes: [], invoices: [], sessions: [], developerConfig: {base_url:'',endpoints:{}}, developerModels: [], preferences: {}, autoTopUp: {}, billingProfile: {}, mfa: {totp_enabled:false,passkeys:[]}, pendingMFA: null, authConfig: {}, playgroundKey: '', playgroundController: null, detailModel: null, detailUsage: null @@ -32,7 +32,13 @@ function setConnected(connected) { $('#connection-state').textContent = state.actor.role?.replaceAll('_', ' ') || 'connected'; $('#actor-label').textContent = state.actor.display_name || state.actor.email || 'Operator'; } -function formJSON(form) { return Object.fromEntries(new FormData(form).entries()); } +function formJSON(form) { + const data = new FormData(form); + // Password-manager username hints are intentionally decoys. They must not + // be sent to strict JSON endpoints that only accept the documented fields. + data.delete('username'); + return Object.fromEntries(data.entries()); +} function selectOptions(items, valueKey, labelKey, empty = 'Select…') { return `${items.map(item => ``).join('')}`; } function decimalToScaled(value, digits) { const match = String(value).trim().match(/^(-?)(\d+)(?:\.(\d+))?$/); if (!match) throw new Error('Enter a valid decimal amount'); @@ -102,7 +108,7 @@ async function loadAll(knownSession = null) { const results = await Promise.all([ permitted('tenants.read','/tenants'), permitted('projects.read','/projects'), permitted('keys.read','/keys'), permitted('platform.read','/providers'), permitted('platform.read','/models'), permitted('overview.read','/developer/config'), - permitted('overview.read','/developer/models'), can('preferences.read') ? api('/developer/preferences') : {}, permitted('usage.read',`/usage?${usageQuery}`), + permitted('overview.read','/developer/models'), can('preferences.read') ? api('/developer/preferences') : {}, permitted('usage.read',`/usage?${usageQuery}&limit=50`), permitted('usage.read','/usage/summary'), permitted('limits.read','/limits'), permitted('users.read','/users'), permitted('audit.read','/audit'), state.overview.billing_enabled ? permitted('billing.read','/billing/accounts') : [], state.overview.billing_enabled ? permitted('billing.read','/billing/ledger') : [], @@ -117,6 +123,7 @@ async function loadAll(knownSession = null) { state.overview.billing_enabled && state.actor.tenant_id && can('billing.read') ? api('/billing/profile') : {} ]); [state.tenants,state.projects,state.keys,state.providers,state.models,state.developerConfig,state.developerModels,state.preferences,state.usage,state.usageSummary,state.limits,state.users,state.audit,state.billingAccounts,state.ledger,state.mfa,state.sessions,state.orders,state.refunds,state.disputes,state.invoices,state.usageDaily,state.usageAnalytics,state.autoTopUp,state.billingProfile] = results; + const usagePage=state.usage;state.usage=Array.isArray(usagePage)?usagePage:(usagePage?.data||[]);state.usagePaging={cursor:'',nextCursor:usagePage?.next_cursor||'',history:[]}; renderAll(); setConnected(true); return true; } catch (error) { setConnected(false); if (error.status !== 401) toast(error.message, true); return false; } } @@ -148,8 +155,9 @@ function renderOverview() { function goTo(section) { const node=$(`.tab[data-section="${section}"]`); if(node){ node.click(); window.scrollTo({top:0,behavior:'smooth'}); } } function latestAvailableBalance() { return state.billingAccounts.find(item => !state.actor.tenant_id || item.tenant_id===state.actor.tenant_id)?.available_micros || 0; } function inferenceBaseURL() { return String(state.developerConfig.base_url||window.location.origin).replace(/\/$/,''); } -function developerEndpoint(name) { const defaults={chat_completions:'/v1/chat/completions',responses:'/v1/responses',messages:'/anthropic/v1/messages',models:'/v1/models'};const path=state.developerConfig.endpoints?.[name]||defaults[name]||'';return `${inferenceBaseURL()}${path}`; } +function developerEndpoint(name) { const defaults={chat_completions:'/v1/chat/completions',responses:'/v1/responses',embeddings:'/v1/embeddings',messages:'/anthropic/v1/messages',models:'/v1/models'};const path=state.developerConfig.endpoints?.[name]||defaults[name]||'';return `${inferenceBaseURL()}${path}`; } function modelIsUnavailable(model) { return model?.health_status==='unavailable'; } +function apiKeyDisplay(item) { return `${item?.key_prefix||''}${item?.key_suffix||''}`; } function developerModelOption(item) { const unavailable=modelIsUnavailable(item);return ``; } function providerHealthSummary(item) { const total=Number(item.provider_count||0);const available=Number(item.available_provider_count||0);const status=item.health_status||'online'; @@ -193,7 +201,7 @@ function renderDeveloperAccess() { const model=selectedDeveloperModel();$('#starter-model-label').textContent=model.public_id||'No model'; form.classList.toggle('hidden',!state.actor.tenant_id||!can('keys.write')); $('#starter-key-submit').disabled=!select.value||!model.public_id||modelIsUnavailable(model); - const rows=[['OpenAI SDK base',`${inferenceBaseURL()}/v1`],['Anthropic SDK base',`${inferenceBaseURL()}/anthropic`],['Chat Completions',developerEndpoint('chat_completions')],['Responses',developerEndpoint('responses')],['Anthropic Messages',developerEndpoint('messages')],['Models',developerEndpoint('models')]]; + const rows=[['OpenAI SDK base',`${inferenceBaseURL()}/v1`],['Anthropic SDK base',`${inferenceBaseURL()}/anthropic`],['Chat Completions',developerEndpoint('chat_completions')],['Responses',developerEndpoint('responses')],['Embeddings',developerEndpoint('embeddings')],['Anthropic Messages',developerEndpoint('messages')],['Models',developerEndpoint('models')]]; $('#endpoint-list').innerHTML=rows.map(([label,value],index)=>`
${esc(label)}${esc(value)}
`).join(''); $('#endpoint-list').dataset.values=JSON.stringify(rows.map(([,value])=>value)); $('#endpoint-env').textContent=`export AIGW_API_KEY="your-key"\nexport OPENAI_BASE_URL="${inferenceBaseURL()}/v1"\nexport ANTHROPIC_BASE_URL="${inferenceBaseURL()}/anthropic"`; @@ -208,22 +216,27 @@ function renderPreferences() { const canWriteDeveloper=can('developer.preferences.write'); defaultSelect.disabled=!canWriteDeveloper; fallbackSelect.disabled=!canWriteDeveloper; $('#low-balance-enabled').checked=preferences.low_balance_enabled!==false; $('#low-balance-threshold').value=scaledToDecimal(preferences.low_balance_threshold_micros||0,6); - const canWriteBilling=can('billing.preferences.write'); $('#low-balance-enabled').disabled=!canWriteBilling; $('#low-balance-threshold').disabled=!canWriteBilling; - $('#balance-alert-status').textContent=state.authConfig.email_delivery_enabled?'Verified billing members receive at most one low-balance alert per day.':'The preference is saved now and activates when SMTP delivery is configured.'; + $('#spend-anomaly-enabled').checked=preferences.spend_anomaly_enabled!==false; + $('#spend-anomaly-multiplier').value=Number(preferences.spend_anomaly_multiplier||3); + $('#spend-anomaly-minimum').value=scaledToDecimal(preferences.spend_anomaly_min_micros||0,6); + const canWriteBilling=can('billing.preferences.write');for(const id of ['low-balance-enabled','low-balance-threshold','spend-anomaly-enabled','spend-anomaly-multiplier','spend-anomaly-minimum'])$(`#${id}`).disabled=!canWriteBilling; + $('#balance-alert-status').textContent=state.authConfig.email_delivery_enabled?'Verified billing members receive at most one alert of each type per day. Anomaly spend compares today with the previous seven-day daily average.':'These preferences are saved now and activate when SMTP delivery is configured.'; } function selectedDeveloperModel() { return state.developerModels.find(item=>item.public_id===$('#quickstart-model')?.value) || state.developerModels.find(item=>!modelIsUnavailable(item)) || state.developerModels[0] || {}; } function syncQuickstartProtocols() { const model=selectedDeveloperModel(); const select=$('#quickstart-protocol'); const previous=select.value; - select.innerHTML=(model.supported_wire_apis||[]).map(apiName=>``).join('')||''; + select.innerHTML=(model.supported_wire_apis||[]).map(apiName=>``).join('')||''; if((model.supported_wire_apis||[]).includes(previous))select.value=previous; syncProviderSelect('#quickstart-provider',model,select.value); } function selectedPlaygroundModel() { return state.developerModels.find(item=>item.public_id===$('#playground-model')?.value) || state.developerModels.find(item=>!modelIsUnavailable(item)) || state.developerModels[0] || {}; } function syncPlaygroundProtocols() { const model=selectedPlaygroundModel(); const select=$('#playground-protocol'); const previous=select.value; - select.innerHTML=(model.supported_wire_apis||[]).map(apiName=>``).join('')||''; + select.innerHTML=(model.supported_wire_apis||[]).map(apiName=>``).join('')||''; if((model.supported_wire_apis||[]).includes(previous))select.value=previous; syncProviderSelect('#playground-provider',model,select.value); + const embeddings=select.value==='embeddings';const maxField=$('#playground-max-output-field'); + maxField.classList.toggle('hidden',embeddings);$('#playground-max-output').required=!embeddings; const max=Number(model.max_output_tokens||4096);$('#playground-max-output').max=String(max>0?max:4096); if(Number($('#playground-max-output').value)>max&&max>0)$('#playground-max-output').value=String(max); } @@ -233,24 +246,31 @@ function syncProviderSelect(selector,model,wire) { if(providers.some(item=>item.slug===previous&&item.state!=='open'))select.value=previous; } function modelSelector(model,providerSelector) { const provider=$(providerSelector)?.value||'';const publicID=model.public_id||'model-id';return provider?`${publicID}:${provider}`:publicID; } -function quickstartEndpoint(wire) { return wire==='messages'?developerEndpoint('messages'):wire==='responses'?developerEndpoint('responses'):developerEndpoint('chat_completions'); } +function quickstartEndpoint(wire) { return wire==='messages'?developerEndpoint('messages'):wire==='responses'?developerEndpoint('responses'):wire==='embeddings'?developerEndpoint('embeddings'):developerEndpoint('chat_completions'); } function renderQuickstartCode() { const model=selectedDeveloperModel(); const selectedModel=modelSelector(model,'#quickstart-provider'); const wire=$('#quickstart-protocol').value||'chat_completions'; const language=$('#quickstart-language').value||'curl'; const endpoint=quickstartEndpoint(wire); const key='$AIGW_API_KEY'; - let code=''; + const base=String(state.developerConfig.base_url||window.location.origin).replace(/\/$/,'');let code=''; if(language==='curl') { const headers=wire==='messages'?`-H "x-api-key: ${key}"\n -H "anthropic-version: 2023-06-01"`:`-H "Authorization: Bearer ${key}"`; - const body=wire==='messages'?`{"model":"${selectedModel}","max_tokens":256,"messages":[{"role":"user","content":"Say hello in one sentence."}]}`:wire==='responses'?`{"model":"${selectedModel}","input":"Say hello in one sentence."}`:`{"model":"${selectedModel}","messages":[{"role":"user","content":"Say hello in one sentence."}],"stream":false}`; + const body=wire==='messages'?`{"model":"${selectedModel}","max_tokens":256,"messages":[{"role":"user","content":"Say hello in one sentence."}]}`:wire==='responses'?`{"model":"${selectedModel}","input":"Say hello in one sentence."}`:wire==='embeddings'?`{"model":"${selectedModel}","input":"AIGW semantic search"}`:`{"model":"${selectedModel}","messages":[{"role":"user","content":"Say hello in one sentence."}],"stream":false}`; code=`export AIGW_API_KEY="your-key"\ncurl ${endpoint} \\\n ${headers} \\\n -H "Content-Type: application/json" \\\n -d '${body}'`; } else if(language==='python') { - code=wire==='messages'?`import os\nfrom anthropic import Anthropic\n\nclient = Anthropic(api_key=os.environ["AIGW_API_KEY"], base_url="${String(state.developerConfig.base_url||window.location.origin).replace(/\/$/,'')}/anthropic")\nmessage = client.messages.create(model="${selectedModel}", max_tokens=256, messages=[{"role":"user", "content":"Say hello in one sentence."}])\nprint(message.content[0].text)`:wire==='responses'?`import os\nfrom openai import OpenAI\n\nclient = OpenAI(api_key=os.environ["AIGW_API_KEY"], base_url="${String(state.developerConfig.base_url||window.location.origin).replace(/\/$/,'')}/v1")\nresponse = client.responses.create(model="${selectedModel}", input="Say hello in one sentence.")\nprint(response.output_text)`: `import os\nfrom openai import OpenAI\n\nclient = OpenAI(api_key=os.environ["AIGW_API_KEY"], base_url="${String(state.developerConfig.base_url||window.location.origin).replace(/\/$/,'')}/v1")\nresponse = client.chat.completions.create(model="${selectedModel}", messages=[{"role":"user", "content":"Say hello in one sentence."}])\nprint(response.choices[0].message.content)`; + if(wire==='messages')code=`import os\nfrom anthropic import Anthropic\n\nclient = Anthropic(api_key=os.environ["AIGW_API_KEY"], base_url="${base}/anthropic")\nmessage = client.messages.create(model="${selectedModel}", max_tokens=256, messages=[{"role":"user", "content":"Say hello in one sentence."}])\nprint(message.content[0].text)`; + else if(wire==='responses')code=`import os\nfrom openai import OpenAI\n\nclient = OpenAI(api_key=os.environ["AIGW_API_KEY"], base_url="${base}/v1")\nresponse = client.responses.create(model="${selectedModel}", input="Say hello in one sentence.")\nprint(response.output_text)`; + else if(wire==='embeddings')code=`import os\nfrom openai import OpenAI\n\nclient = OpenAI(api_key=os.environ["AIGW_API_KEY"], base_url="${base}/v1")\nresponse = client.embeddings.create(model="${selectedModel}", input="AIGW semantic search")\nprint(len(response.data[0].embedding))`; + else code=`import os\nfrom openai import OpenAI\n\nclient = OpenAI(api_key=os.environ["AIGW_API_KEY"], base_url="${base}/v1")\nresponse = client.chat.completions.create(model="${selectedModel}", messages=[{"role":"user", "content":"Say hello in one sentence."}])\nprint(response.choices[0].message.content)`; } else { - code=wire==='messages'?`import Anthropic from "@anthropic-ai/sdk";\n\nconst client = new Anthropic({ apiKey: process.env.AIGW_API_KEY, baseURL: "${String(state.developerConfig.base_url||window.location.origin).replace(/\/$/,'')}/anthropic" });\nconst message = await client.messages.create({ model: "${selectedModel}", max_tokens: 256, messages: [{ role: "user", content: "Say hello in one sentence." }] });\nconsole.log(message.content[0].text);`: `import OpenAI from "openai";\n\nconst client = new OpenAI({ apiKey: process.env.AIGW_API_KEY, baseURL: "${String(state.developerConfig.base_url||window.location.origin).replace(/\/$/,'')}/v1" });\nconst response = await client.${wire==='responses'?'responses.create({ model: "'+selectedModel+'", input: "Say hello in one sentence." })':'chat.completions.create({ model: "'+selectedModel+'", messages: [{ role: "user", content: "Say hello in one sentence." }] })'};\nconsole.log(${wire==='responses'?'response.output_text':'response.choices[0].message.content'});`; + if(wire==='messages')code=`import Anthropic from "@anthropic-ai/sdk";\n\nconst client = new Anthropic({ apiKey: process.env.AIGW_API_KEY, baseURL: "${base}/anthropic" });\nconst message = await client.messages.create({ model: "${selectedModel}", max_tokens: 256, messages: [{ role: "user", content: "Say hello in one sentence." }] });\nconsole.log(message.content[0].text);`; + else if(wire==='responses')code=`import OpenAI from "openai";\n\nconst client = new OpenAI({ apiKey: process.env.AIGW_API_KEY, baseURL: "${base}/v1" });\nconst response = await client.responses.create({ model: "${selectedModel}", input: "Say hello in one sentence." });\nconsole.log(response.output_text);`; + else if(wire==='embeddings')code=`import OpenAI from "openai";\n\nconst client = new OpenAI({ apiKey: process.env.AIGW_API_KEY, baseURL: "${base}/v1" });\nconst response = await client.embeddings.create({ model: "${selectedModel}", input: "AIGW semantic search" });\nconsole.log(response.data[0].embedding.length);`; + else code=`import OpenAI from "openai";\n\nconst client = new OpenAI({ apiKey: process.env.AIGW_API_KEY, baseURL: "${base}/v1" });\nconst response = await client.chat.completions.create({ model: "${selectedModel}", messages: [{ role: "user", content: "Say hello in one sentence." }] });\nconsole.log(response.choices[0].message.content);`; } $('#quickstart-code-block code').textContent=code; } function playgroundBody(wire,model,prompt,maxOutput) { if(wire==='messages')return {model,max_tokens:maxOutput,messages:[{role:'user',content:prompt}]}; if(wire==='responses')return {model,input:prompt,max_output_tokens:maxOutput,stream:false}; + if(wire==='embeddings')return {model,input:prompt}; return {model,messages:[{role:'user',content:prompt}],max_tokens:maxOutput,stream:false}; } function playgroundTokenCount(payload) { @@ -278,7 +298,7 @@ function showPlaygroundDiagnostic(status,network=false) { async function runPlayground(event) { event.preventDefault(); const key=$('#playground-key').value.trim();const model=modelSelector(selectedPlaygroundModel(),'#playground-provider');const wire=$('#playground-protocol').value;const prompt=$('#playground-prompt').value.trim();const maxOutput=Number($('#playground-max-output').value); - if(!key||!model||!prompt||!Number.isSafeInteger(maxOutput)||maxOutput<1)return toast('Complete the API request fields',true); + if(!key||!model||!prompt||(wire!=='embeddings'&&(!Number.isSafeInteger(maxOutput)||maxOutput<1)))return toast('Complete the API request fields',true); state.playgroundKey=key; const endpoint=quickstartEndpoint(wire);const headers={'Content-Type':'application/json'}; if(wire==='messages'){headers['X-API-Key']=key;headers['Anthropic-Version']='2023-06-01';}else headers.Authorization=`Bearer ${key}`; @@ -309,14 +329,18 @@ function renderCatalog() { $('#catalog-count').textContent=`${models.length} of ${state.developerModels.length} models`; $('#catalog-grid').innerHTML=models.map(item=>{const health=providerHealthSummary(item);return `
${esc(item.owned_by||'MODEL')}

${esc(item.display_name||item.public_id)}

${esc(item.public_id)}
${esc(health.label)}${esc(health.detail)}

${esc(item.description||'No description provided.')}

${(item.supported_wire_apis||[]).map(apiName=>`${esc(apiName)}`).join('')}${(item.input_modalities||[]).map(modality=>`${esc(modality)} input`).join('')}
${money(item.input_price_micros_per_million,item.price_currency)} in · ${money(item.output_price_micros_per_million,item.price_currency)} out / 1M tokens
${integer(item.context_window)} context · ${integer(item.max_output_tokens)} max output
`;}).join('')||'
No models match these filters.
'; } -function modelProtocolLabel(value) { return value==='chat_completions'?'OpenAI Chat Completions':value==='responses'?'OpenAI Responses':value==='messages'?'Anthropic Messages':value; } +function modelProtocolLabel(value) { return value==='chat_completions'?'OpenAI Chat Completions':value==='responses'?'OpenAI Responses':value==='embeddings'?'OpenAI Embeddings':value==='messages'?'Anthropic Messages':value; } function showModelDetails(publicID) { const model=state.developerModels.find(item=>item.public_id===publicID);if(!model)return; state.detailModel=model;$('#model-dialog-title').textContent=model.display_name||model.public_id;$('#model-dialog-id').textContent=model.public_id; const lifecycle=model.lifecycle||'active';const status=lifecycle==='retired'?'retired':lifecycle==='deprecated'?'deprecated':lifecycle==='preview'?'preview':'available'; const health=providerHealthSummary(model);const rows=[['Status',status],['Runtime',health.label],['Providers',health.detail],['Owner',model.owned_by||'—'],['Protocols',(model.supported_wire_apis||[]).map(modelProtocolLabel).join(', ')||'—'],['Input',(model.input_modalities||[]).join(', ')||'—'],['Output',(model.output_modalities||[]).join(', ')||'—'],['Context',`${integer(model.context_window)} tokens`],['Max output',`${integer(model.max_output_tokens)} tokens`],['Released',date(model.released_at)],['Capabilities',(model.capabilities||[]).join(', ')||'—'],['Aliases',(model.aliases||[]).join(', ')||'—']]; $('#model-detail-grid').innerHTML=rows.map(([label,value])=>`
${esc(label)}${esc(value)}
`).join(''); - const providers=model.providers||[];$('#model-provider-health').classList.toggle('hidden',providers.length===0);$('#model-provider-health-body').innerHTML=providers.map(item=>{const measured=Number(item.recent_samples||0)>0;const availability=measured?`${Number(item.availability_percent||0).toFixed(1)}%`:'Not measured';const latency=Number(item.header_latency_ewma_ms||0)>0?`${integer(item.header_latency_ewma_ms)} ms`:'Not measured';const retry=item.circuit_open_until?`Retry ${date(item.circuit_open_until)}`:'';return `${esc(item.name||item.slug)}${esc(item.slug)} · ${esc(modelProtocolLabel(item.wire_api)||item.protocol||'')}${esc(item.state==='open'?'Circuit open':item.state)}${retry?`${esc(retry)}`:''}${esc(availability)}${esc(latency)}${integer(item.attempts||0)}`;}).join(''); + const providers=model.providers||[];$('#model-provider-health').classList.toggle('hidden',providers.length===0); + const activeProbeCount=providers.reduce((sum,item)=>sum+Number(item.active_probes||0),0); + const sharedCount=providers.reduce((sum,item)=>sum+Number(item.shared_attempts||0)+Number(item.shared_ttft_samples||0),0); + $('#model-provider-health .section-heading small').textContent=sharedCount>0?'Measured from requests shared across gateway instances.':activeProbeCount>0?'Measured from requests and authenticated availability probes on this gateway instance.':'Measured from requests handled by this gateway instance.'; + $('#model-provider-health-body').innerHTML=providers.map(item=>{const measured=Number(item.recent_samples||0)>0;const availability=measured?`${Number(item.availability_percent||0).toFixed(1)}%`:'Not measured';const ttft=Number(item.ttft_ewma_ms||0)>0?`${integer(item.ttft_ewma_ms)} ms`:'Not measured';const latency=Number(item.header_latency_ewma_ms||0)>0?`${integer(item.header_latency_ewma_ms)} ms`:'Not measured';const retry=item.circuit_open_until?`Retry ${date(item.circuit_open_until)}`:'';const probes=Number(item.active_probes||0)>0?` · ${integer(item.active_probes)} probes`:'';const shared=Number(item.shared_attempts||0)>0?` · ${integer(item.shared_attempts)} shared`:'';return `${esc(item.name||item.slug)}${esc(item.slug)} · ${esc(modelProtocolLabel(item.wire_api)||item.protocol||'')}${esc(item.state==='open'?'Circuit open':item.state)}${retry?`${esc(retry)}`:''}${esc(availability)}${esc(ttft)}${esc(latency)}${integer(item.attempts||0)}${esc(shared)}${esc(probes)}`;}).join(''); $('#model-dialog-playground').disabled=modelIsUnavailable(model);$('#estimate-input').value='1000';$('#estimate-output').value='500';$('#estimate-cache-read').value='0';$('#estimate-cache-write').value='0';renderModelEstimate();$('#model-dialog').showModal(); } function renderModelEstimate() { @@ -333,10 +357,12 @@ function renderKeyProjects() { const tenant=$('#key-tenant').value;const node=$( function renderKeys() { const picker=$('#key-models');const selected=new Set([...picker.selectedOptions].map(option=>option.value));picker.innerHTML=state.developerModels.map(item=>``).join('');[...picker.options].forEach(option=>{option.selected=selected.has(option.value);}); $('#keys-body').innerHTML = state.keys.map(item => { - const models=item.allowed_models||[];const expires=item.expires_at?date(item.expires_at):'Never';const tags=(item.tags||[]).map(tag=>`${esc(tag)}`).join(''); - const spent=Number(item.current_month_spend_micros||0);const reserved=Number(item.current_month_reserved_micros||0);const cap=Number(item.monthly_spend_micros||0);const remaining=Math.max(0,cap-spent-reserved); - const effectiveStatus=item.status==='active'&&item.expires_at&&new Date(item.expires_at)<=new Date()?'expired':item.status; - return `${esc(item.name)}${tags?`${tags}`:''}${esc(item.key_prefix)}${shortID(item.project_id)}${(item.scopes||[]).map(scope=>`${esc(scope)}`).join('')}${models.length?`${integer(models.length)} selected model${models.length===1?'':'s'}`:'All visible models'}${money(spent)}${integer(item.current_month_requests)} requests · ${money(reserved)} reserved${cap?money(cap):'Unlimited'}${cap?`${money(remaining)} remaining`:'No key-level cap'} · Expires: ${esc(expires)}${esc(effectiveStatus)}Last used: ${esc(date(item.last_used_at))}${item.status==='active'&&can('keys.write')?``:''}`; + const models=item.allowed_models||[];const expires=item.expires_at?date(item.expires_at):'Never';const tags=(item.tags||[]).map(tag=>`${esc(tag)}`).join(''); + const spent=Number(item.current_month_spend_micros||0),reserved=Number(item.current_month_reserved_micros||0),cap=Number(item.monthly_spend_micros||0),remaining=Math.max(0,cap-spent-reserved); + const daySpent=Number(item.current_day_spend_micros||0),dayReserved=Number(item.current_day_reserved_micros||0),dayCap=Number(item.daily_spend_micros||0),dayRemaining=Math.max(0,dayCap-daySpent-dayReserved); + const effectiveStatus=item.status==='active'&&item.expires_at&&new Date(item.expires_at)<=new Date()?'expired':item.status; + const actions=can('keys.write')&&item.status!=='revoked'?`${item.status==='active'?``:``} `:''; + return `${esc(item.name)}${tags?`${tags}`:''}${esc(apiKeyDisplay(item))}${shortID(item.project_id)}${(item.scopes||[]).map(scope=>`${esc(scope)}`).join('')}${models.length?`${integer(models.length)} selected model${models.length===1?'':'s'}`:'All visible models'}${money(spent)} this month${integer(item.current_month_requests)} requests · ${money(reserved)} reserved${money(daySpent)} today · ${integer(item.current_day_requests)} requests${cap?`${money(remaining)} monthly left`:'Monthly unlimited'}${dayCap?`${money(dayRemaining)} daily left`:'Daily unlimited'}${item.requests_per_minute?`${integer(item.requests_per_minute)} RPM`:'RPM unlimited'} · ${item.tokens_per_minute?`${integer(item.tokens_per_minute)} TPM`:'TPM unlimited'} · Expires: ${esc(expires)}${esc(effectiveStatus)}Last used: ${esc(date(item.last_used_at))}${actions}`; }).join('') || emptyRow(8); } function renderProviders() { $('#providers-body').innerHTML = state.providers.map(item => `${esc(item.name)}${esc(item.slug)}${esc(item.protocol)}${esc(item.wire_api)}${esc(item.base_url)}${integer(item.route_count)}${item.enabled?'enabled':'disabled'}${can('platform.write')?``:''}`).join('') || emptyRow(6); } @@ -344,7 +370,7 @@ function renderModels() { $('#models-body').innerHTML = state.models.map(item => function renderBilling() { $('#billing-currency').textContent=(state.overview.billing_currency||'').toUpperCase(); $('#billing-accounts-body').innerHTML=state.billingAccounts.map(item=>`${esc(item.tenant_name)}
${shortID(item.tenant_id)}${money(item.balance_micros,item.currency)}${money(item.reserved_micros,item.currency)}${money(item.available_micros,item.currency)}${date(item.updated_at)}`).join('')||emptyRow(5); - $('#billing-ledger-body').innerHTML=state.ledger.map(item=>`${date(item.created_at)}${shortID(item.tenant_id)}${esc(item.kind)}${money(item.amount_micros,item.currency)}${money(item.balance_after_micros,item.currency)}${shortID(item.source_id)}`).join('')||emptyRow(6); + $('#billing-ledger-body').innerHTML=state.ledger.map(item=>{const requestReference=item.kind==='usage'&&item.source_type==='request'&&item.source_id;const reference=requestReference?``:`${shortID(item.source_id)}`;return `${date(item.created_at)}${shortID(item.tenant_id)}${esc(item.kind)}${money(item.amount_micros,item.currency)}${money(item.balance_after_micros,item.currency)}${reference}`;}).join('')||emptyRow(6); const orders=state.orders||[];$('#billing-orders-body').innerHTML=orders.map(item=>`${date(item.created_at)}${item.trigger_type==='auto'?'automatic':''}${money(item.amount_micros,item.currency)}${esc(item.status)}${esc(item.reconciliation_status||'unknown')}${item.invoice_url?`Invoice`:''} ${item.invoice_pdf_url?`PDF`:''} ${item.receipt_url?`Receipt`:''}${item.trigger_type!=='auto'&&['failed','expired'].includes(item.status)?``:''}${can('billing.adjust')&&item.status==='pending'&&item.reconciliation_status==='missing'?` `:''}${can('billing.adjust')&&item.status==='paid'&&item.reconciliation_status==='missing'&&!item.stripe_payment_intent_id?` `:''}${can('billing.adjust')&&['paid','partially_refunded'].includes(item.status)?` `:''}`).join('')||emptyRow(6); $('#refunds-body').innerHTML=(state.refunds||[]).map(item=>`${date(item.created_at)}${shortID(item.topup_order_id)}${money(item.amount_micros,item.currency)}${esc(item.status)}${esc(item.last_error||'—')}`).join('')||emptyRow(5); $('#disputes-body').innerHTML=(state.disputes||[]).map(item=>`${date(item.updated_at)}${money(item.amount_micros,item.currency)}${esc(item.status)}${esc(item.reason)}${date(item.due_by)}`).join('')||emptyRow(5); @@ -379,30 +405,47 @@ function renderSecurity() { $('#orders-body').innerHTML=(state.orders||[]).map(item=>`${date(item.created_at)}${money(item.amount_micros,item.currency)}${esc(item.status)}${shortID(item.id)}`).join('')||emptyRow(4); } function renderUsage() { - const project=$('#usage-project'),key=$('#usage-key'),model=$('#usage-model'),provider=$('#usage-provider');const projectValue=project.value,keyValue=key.value,modelValue=model.value,providerValue=provider.value; - const providers=new Map();state.developerModels.forEach(item=>(item.providers||[]).forEach(route=>providers.set(route.slug,route.name||route.slug))); - project.innerHTML=`${state.projects.map(item=>``).join('')}`;key.innerHTML=`${state.keys.map(item=>``).join('')}`;model.innerHTML=`${state.developerModels.map(item=>``).join('')}`;provider.innerHTML=`${[...providers].sort((a,b)=>a[1].localeCompare(b[1])).map(([slug,name])=>``).join('')}`; - if(projectValue)project.value=projectValue;if(keyValue)key.value=keyValue;if(modelValue)model.value=modelValue;if(providerValue)provider.value=providerValue; + const platformDiagnostics=!state.actor.tenant_id&&can('platform.read'); + const project=$('#usage-project'),key=$('#usage-key'),model=$('#usage-model'),provider=$('#usage-provider'); + const projectValue=project.value,keyValue=key.value,modelValue=model.value,providerValue=provider.value; + const providers=new Map();if(platformDiagnostics)state.developerModels.forEach(item=>(item.providers||[]).forEach(route=>providers.set(route.slug,route.name||route.slug))); + project.innerHTML=`${state.projects.map(item=>``).join('')}`; + key.innerHTML=`${state.keys.map(item=>``).join('')}`; + model.innerHTML=`${state.developerModels.map(item=>``).join('')}`; + provider.innerHTML=`${[...providers].sort((a,b)=>a[1].localeCompare(b[1])).map(([slug,name])=>``).join('')}`; + if(projectValue)project.value=projectValue;if(keyValue)key.value=keyValue;if(modelValue)model.value=modelValue;if(providerValue&&platformDiagnostics)provider.value=providerValue; + $('#usage-provider-filter').classList.toggle('hidden',!platformDiagnostics);$('#usage-provider-panel').classList.toggle('hidden',!platformDiagnostics); if(!$('#usage-from').value){const query=new URLSearchParams(defaultUsageQuery());$('#usage-from').value=query.get('from');$('#usage-to').value=query.get('to');} - const points=state.usageDaily||[];const totals=points.reduce((acc,item)=>{acc.requests+=Number(item.request_count||0);acc.success+=Number(item.successful_requests||0);acc.tokens+=Number(item.total_tokens||0);acc.charged+=Number(item.charged_micros||0);acc.duration+=Number(item.average_duration_ms||0)*Number(item.request_count||0);acc.p95=Math.max(acc.p95,Number(item.p95_duration_ms||0));return acc;},{requests:0,success:0,tokens:0,charged:0,duration:0,p95:0}); - const metrics=[['Requests',integer(totals.requests),'selected range'],['Success rate',percent(totals.success,totals.requests),'completed requests'],['Tokens',integer(totals.tokens),'input and output'],['Charged',money(totals.charged),'wallet debit'],['Average latency',totals.requests?`${integer(Math.round(totals.duration/totals.requests))} ms`:'—','request weighted'],['P95 latency',totals.p95?`${integer(totals.p95)} ms`:'—','highest daily P95']];$('#usage-metrics').innerHTML=metrics.map(([label,value,sub])=>`
${label}${esc(value)}${esc(sub)}
`).join(''); + const points=state.usageDaily||[];const totals=points.reduce((acc,item)=>{acc.requests+=Number(item.request_count||0);acc.success+=Number(item.successful_requests||0);acc.tokens+=Number(item.total_tokens||0);acc.charged+=Number(item.charged_micros||0);acc.duration+=Number(item.average_duration_ms||0)*Number(item.request_count||0);acc.p50=Math.max(acc.p50,Number(item.p50_duration_ms||0));acc.p95=Math.max(acc.p95,Number(item.p95_duration_ms||0));acc.ttftP95=Math.max(acc.ttftP95,Number(item.p95_ttft_ms||0));return acc;},{requests:0,success:0,tokens:0,charged:0,duration:0,p50:0,p95:0,ttftP95:0}); + const metrics=[['Requests',integer(totals.requests),'selected range'],['Success rate',percent(totals.success,totals.requests),'completed requests'],['Tokens',integer(totals.tokens),'input and output'],['Charged',money(totals.charged),'wallet debit'],['Average latency',totals.requests?`${integer(Math.round(totals.duration/totals.requests))} ms`:'—','request weighted'],['P50 latency',totals.p50?`${integer(totals.p50)} ms`:'—','highest daily median'],['P95 latency',totals.p95?`${integer(totals.p95)} ms`:'—','highest daily P95'],['P95 TTFT',totals.ttftP95?`${integer(totals.ttftP95)} ms`:'—','highest daily P95']];$('#usage-metrics').innerHTML=metrics.map(([label,value,sub])=>`
${label}${esc(value)}${esc(sub)}
`).join(''); const maxRequests=Math.max(1,...points.map(item=>Number(item.request_count||0)));$('#usage-chart').innerHTML=points.length?`
${points.map(item=>`
${new Date(item.day).toLocaleDateString(undefined,{month:'short',day:'numeric'})}
`).join('')}
`:'
No usage in this range
'; $('#usage-summary-body').innerHTML=state.usageSummary.map(item=>`${new Date(item.period_start).toLocaleDateString(undefined,{year:'numeric',month:'short'})}${esc(item.project_name)}${integer(item.request_count)}${percent(item.successful_requests,item.request_count)}${integer(item.input_tokens)}${integer(item.output_tokens)}${money(item.cost_micros)}`).join('')||emptyRow(7); - const analytics=Array.isArray(state.usageAnalytics)?{models:[],providers:[]}:state.usageAnalytics||{models:[],providers:[]}; + const analytics=Array.isArray(state.usageAnalytics)?{models:[],providers:[],keys:[]}:state.usageAnalytics||{models:[],providers:[],keys:[]}; const changeLabel=item=>item.charge_change_percent==null?(Number(item.previous_charged_micros||0)===0&&Number(item.charged_micros||0)>0?'New':'—'):`${Number(item.charge_change_percent)>=0?'+':''}${Number(item.charge_change_percent).toFixed(1)}%`; const changeClass=item=>item.charge_change_percent==null?(Number(item.charged_micros||0)>0?'positive':''):Number(item.charge_change_percent)>0?'money-negative':'positive'; const cacheRate=item=>{const denominator=Number(item.input_tokens||0)+Number(item.cache_read_input_tokens||0)+Number(item.cache_creation_input_tokens||0);return denominator?percent(item.cache_read_input_tokens,denominator):'—';}; - $('#usage-model-analytics-body').innerHTML=(analytics.models||[]).map(item=>`${esc(item.public_model)}${integer(item.provider_count)} provider${Number(item.provider_count)===1?'':'s'}${integer(item.request_count)}${percent(item.successful_requests,item.request_count)}${integer(item.total_tokens)}${money(item.charged_micros)}${changeLabel(item)}${integer(item.p95_duration_ms)} ms${integer(item.missing_usage_requests)}`).join('')||emptyRow(8); - $('#usage-provider-analytics-body').innerHTML=(analytics.providers||[]).map(item=>`${esc(item.provider_name)}${esc(item.wire_api||'unknown')}${integer(item.request_count)}${percent(item.successful_requests,item.request_count)}${integer(item.model_count)}${cacheRate(item)}${money(item.charged_micros)}${changeLabel(item)}${integer(item.p95_duration_ms)} ms`).join('')||emptyRow(8); - $('#usage-events-body').innerHTML=state.usage.map((item,index)=>`${date(item.started_at)}${shortID(item.request_id)}${esc(item.protocol)}${item.attempts>1?` · ${item.attempts} attempts`:''}${esc(item.project_name||shortID(item.project_id))}${esc(item.key_name||shortID(item.key_id))}${esc(item.public_model)}${esc(item.provider_name||'—')}${item.status_code}${item.error_type?`${esc(item.error_type)}`:''}${esc(item.metering_status||'')}${integer(item.total_tokens)}${money(item.charged_micros)}${item.uncollected_micros?`${money(item.uncollected_micros)} uncollected`:''}${integer(item.duration_ms)} ms`).join('')||emptyRow(9); + const latencyPair=item=>`${integer(item.p50_duration_ms)} / ${integer(item.p95_duration_ms)} ms`; + const ttftPair=item=>item.p95_ttft_ms?`${integer(item.p50_ttft_ms)} / ${integer(item.p95_ttft_ms)} ms`:'—'; + $('#usage-model-analytics-body').innerHTML=(analytics.models||[]).map(item=>`${esc(item.public_model)}${platformDiagnostics?`${integer(item.provider_count)} provider${Number(item.provider_count)===1?'':'s'}`:''}${integer(item.request_count)}${percent(item.successful_requests,item.request_count)}${integer(item.total_tokens)}${money(item.charged_micros)}${changeLabel(item)}${latencyPair(item)}${ttftPair(item)}${integer(item.missing_usage_requests)}`).join('')||emptyRow(9); + $('#usage-key-analytics-body').innerHTML=(analytics.keys||[]).map(item=>`${esc(item.key_name)}${shortID(item.key_id)}${integer(item.request_count)}${percent(item.successful_requests,item.request_count)}${integer(item.model_count)}${integer(item.total_tokens)}${money(item.charged_micros)}${latencyPair(item)}${ttftPair(item)}`).join('')||emptyRow(8); + $('#usage-provider-analytics-body').innerHTML=(analytics.providers||[]).map(item=>`${esc(item.provider_name)}${esc(item.wire_api||'unknown')}${integer(item.request_count)}${percent(item.successful_requests,item.request_count)}${integer(item.model_count)}${cacheRate(item)}${money(item.charged_micros)}${changeLabel(item)}${latencyPair(item)}${ttftPair(item)}`).join('')||emptyRow(9); + $('#usage-events-body').innerHTML=state.usage.map((item,index)=>`${date(item.started_at)}${shortID(item.request_id)}${esc(item.protocol)}${item.attempts>1?` · ${item.attempts} attempts`:''}${esc(item.project_name||shortID(item.project_id))}${esc(item.key_name||shortID(item.key_id))}${esc(item.public_model)}${platformDiagnostics&&item.provider_name?`${esc(item.provider_name)}`:''}${item.status_code}${item.error_type?`${esc(item.error_type)}`:''}${esc(item.metering_status||'')}${integer(item.total_tokens)}${money(item.charged_micros)}${item.uncollected_micros?`${money(item.uncollected_micros)} uncollected`:''}${integer(item.duration_ms)} msTTFT ${item.ttft_ms?`${integer(item.ttft_ms)} ms`:'—'}`).join('')||emptyRow(9); + $('#usage-page-label').textContent=`Page ${state.usagePaging.history.length+1}`;$('#usage-page-prev').disabled=state.usagePaging.history.length===0;$('#usage-page-next').disabled=!state.usagePaging.nextCursor; } function usageDiagnostic(item) { - return {request_id:item.request_id,started_at:item.started_at,status_code:item.status_code,success:Boolean(item.success),error_type:item.error_type||'',project:{id:item.project_id,name:item.project_name||''},api_key:{id:item.key_id,name:item.key_name||''},model:{public_id:item.public_model,provider:item.provider_name||item.provider_id||'',upstream_id:item.upstream_model||''},transport:{protocol:item.protocol,stream:Boolean(item.stream),attempts:Number(item.attempts||0),duration_ms:Number(item.duration_ms||0)},usage:{input_tokens:Number(item.input_tokens||0),output_tokens:Number(item.output_tokens||0),total_tokens:Number(item.total_tokens||0),cache_read_input_tokens:Number(item.cache_read_input_tokens||0),cache_creation_input_tokens:Number(item.cache_creation_input_tokens||0),reported:Boolean(item.usage_reported)},billing:{cost_micros:Number(item.cost_micros||0),charged_micros:Number(item.charged_micros||0),uncollected_micros:Number(item.uncollected_micros||0),metering_status:item.metering_status||''}}; + const diagnostic={request_id:item.request_id,started_at:item.started_at,status_code:item.status_code,success:Boolean(item.success),error_type:item.error_type||'',project:{id:item.project_id,name:item.project_name||''},api_key:{id:item.key_id,name:item.key_name||''},model:{public_id:item.public_model},transport:{protocol:item.protocol,stream:Boolean(item.stream),attempts:Number(item.attempts||0),duration_ms:Number(item.duration_ms||0),ttft_ms:Number(item.ttft_ms||0)},usage:{input_tokens:Number(item.input_tokens||0),output_tokens:Number(item.output_tokens||0),total_tokens:Number(item.total_tokens||0),cache_read_input_tokens:Number(item.cache_read_input_tokens||0),cache_creation_input_tokens:Number(item.cache_creation_input_tokens||0),reported:Boolean(item.usage_reported)},billing:{cost_micros:Number(item.cost_micros||0),charged_micros:Number(item.charged_micros||0),uncollected_micros:Number(item.uncollected_micros||0),metering_status:item.metering_status||''}}; + if(!state.actor.tenant_id&&can('platform.read'))diagnostic.model.route={provider:item.provider_name||item.provider_id||'',upstream_id:item.upstream_model||''};return diagnostic; +} +function openUsageDetails(item) { + if(!item)return;state.detailUsage=item;$('#usage-dialog-title').textContent=item.public_model||'Request details';$('#usage-dialog-id').textContent=item.request_id; + const throughput=item.duration_ms>0&&item.output_tokens>0?`${(Number(item.output_tokens)*1000/Number(item.duration_ms)).toFixed(1)} output tokens/s`:'—';const rows=[['Time',date(item.started_at)],['Status',`${item.status_code} · ${item.success?'success':'error'}`],['Project',item.project_name||shortID(item.project_id)],['API key',item.key_name||shortID(item.key_id)],['Protocol',`${item.protocol}${item.stream?' · stream':''}`],['Public model',item.public_model],['Attempts',integer(item.attempts)],['Latency',`${integer(item.duration_ms)} ms`],['TTFT',item.ttft_ms?`${integer(item.ttft_ms)} ms`:'—'],['Throughput',throughput],['Tokens',`${integer(item.input_tokens)} in · ${integer(item.output_tokens)} out`],['Cache',`${integer(item.cache_read_input_tokens)} read · ${integer(item.cache_creation_input_tokens)} write`],['Charged',money(item.charged_micros)],['Metering',item.metering_status||'—'],['Error',item.error_type||'—']]; + if(!state.actor.tenant_id&&can('platform.read'))rows.splice(6,0,['Provider',item.provider_name||item.provider_id||'—'],['Upstream model',item.upstream_model||'—']); + $('#usage-detail-grid').innerHTML=rows.map(([label,value])=>`
${esc(label)}${esc(value)}
`).join('');$('#usage-diagnostic').textContent=JSON.stringify(usageDiagnostic(item),null,2);const releaseForm=$('#release-reservation-form');releaseForm.classList.toggle('hidden',!can('billing.adjust')||item.metering_status!=='missing');releaseForm.reset();$('#usage-dialog').showModal(); } -function showUsageDetails(index) { - const item=state.usage[Number(index)];if(!item)return;state.detailUsage=item;$('#usage-dialog-title').textContent=item.public_model||'Request details';$('#usage-dialog-id').textContent=item.request_id; - const throughput=item.duration_ms>0&&item.output_tokens>0?`${(Number(item.output_tokens)*1000/Number(item.duration_ms)).toFixed(1)} output tokens/s`:'—';const rows=[['Time',date(item.started_at)],['Status',`${item.status_code} · ${item.success?'success':'error'}`],['Project',item.project_name||shortID(item.project_id)],['API key',item.key_name||shortID(item.key_id)],['Protocol',`${item.protocol}${item.stream?' · stream':''}`],['Public model',item.public_model],['Provider',item.provider_name||item.provider_id||'—'],['Upstream model',item.upstream_model||'—'],['Attempts',integer(item.attempts)],['Latency',`${integer(item.duration_ms)} ms`],['Throughput',throughput],['Tokens',`${integer(item.input_tokens)} in · ${integer(item.output_tokens)} out`],['Cache',`${integer(item.cache_read_input_tokens)} read · ${integer(item.cache_creation_input_tokens)} write`],['Charged',money(item.charged_micros)],['Metering',item.metering_status||'—'],['Error',item.error_type||'—']]; - $('#usage-detail-grid').innerHTML=rows.map(([label,value])=>`
${esc(label)}${esc(value)}
`).join('');$('#usage-diagnostic').textContent=JSON.stringify(usageDiagnostic(item),null,2);$('#usage-dialog').showModal(); +function showUsageDetails(index) { openUsageDetails(state.usage[Number(index)]); } +async function showLedgerRequest(requestID,tenantID) { + const params=new URLSearchParams({request_id:requestID,limit:'1'});if(!state.actor.tenant_id&&tenantID)params.set('tenant_id',tenantID); + const page=await api(`/usage?${params}`);const items=Array.isArray(page)?page:(page?.data||[]);if(!items.length)throw new Error('The request record is no longer available');openUsageDetails(items[0]); } function renderLimits() { const existing=new Map(state.limits.map(item=>[item.project_id,item])); @@ -426,9 +469,13 @@ document.addEventListener('click',async(event)=>{ const details=event.target.closest('[data-model-details]');if(details){showModelDetails(details.dataset.modelDetails);return;} const useModel=event.target.closest('[data-use-model]');if(useModel){useModelInPlayground(useModel.dataset.useModel);return;} const usageDetails=event.target.closest('[data-usage-details]');if(usageDetails){showUsageDetails(usageDetails.dataset.usageDetails);return;} + const ledgerRequest=event.target.closest('[data-ledger-request]');if(ledgerRequest){try{await showLedgerRequest(ledgerRequest.dataset.ledgerRequest,ledgerRequest.dataset.ledgerTenant);}catch(error){toast(error.message,true);}return;} if(event.target.id==='reload'){try{await api('/reload',{method:'POST',body:'{}'});await loadAll();toast('Snapshot reloaded');}catch(error){toast(error.message,true);}} if(event.target.id==='add-route')addRoute();if(event.target.closest('.remove-route'))event.target.closest('.route-row').remove(); - const revokeKey=event.target.closest('[data-revoke-key]');if(revokeKey&&confirm('Revoke this API key?')){try{await api(`/keys/${revokeKey.dataset.revokeKey}/revoke`,{method:'POST',body:'{}'});await loadAll();toast('API key revoked');}catch(error){toast(error.message,true);}} + const disableKey=event.target.closest('[data-disable-key]');if(disableKey&&confirm('Disable this API key immediately?')){try{await api(`/keys/${disableKey.dataset.disableKey}/disable`,{method:'POST',body:'{}'});await loadAll();toast('API key disabled');}catch(error){toast(error.message,true);}} + const enableKey=event.target.closest('[data-enable-key]');if(enableKey){try{await api(`/keys/${enableKey.dataset.enableKey}/enable`,{method:'POST',body:'{}'});await loadAll();toast('API key enabled');}catch(error){toast(error.message,true);}} + const rotateKey=event.target.closest('[data-rotate-key]');if(rotateKey&&confirm('Rotate this API key? The current secret will stop working immediately.')){try{const result=await api(`/keys/${rotateKey.dataset.rotateKey}/rotate`,{method:'POST',body:'{}'});state.playgroundKey=result.key;showSecret('API key rotated',result.key);await loadAll();toast(result.runtime_sync_status==='applied'?'API key rotated':'API key rotated; runtime reload is pending');}catch(error){toast(error.message,true);}} + const revokeKey=event.target.closest('[data-revoke-key]');if(revokeKey&&confirm('Permanently revoke this API key?')){try{await api(`/keys/${revokeKey.dataset.revokeKey}/revoke`,{method:'POST',body:'{}'});await loadAll();toast('API key revoked');}catch(error){toast(error.message,true);}} const revokeUser=event.target.closest('[data-revoke-user]');if(revokeUser&&confirm('Revoke this console credential?')){try{await api(`/users/${revokeUser.dataset.revokeUser}/revoke`,{method:'POST',body:'{}'});await loadAll();toast('Console credential revoked');}catch(error){toast(error.message,true);}} const provider=event.target.closest('[data-toggle-provider]');if(provider){try{await api(`/providers/${provider.dataset.toggleProvider}/toggle`,{method:'POST',body:JSON.stringify({enabled:provider.dataset.enabled==='true'})});await loadAll();toast('Provider updated');}catch(error){toast(error.message,true);}} const model=event.target.closest('[data-toggle-model]');if(model){try{await api(`/models/${model.dataset.toggleModel}/toggle`,{method:'POST',body:JSON.stringify({enabled:model.dataset.enabled==='true'})});await loadAll();toast('Model updated');}catch(error){toast(error.message,true);}} @@ -458,9 +505,9 @@ $('#bootstrap-pane').addEventListener('submit',async(event)=>{event.preventDefau $('#sign-out').addEventListener('click',async()=>{try{if(!state.token)await api('/auth/logout',{method:'POST',body:'{}'});}catch(error){if(error.status!==401)toast(error.message,true);}state.token='';state.csrf='';state.actor={};state.permissions=new Set();setConnected(false);}); $('#tenant-form').addEventListener('submit',async(event)=>{event.preventDefault();try{await api('/tenants',{method:'POST',body:JSON.stringify(formJSON(event.target))});event.target.reset();await loadAll();toast('Tenant created');}catch(error){toast(error.message,true);}}); $('#project-form').addEventListener('submit',async(event)=>{event.preventDefault();try{await api('/projects',{method:'POST',body:JSON.stringify(formJSON(event.target))});event.target.reset();await loadAll();toast('Project created');}catch(error){toast(error.message,true);}}); -$('#key-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const form=event.target;const data=formJSON(form);data.scopes=data.scopes.split(',').map(value=>value.trim()).filter(Boolean);data.tags=data.tags.split(',').map(value=>value.trim()).filter(Boolean);data.allowed_models=[...$('#key-models').selectedOptions].map(option=>option.value);data.monthly_spend_micros=data.monthly_spend.trim()?decimalToScaled(data.monthly_spend,6):0;delete data.monthly_spend;data.expires_at=data.expires_at?new Date(data.expires_at).toISOString():null;const result=await api('/keys',{method:'POST',body:JSON.stringify(data)});state.playgroundKey=result.key;form.reset();showSecret('API key created',result.key);await loadAll();goTo('quickstart');}catch(error){toast(error.message,true);}}); -$('#starter-key-form').addEventListener('submit',async(event)=>{event.preventDefault();const button=$('#starter-key-submit');try{const projectID=$('#starter-project').value;const name=$('#starter-key-name').value.trim();const model=selectedDeveloperModel();if(!projectID||!name||!model.public_id)throw new Error('An active project and model are required');button.disabled=true;const result=await api('/keys',{method:'POST',body:JSON.stringify({tenant_id:state.actor.tenant_id,project_id:projectID,name,scopes:['inference'],tags:['quickstart'],allowed_models:[model.public_id],monthly_spend_micros:0,expires_at:null})});state.playgroundKey=result.key;await loadAll();$('#playground-key').value=result.key;showSecret('Starter API key created',result.key);toast('Starter key is ready in the Playground');}catch(error){toast(error.message,true);}finally{button.disabled=false;}}); -function syncProviderWireAPI(){const form=$('#provider-form');const wire=form.elements.wire_api;const protocol=form.elements.protocol.value;wire.innerHTML=protocol==='anthropic'?'':'';} +$('#key-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const form=event.target;const data=formJSON(form);data.scopes=data.scopes.split(',').map(value=>value.trim()).filter(Boolean);data.tags=data.tags.split(',').map(value=>value.trim()).filter(Boolean);data.allowed_models=[...$('#key-models').selectedOptions].map(option=>option.value);data.monthly_spend_micros=data.monthly_spend.trim()?decimalToScaled(data.monthly_spend,6):0;data.daily_spend_micros=data.daily_spend.trim()?decimalToScaled(data.daily_spend,6):0;data.requests_per_minute=Number(data.requests_per_minute||0);data.tokens_per_minute=Number(data.tokens_per_minute||0);delete data.monthly_spend;delete data.daily_spend;data.expires_at=data.expires_at?new Date(data.expires_at).toISOString():null;const result=await api('/keys',{method:'POST',body:JSON.stringify(data)});state.playgroundKey=result.key;form.reset();showSecret('API key created',result.key);await loadAll();goTo('quickstart');if(result.runtime_sync_status!=='applied')toast('API key created; runtime reload is pending');}catch(error){toast(error.message,true);}}); +$('#starter-key-form').addEventListener('submit',async(event)=>{event.preventDefault();const button=$('#starter-key-submit');try{const projectID=$('#starter-project').value;const name=$('#starter-key-name').value.trim();const model=selectedDeveloperModel();if(!projectID||!name||!model.public_id)throw new Error('An active project and model are required');button.disabled=true;const result=await api('/keys',{method:'POST',body:JSON.stringify({tenant_id:state.actor.tenant_id,project_id:projectID,name,scopes:['inference'],tags:['quickstart'],allowed_models:[model.public_id],monthly_spend_micros:0,daily_spend_micros:0,requests_per_minute:0,tokens_per_minute:0,expires_at:null})});state.playgroundKey=result.key;await loadAll();$('#playground-key').value=result.key;showSecret('Starter API key created',result.key);toast(result.runtime_sync_status==='applied'?'Starter key is ready in the Playground':'Starter key created; runtime reload is pending');}catch(error){toast(error.message,true);}finally{button.disabled=false;}}); +function syncProviderWireAPI(){const form=$('#provider-form');const wire=form.elements.wire_api;const protocol=form.elements.protocol.value;wire.innerHTML=protocol==='anthropic'?'':'';} $('#provider-form [name=protocol]').addEventListener('change',syncProviderWireAPI); $('#provider-form').addEventListener('submit',async(event)=>{event.preventDefault();try{await api('/providers',{method:'POST',body:JSON.stringify(formJSON(event.target))});event.target.reset();syncProviderWireAPI();await loadAll();toast('Provider added');}catch(error){toast(error.message,true);}}); $('#model-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);data.input_price_micros_per_million=decimalToScaled(data.input_price,6);data.output_price_micros_per_million=decimalToScaled(data.output_price,6);data.cache_read_price_micros_per_million=decimalToScaled(data.cache_read_price,6);data.cache_write_price_micros_per_million=decimalToScaled(data.cache_write_price,6);delete data.input_price;delete data.output_price;delete data.cache_read_price;delete data.cache_write_price;for(const field of ['capabilities','input_modalities','output_modalities','regions','aliases','allowed_tenant_ids','allowed_key_ids'])data[field]=String(data[field]||'').split(',').map(value=>value.trim()).filter(Boolean);data.context_window=Number(data.context_window||0);data.max_output_tokens=Number(data.max_output_tokens||0);data.price_currency=state.overview.billing_currency||'usd';data.routes=$$('.route-row').map(row=>({provider_id:row.querySelector('.route-provider').value,upstream_model:row.querySelector('.route-upstream').value,priority:Number(row.querySelector('.route-priority').value),weight:Number(row.querySelector('.route-weight').value)}));await api('/models',{method:'POST',body:JSON.stringify(data)});event.target.reset();$('#route-editor').innerHTML='';renderRouteEditor();await loadAll();toast('Model created');}catch(error){toast(error.message,true);}}); @@ -469,7 +516,7 @@ $('#billing-profile-form').addEventListener('submit',async(event)=>{event.preven $('#auto-topup-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const currency=state.autoTopUp.currency||state.overview.billing_currency||'usd';const result=await api('/billing/auto-topup',{method:'PUT',body:JSON.stringify({enabled:$('#auto-topup-enabled').checked,threshold_micros:decimalToScaled($('#auto-topup-threshold').value,6),topup_amount_minor:decimalToScaled($('#auto-topup-amount').value,currencyDigits(currency))})});state.autoTopUp=result;renderAutoTopUp();renderQuickstart();toast(result.enabled?'Automatic top-up enabled':'Automatic top-up settings saved');}catch(error){toast(error.message,true);}}); $('#auto-topup-payment-setup').addEventListener('click',async()=>{try{const result=await api('/billing/auto-topup/setup-sessions',{method:'POST',body:'{}'});window.location.assign(result.url);}catch(error){toast(error.message,true);}}); $('#developer-preferences-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);await api('/developer/preferences',{method:'PUT',body:JSON.stringify({default_model:data.default_model||'',fallback_model:data.fallback_model||''})});await loadAll();toast('API defaults saved');}catch(error){toast(error.message,true);}}); -$('#billing-preferences-form').addEventListener('submit',async(event)=>{event.preventDefault();try{await api('/developer/preferences/billing',{method:'PUT',body:JSON.stringify({low_balance_enabled:$('#low-balance-enabled').checked,low_balance_threshold_micros:decimalToScaled($('#low-balance-threshold').value,6)})});await loadAll();toast('Balance alert saved');}catch(error){toast(error.message,true);}}); +$('#billing-preferences-form').addEventListener('submit',async(event)=>{event.preventDefault();try{await api('/developer/preferences/billing',{method:'PUT',body:JSON.stringify({low_balance_enabled:$('#low-balance-enabled').checked,low_balance_threshold_micros:decimalToScaled($('#low-balance-threshold').value,6),spend_anomaly_enabled:$('#spend-anomaly-enabled').checked,spend_anomaly_multiplier:Number($('#spend-anomaly-multiplier').value),spend_anomaly_min_micros:decimalToScaled($('#spend-anomaly-minimum').value,6)})});await loadAll();toast('Billing alerts saved');}catch(error){toast(error.message,true);}}); $('#adjustment-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);await api('/billing/adjustments',{method:'POST',body:JSON.stringify({tenant_id:data.tenant_id,amount_micros:decimalToScaled(data.amount,6),description:data.description})});event.target.reset();await loadAll();toast('Balance adjusted');}catch(error){toast(error.message,true);}}); $('#user-form').addEventListener('submit',async(event)=>{event.preventDefault();try{await api('/users',{method:'POST',body:JSON.stringify(formJSON(event.target))});event.target.reset();await loadAll();toast('Invitation sent');}catch(error){toast(error.message,true);}}); $('#password-form').addEventListener('submit',async(event)=>{event.preventDefault();try{const data=formJSON(event.target);await api('/auth/password',{method:'POST',body:JSON.stringify({current_password:data.current_password,new_password:data.new_password})});event.target.reset();state.csrf='';state.actor={};state.permissions=new Set();setConnected(false);authError('Password changed. Sign in again.',true);}catch(error){toast(error.message,true);}}); @@ -480,6 +527,7 @@ $('#revoke-other-sessions').addEventListener('click',async()=>{if(!confirm('Sign $('#close-dialog').addEventListener('click',()=>$('#secret-dialog').close());$('#copy-secret').addEventListener('click',async()=>{await navigator.clipboard.writeText($('#created-secret').textContent);toast('Credential copied');}); $('#close-model-dialog').addEventListener('click',()=>$('#model-dialog').close());$('#model-dialog-close').addEventListener('click',()=>$('#model-dialog').close());$('#model-dialog-playground').addEventListener('click',()=>{if(state.detailModel)useModelInPlayground(state.detailModel.public_id);});$('#copy-model-id').addEventListener('click',async()=>{await navigator.clipboard.writeText($('#model-dialog-id').textContent);toast('Model ID copied');}); $('#close-usage-dialog').addEventListener('click',()=>$('#usage-dialog').close());$('#usage-dialog-close').addEventListener('click',()=>$('#usage-dialog').close());$('#copy-request-id').addEventListener('click',async()=>{await navigator.clipboard.writeText(state.detailUsage?.request_id||'');toast('Request ID copied');});$('#copy-request-diagnostic').addEventListener('click',async()=>{await navigator.clipboard.writeText($('#usage-diagnostic').textContent);toast('Diagnostic copied');}); +$('#release-reservation-form').addEventListener('submit',async(event)=>{event.preventDefault();const item=state.detailUsage;if(!item)return;const reason=event.target.elements.reason.value.trim();if(!reason)return;const params=new URLSearchParams({tenant_id:item.tenant_id});try{await api(`/billing/reservations/${encodeURIComponent(item.request_id)}/release?${params}`,{method:'POST',body:JSON.stringify({reason})});$('#usage-dialog').close();await loadAll();toast('Reservation hold released');}catch(error){toast(error.message,true);}}); $('#model-estimate-form').addEventListener('submit',event=>event.preventDefault()); ['estimate-input','estimate-output','estimate-cache-read','estimate-cache-write'].forEach(id=>$('#'+id).addEventListener('input',renderModelEstimate)); $('#quickstart-model').addEventListener('change',()=>{syncQuickstartProtocols();renderQuickstartCode();renderDeveloperAccess();}); @@ -487,16 +535,26 @@ $('#quickstart-protocol').addEventListener('change',()=>{syncProviderSelect('#qu $('#quickstart-provider').addEventListener('change',renderQuickstartCode); $('#quickstart-language').addEventListener('change',renderQuickstartCode); $('#playground-model').addEventListener('change',syncPlaygroundProtocols); -$('#playground-protocol').addEventListener('change',()=>syncProviderSelect('#playground-provider',selectedPlaygroundModel(),$('#playground-protocol').value)); +$('#playground-protocol').addEventListener('change',syncPlaygroundProtocols); $('#playground-form').addEventListener('submit',runPlayground); $('#playground-stop').addEventListener('click',()=>state.playgroundController?.abort()); $('#playground-diagnostic-action').addEventListener('click',event=>goTo(event.currentTarget.dataset.target)); $('#catalog-search').addEventListener('input',renderCatalog);$('#catalog-protocol').addEventListener('change',renderCatalog);$('#catalog-input').addEventListener('change',renderCatalog);$('#catalog-owner').addEventListener('change',renderCatalog);$('#catalog-sort').addEventListener('change',renderCatalog); $('#copy-quickstart').addEventListener('click',async()=>{try{await navigator.clipboard.writeText($('#quickstart-code-block code').textContent);toast('Example copied');}catch(error){toast('Copy failed; select the example manually',true);}}); $('#copy-endpoint-env').addEventListener('click',async()=>{try{await navigator.clipboard.writeText($('#endpoint-env').textContent);toast('Environment copied');}catch(error){toast('Copy failed; select the environment manually',true);}}); -async function loadUsageFilters(){const params=new URLSearchParams(new FormData($('#usage-filter-form')));for(const [key,value] of [...params.entries()])if(!String(value).trim())params.delete(key);const query=params.toString();try{[state.usage,state.usageDaily,state.usageAnalytics]=await Promise.all([api(`/usage${query?`?${query}`:''}`),api(`/usage/daily${query?`?${query}`:''}`),api(`/usage/analytics${query?`?${query}`:''}`)]);renderUsage();}catch(error){toast(error.message,true);}} +function usageFilterParams(){const params=new URLSearchParams(new FormData($('#usage-filter-form')));for(const [key,value] of [...params.entries()])if(!String(value).trim())params.delete(key);if(state.actor.tenant_id)params.delete('provider');return params;} +async function loadUsageFilters({cursor='',history=[],refreshAggregates=true}={}){ + const params=usageFilterParams(),pageParams=new URLSearchParams(params);pageParams.set('limit','50');if(cursor)pageParams.set('cursor',cursor); + try{ + const pagePromise=api(`/usage?${pageParams}`);let page; + if(refreshAggregates){[page,state.usageDaily,state.usageAnalytics]=await Promise.all([pagePromise,api(`/usage/daily${params.size?`?${params}`:''}`),api(`/usage/analytics${params.size?`?${params}`:''}`)]);}else page=await pagePromise; + state.usage=Array.isArray(page)?page:(page.data||[]);state.usagePaging={cursor,nextCursor:page.next_cursor||'',history};renderUsage(); + }catch(error){toast(error.message,true);} +} $('#usage-filter-form').addEventListener('submit',async(event)=>{event.preventDefault();await loadUsageFilters();}); $('#usage-filter-reset').addEventListener('click',async()=>{const form=$('#usage-filter-form');form.reset();const query=new URLSearchParams(defaultUsageQuery());$('#usage-from').value=query.get('from');$('#usage-to').value=query.get('to');await loadUsageFilters();}); +$('#usage-page-next').addEventListener('click',async()=>{if(!state.usagePaging.nextCursor)return;await loadUsageFilters({cursor:state.usagePaging.nextCursor,history:[...state.usagePaging.history,state.usagePaging.cursor],refreshAggregates:false});}); +$('#usage-page-prev').addEventListener('click',async()=>{if(!state.usagePaging.history.length)return;const history=state.usagePaging.history.slice(0,-1);await loadUsageFilters({cursor:state.usagePaging.history.at(-1)||'',history,refreshAggregates:false});}); async function pollTopUp(orderID){for(let attempt=0;attempt<20;attempt++){const order=await api(`/billing/orders/${encodeURIComponent(orderID)}`);if(order.status==='paid'){await loadAll();toast('Balance credited');return;}if(['failed','expired'].includes(order.status)){toast(`Top-up ${order.status}`,true);return;}await new Promise(resolve=>setTimeout(resolve,1500));}toast('Payment is still processing');} async function pollAutoTopUpSetup(){for(let attempt=0;attempt<20;attempt++){const settings=await api('/billing/auto-topup');if(settings.payment_method_configured){state.autoTopUp=settings;renderBilling();renderQuickstart();toast('Payment method saved; review the threshold and enable automatic top-up');return;}await new Promise(resolve=>setTimeout(resolve,1500));}toast('Payment method setup is still processing',true);} async function start(){try{state.authConfig=await api('/auth/config');$('#register-tab').classList.toggle('hidden',!state.authConfig.registration_enabled);if(!state.authConfig.registration_enabled&&$('#register-tab').classList.contains('active'))showAuthPane('login-pane');const params=new URLSearchParams(location.search);const action=params.get('action');const token=params.get('token');if(params.get('auth')==='register'&&state.authConfig.registration_enabled)showAuthPane('register-pane');if(action==='reset-password'&&token){$('#reset-token').value=token;showAuthPane('reset-complete-pane');setConnected(false);return;}if(action==='accept-invite'&&token){$('#invite-token').value=token;showAuthPane('invite-pane');setConnected(false);return;}state.csrf=cookie('aigw_csrf');if(action==='verify-email'&&token){const result=await api('/auth/email/verify',{method:'POST',body:JSON.stringify({token})});history.replaceState({},'',location.pathname);await completeBrowserLogin(result);return;}const session=await api('/auth/session');if(session.authenticated){await loadAll(session);const orderID=params.get('order_id');if(params.get('topup')==='success'&&orderID){history.replaceState({},'',location.pathname);pollTopUp(orderID).catch(error=>toast(error.message,true));}else if(params.get('topup')==='cancel'){history.replaceState({},'',location.pathname);toast('Top-up cancelled');}else if(params.get('autotopup')==='setup'){history.replaceState({},'',location.pathname);pollAutoTopUpSetup().catch(error=>toast(error.message,true));}else if(params.get('autotopup')==='cancel'){history.replaceState({},'',location.pathname);toast('Payment method setup cancelled');}}else setConnected(false);}catch(error){setConnected(false);authError(error.message);}} diff --git a/internal/adminui/assets/index.html b/internal/adminui/assets/index.html index 8f31e1c..7d367ca 100644 --- a/internal/adminui/assets/index.html +++ b/internal/adminui/assets/index.html @@ -123,7 +123,7 @@ -
+
RECENT ACTIVITY

Latest requests

@@ -154,7 +157,7 @@
DISCOVER

Model catalog

-
+
@@ -173,8 +176,8 @@ - - + + @@ -184,12 +187,13 @@
-
COST RANKING

Model cost

Current range vs previous
ModelRequestsSuccessTokensChargedChangeP95Missing usage
-
ROUTE HEALTH

Provider performance

Cache hit and latency
ProviderRequestsSuccessModelsCache hitChargedChangeP95
+
COST RANKING

Model cost

Current range vs previous
ModelRequestsSuccessTokensChargedChangeLatency P50 / P95TTFT P50 / P95Missing usage
+
KEY ATTRIBUTION

API key usage

Spend and performance
API keyRequestsSuccessModelsTokensChargedLatency P50 / P95TTFT P50 / P95
+
ROUTE HEALTH

Provider performance

Platform diagnostics
ProviderRequestsSuccessModelsCache hitChargedChangeLatency P50 / P95TTFT P50 / P95
PeriodProjectRequestsSuccessInputOutputCost
-
REQUESTS

Recent events

-
TimeRequestProject / keyModel / providerStatusTokensChargedLatency
+
REQUESTS

Recent events

Page 1
+
TimeRequestProject / keyModelStatusTokensChargedLatency / TTFT
@@ -212,18 +216,21 @@ + + +
Key visibilityThe secret is shown only once after creation.
-
NamePrefixProjectRestrictionsMonth usageMonthly capActivity
+
NameKeyProjectRestrictionsUsageLimitsActivity
UPSTREAMS

Providers

-
+
NameProtocolBase URLRoutesStatus
@@ -303,8 +310,8 @@
-
MODEL DETAIL

Model

Provider health

Measured from requests handled by this gateway instance.
ProviderStateRecent availabilityHeader latencySamples

Cost estimate

-
API REQUEST

Request details

+
MODEL DETAIL

Model

Provider health

Measured from recent gateway requests.
ProviderStateRecent availabilityTTFTHeader latencySamples

Cost estimate

+
API REQUEST

Request details

ONE-TIME SECRET

Credential created

Copy this credential now. It will not be shown again.

diff --git a/internal/adminui/assets/models.css b/internal/adminui/assets/models.css index dc7c9c9..c373c2c 100644 --- a/internal/adminui/assets/models.css +++ b/internal/adminui/assets/models.css @@ -60,5 +60,24 @@ dialog::backdrop { background:rgba(16,42,58,.5); } .estimate-inputs { display:grid; grid-template-columns:1fr 1fr 1fr; align-items:end; gap:12px; } .estimate-inputs output { min-height:40px; display:flex; align-items:center; justify-content:center; color:var(--good); background:#f0fbf5; border:1px solid #c7e9d9; font-weight:750; } .model-dialog footer { display:flex; justify-content:flex-end; gap:8px; margin-top:22px; } +.model-page { max-width:1040px; } +.back-link { display:inline-flex; margin-bottom:22px; color:var(--accent); text-decoration:none; font-weight:700; } +.model-page-header { display:flex; align-items:flex-start; justify-content:space-between; gap:20px; padding-bottom:20px; border-bottom:1px solid var(--line); } +.model-page-header code { display:block; margin-top:8px; color:#486071; overflow-wrap:anywhere; } +.model-page-description { max-width:760px; margin:22px 0; color:#526673; font-size:16px; line-height:1.6; } +.model-page-layout { display:grid; grid-template-columns:minmax(0,1.55fr) minmax(250px,.7fr); gap:18px; margin-top:24px; } +.model-specs,.model-start,.code-example { background:var(--panel); border:1px solid var(--line); } +.model-specs { padding:22px; } +.model-specs h2,.model-start h2,.code-examples h2 { margin:6px 0 0; font-size:20px; } +.model-start { align-self:start; padding:22px; } +.model-start > code { display:block; margin:18px 0; color:#486071; overflow-wrap:anywhere; } +.model-start-actions { display:grid; gap:8px; } +.code-examples { margin-top:34px; } +.code-example-grid { display:grid; gap:14px; margin-top:17px; } +.code-example header { display:flex; align-items:center; justify-content:space-between; gap:16px; padding:13px 16px; border-bottom:1px solid var(--line); } +.code-example header code { color:#486071; font-size:12px; } +.code-example pre { margin:0; padding:18px; overflow:auto; background:#102a3a; color:#e9f7fb; white-space:pre-wrap; word-break:break-word; } +.code-example pre code { font-size:12px; line-height:1.6; } @media (max-width:900px) { .catalog-filters { grid-template-columns:repeat(2,minmax(0,1fr)); } .search-field { grid-column:1/-1; } .catalog-grid { grid-template-columns:repeat(2,minmax(0,1fr)); } } -@media (max-width:620px) { .catalog-header { padding:12px; } .catalog-brand small { display:none; } .catalog-header .button { padding:0 11px; } main { width:calc(100% - 24px); margin-top:20px; } .catalog-intro { align-items:start; flex-direction:column; } .catalog-stats { width:100%; justify-content:space-between; } .catalog-filters,.catalog-grid,.detail-grid,.estimate-inputs { grid-template-columns:1fr; } .search-field { grid-column:auto; } .model-card > p { min-height:0; } .model-dialog { padding:20px; } .model-dialog footer { flex-direction:column; } } +@media (max-width:760px) { .model-page-layout { grid-template-columns:1fr; } } +@media (max-width:620px) { .catalog-header { padding:12px; } .catalog-brand small { display:none; } .catalog-header .button { padding:0 11px; } main { width:calc(100% - 24px); margin-top:20px; } .catalog-intro,.model-page-header { align-items:start; flex-direction:column; } .catalog-stats { width:100%; justify-content:space-between; } .catalog-filters,.catalog-grid,.detail-grid,.estimate-inputs { grid-template-columns:1fr; } .search-field { grid-column:auto; } .model-card > p { min-height:0; } .model-dialog { padding:20px; } .model-dialog footer { flex-direction:column; } .code-example header { align-items:flex-start; flex-direction:column; } } diff --git a/internal/adminui/assets/models.js b/internal/adminui/assets/models.js index 4057ac5..a783d94 100644 --- a/internal/adminui/assets/models.js +++ b/internal/adminui/assets/models.js @@ -4,7 +4,7 @@ const state = { models: [], selected: null, registrationEnabled: false }; const $ = selector => document.querySelector(selector); const esc = value => String(value ?? '').replace(/[&<>'"]/g, char => ({'&':'&','<':'<','>':'>',"'":''','"':'"'}[char])); const integer = value => new Intl.NumberFormat().format(Number(value || 0)); -const protocolName = value => ({chat_completions:'Chat Completions',responses:'Responses',messages:'Anthropic Messages'}[value] || value); +const protocolName = value => ({chat_completions:'Chat Completions',responses:'Responses',embeddings:'Embeddings',messages:'Anthropic Messages'}[value] || value); const price = (micros, currency='usd') => new Intl.NumberFormat(undefined, {style:'currency', currency:String(currency).toUpperCase(), minimumFractionDigits:2, maximumFractionDigits:6}).format(Number(micros || 0) / 1_000_000); const date = value => value ? new Intl.DateTimeFormat(undefined, {year:'numeric',month:'short',day:'numeric'}).format(new Date(value)) : 'Not published'; @@ -59,7 +59,7 @@ function renderCatalog() {

${esc(model.description || 'No description published.')}

${(model.supported_wire_apis || []).map(item => `${esc(protocolName(item))}`).join('')}${(model.input_modalities || []).map(item => `${esc(item)}`).join('')}
Input
${price(model.input_price_micros_per_million, model.price_currency)} / 1M
Output
${price(model.output_price_micros_per_million, model.price_currency)} / 1M
Context
${integer(model.context_window)}
- + View model `; }).join(''); } diff --git a/internal/adminui/assets/style.css b/internal/adminui/assets/style.css index b71dc20..b4c7a0f 100644 --- a/internal/adminui/assets/style.css +++ b/internal/adminui/assets/style.css @@ -6,8 +6,9 @@ .section { display:none; } .section.active { display:block; } .section-heading { display:flex; justify-content:space-between; align-items:flex-end; gap:20px; margin-bottom:19px; } .eyebrow { color:var(--accent); font-size:10px; letter-spacing:0; font-weight:800; } h1 { font-size:28px; line-height:1.1; margin:7px 0 0; letter-spacing:0; } h2 { margin:4px 0 0; font-size:20px; } .metric-grid { display:grid; grid-template-columns:repeat(6,1fr); gap:12px; } .metric { background:var(--panel); border:1px solid var(--line); padding:18px; box-shadow:var(--shadow); } .metric span,.metric small { display:block; color:var(--muted); } .metric strong { display:block; font-size:28px; margin:12px 0 3px; font-weight:750; } .metric small { font-size:11px; } .panel { background:var(--panel); border:1px solid var(--line); box-shadow:var(--shadow); padding:20px; margin-bottom:16px; } .note,.warning { display:flex; align-items:flex-start; gap:12px; } .note-icon { flex:0 0 22px; height:22px; border:1px solid var(--accent); color:var(--accent); display:grid; place-items:center; font-weight:700; } .note p { margin:5px 0 0; color:var(--muted); } .warning { color:#6d5523; background:#fff9e9; border-color:#ead9a9; box-shadow:none; } .warning span { margin-left:8px; color:#887650; } -.form-grid { display:grid; grid-template-columns:repeat(3,minmax(0,1fr)); align-items:end; gap:13px; } label { display:flex; flex-direction:column; gap:7px; color:var(--muted); font-size:12px; font-weight:650; } input,select,textarea { width:100%; border:1px solid var(--line); background:#fff; color:var(--ink); padding:10px 11px; min-height:40px; outline:none; } textarea { resize:vertical; line-height:1.5; } input:focus,select:focus,textarea:focus { border-color:#69a9bf; box-shadow:0 0 0 3px var(--accent-soft); } .button { border:1px solid transparent; min-height:40px; padding:0 15px; font-weight:700; } .button.primary { color:#fff; background:var(--accent); } .button.primary:hover { background:#0d5879; } .button.secondary { color:var(--accent); background:var(--accent-soft); border-color:#c5e1ea; } .button.subtle { color:var(--accent); background:#fff; border-color:var(--line); grid-column:1; } +.form-grid { display:grid; grid-template-columns:repeat(3,minmax(0,1fr)); align-items:end; gap:13px; } label { display:flex; flex-direction:column; gap:7px; color:var(--muted); font-size:12px; font-weight:650; } input,select,textarea { width:100%; border:1px solid var(--line); background:#fff; color:var(--ink); padding:10px 11px; min-height:40px; outline:none; } textarea { resize:vertical; line-height:1.5; } input:focus,select:focus,textarea:focus { border-color:#69a9bf; box-shadow:0 0 0 3px var(--accent-soft); } .button { border:1px solid transparent; min-height:40px; padding:0 15px; font-weight:700; } .button.primary { color:#fff; background:var(--accent); } .button.primary:hover { background:#0d5879; } .button.secondary { color:var(--accent); background:var(--accent-soft); border-color:#c5e1ea; } .button.subtle { color:var(--accent); background:#fff; border-color:var(--line); grid-column:1; } .button.danger { color:var(--danger); background:#fff0f0; border-color:#e7baba; } .table-wrap { overflow:auto; padding:0; } table { width:100%; border-collapse:collapse; min-width:700px; } th,td { padding:14px 18px; text-align:left; border-bottom:1px solid var(--line); vertical-align:middle; } th { color:var(--muted); font-size:11px; font-weight:700; text-transform:uppercase; letter-spacing:0; background:#fbfcfd; } tbody tr:last-child td { border-bottom:0; } td { font-size:13px; } code { font-family:"SFMono-Regular",Consolas,monospace; font-size:12px; color:#486071; } .badge,.tag { display:inline-flex; align-items:center; padding:4px 7px; font-size:11px; line-height:1; } .badge { border:1px solid #d9e0e4; color:var(--muted); } .badge.active { color:#187151; background:#e9f7f0; border-color:#c7e9d9; } .badge.revoked,.badge.suspended,.badge.expired,.badge.unavailable { color:var(--danger); background:#fff0f0; border-color:#f0cccc; } .badge.degraded { color:#765b16; background:#fff8dc; border-color:#e9d990; } .tag { color:#4c6572; background:#eef3f5; margin:2px 3px 2px 0; } .text-button { border:0; background:transparent; color:var(--accent); padding:5px 0; } .text-button.danger { color:var(--danger); } .empty { color:var(--muted); text-align:center; padding:32px; } .truncate { max-width:280px; overflow:hidden; text-overflow:ellipsis; white-space:nowrap; } +.ledger-request-link code { color:inherit; text-decoration:underline; text-underline-offset:3px; } .route-editor { grid-column:1/-1; display:flex; flex-direction:column; gap:8px; } .route-row { display:grid; grid-template-columns:1.4fr 1.4fr .6fr .6fr 34px; gap:8px; } .icon-button { border:1px solid var(--line); background:#fff; color:var(--muted); width:34px; height:34px; font-size:19px; } .route-list { display:flex; flex-direction:column; gap:3px; color:#486071; font-size:12px; } .route-list em { color:var(--muted); font-style:normal; margin-left:4px; } .hidden { display:none !important; } .visually-hidden { position:absolute !important; width:1px !important; height:1px !important; padding:0 !important; margin:-1px !important; overflow:hidden !important; clip:rect(0,0,0,0) !important; white-space:nowrap !important; border:0 !important; } .billing-actions { display:grid; grid-template-columns:1fr 1fr; gap:16px; } .compact-form { grid-template-columns:1fr 1fr; } .compact-form .button { grid-column:1/-1; } .currency-label { color:var(--muted); font-size:12px; text-transform:uppercase; } .ledger-heading { margin-top:28px; } .money-positive { color:#187151; } .money-negative { color:var(--danger); } .account-grid { display:grid; grid-template-columns:1fr 1fr; gap:16px; } .account-grid .panel { align-content:start; } #totp-qr { width:180px; height:180px; border:1px solid var(--line); } #totp-secret { overflow-wrap:anywhere; } .table-wrap > .section-heading { padding:18px; margin:0; align-items:center; border-bottom:1px solid var(--line); } @@ -89,6 +90,10 @@ .analytics-panel { padding-top:0; } .analytics-panel .section-heading { min-width:760px; margin:0; padding:18px 18px 14px; } .analytics-table { min-width:960px; } +.badge.disabled { color:#765b16; background:#fff8dc; border-color:#e9d990; } +.usage-pager { display:flex; align-items:center; gap:10px; color:var(--muted); font-size:12px; } +.usage-pager .button { min-height:34px; padding:0 11px; } +.usage-pager .button:disabled { opacity:.45; cursor:not-allowed; } .positive { color:#187151; } .auto-topup-panel .section-heading { align-items:center; } .billing-profile-panel { margin-top:16px; } diff --git a/internal/auth/static.go b/internal/auth/static.go index 2b44158..67756bf 100644 --- a/internal/auth/static.go +++ b/internal/auth/static.go @@ -27,6 +27,9 @@ type KeyRecord struct { Scopes []string `json:"scopes"` AllowedModels []string `json:"allowed_models,omitempty"` MonthlySpendMicros int64 `json:"monthly_spend_micros,omitempty"` + DailySpendMicros int64 `json:"daily_spend_micros,omitempty"` + RequestsPerMinute int64 `json:"requests_per_minute,omitempty"` + TokensPerMinute int64 `json:"tokens_per_minute,omitempty"` ExpiresAt *time.Time `json:"expires_at,omitempty"` } @@ -78,7 +81,8 @@ func NewStatic(raw string, allowAnonymous bool) (*StaticAuthenticator, error) { hashed = append(hashed, HashedKeyRecord{Hash: hash, Principal: domain.Principal{ KeyID: record.KeyID, TenantID: record.TenantID, ProjectID: record.ProjectID, Scopes: append([]string(nil), record.Scopes...), AllowedModels: allowedModels, - MonthlySpendMicros: record.MonthlySpendMicros, ExpiresAt: record.ExpiresAt, + MonthlySpendMicros: record.MonthlySpendMicros, DailySpendMicros: record.DailySpendMicros, + RequestsPerMinute: record.RequestsPerMinute, TokensPerMinute: record.TokensPerMinute, ExpiresAt: record.ExpiresAt, }}) } if len(hashed) == 0 && !allowAnonymous { diff --git a/internal/auth/static_test.go b/internal/auth/static_test.go index a037ce1..28ca39f 100644 --- a/internal/auth/static_test.go +++ b/internal/auth/static_test.go @@ -78,7 +78,7 @@ func TestStaticAuthenticatorAcceptsAnthropicHeader(t *testing.T) { func TestStaticAuthenticatorLoadsRestrictionsAndRejectsExpiredKey(t *testing.T) { future := time.Now().Add(time.Hour).UTC().Format(time.RFC3339Nano) - authenticator, err := NewStatic(`[{"key":"sk-limited","key_id":"key-1","tenant_id":"tenant-1","project_id":"project-1","allowed_models":["model/allowed"],"monthly_spend_micros":1250000,"expires_at":"`+future+`"}]`, false) + authenticator, err := NewStatic(`[{"key":"sk-limited","key_id":"key-1","tenant_id":"tenant-1","project_id":"project-1","allowed_models":["model/allowed"],"monthly_spend_micros":1250000,"daily_spend_micros":250000,"requests_per_minute":12,"tokens_per_minute":3400,"expires_at":"`+future+`"}]`, false) if err != nil { t.Fatal(err) } @@ -91,6 +91,9 @@ func TestStaticAuthenticatorLoadsRestrictionsAndRejectsExpiredKey(t *testing.T) if principal.MonthlySpendMicros != 1_250_000 { t.Fatalf("monthly spend limit = %d", principal.MonthlySpendMicros) } + if principal.DailySpendMicros != 250_000 || principal.RequestsPerMinute != 12 || principal.TokensPerMinute != 3400 { + t.Fatalf("key spend or rate controls were not loaded: %+v", principal) + } if _, ok := principal.AllowedModels["model/allowed"]; !ok { t.Fatalf("allowed model was not loaded: %+v", principal.AllowedModels) } diff --git a/internal/billing/auto_topup.go b/internal/billing/auto_topup.go index a90405b..7fa4e59 100644 --- a/internal/billing/auto_topup.go +++ b/internal/billing/auto_topup.go @@ -160,6 +160,23 @@ func (s *Service) CreateAutoTopUpSetupSession(ctx context.Context, input AutoTop if err != nil { return AutoTopUpSetupResult{}, err } + params := s.autoTopUpSetupSessionParams(input, customerID) + session, err := s.createStripeCheckout(ctx, params) + if err != nil { + return AutoTopUpSetupResult{}, fmt.Errorf("create automatic top-up setup session: %w", err) + } + if session == nil || session.ID == "" || session.URL == "" { + return AutoTopUpSetupResult{}, errors.New("Stripe returned an incomplete setup session") + } + if _, err := s.db.Exec(ctx, `INSERT INTO tenant_auto_topup_settings (tenant_id,threshold_micros,topup_amount_minor,stripe_setup_session_id) + VALUES ($1,$2,$3,$4) ON CONFLICT (tenant_id) DO UPDATE SET stripe_setup_session_id=EXCLUDED.stripe_setup_session_id,updated_at=now()`, + input.TenantID, s.defaultAutoTopUpThresholdMicros(), s.defaultAutoTopUpAmountMinor(), session.ID); err != nil { + return AutoTopUpSetupResult{}, fmt.Errorf("persist automatic top-up setup session: %w", err) + } + return AutoTopUpSetupResult{SessionID: session.ID, URL: session.URL}, nil +} + +func (s *Service) autoTopUpSetupSessionParams(input AutoTopUpSetupInput, customerID string) *stripe.CheckoutSessionCreateParams { params := &stripe.CheckoutSessionCreateParams{ Mode: stripe.String(string(stripe.CheckoutSessionModeSetup)), Currency: stripe.String(s.currency), @@ -181,19 +198,7 @@ func (s *Service) CreateAutoTopUpSetupSession(ctx context.Context, input AutoTop } } params.SetIdempotencyKey("aigw_autotopup_setup_" + randomHex(16)) - session, err := s.createStripeCheckout(ctx, params) - if err != nil { - return AutoTopUpSetupResult{}, fmt.Errorf("create automatic top-up setup session: %w", err) - } - if session.ID == "" || session.URL == "" { - return AutoTopUpSetupResult{}, errors.New("Stripe returned an incomplete setup session") - } - if _, err := s.db.Exec(ctx, `INSERT INTO tenant_auto_topup_settings (tenant_id,threshold_micros,topup_amount_minor,stripe_setup_session_id) - VALUES ($1,$2,$3,$4) ON CONFLICT (tenant_id) DO UPDATE SET stripe_setup_session_id=EXCLUDED.stripe_setup_session_id,updated_at=now()`, - input.TenantID, s.defaultAutoTopUpThresholdMicros(), s.defaultAutoTopUpAmountMinor(), session.ID); err != nil { - return AutoTopUpSetupResult{}, fmt.Errorf("persist automatic top-up setup session: %w", err) - } - return AutoTopUpSetupResult{SessionID: session.ID, URL: session.URL}, nil + return params } func autoTopUpReturnURL(raw string, success bool) string { @@ -344,16 +349,7 @@ func (s *Service) processAutoTopUpOnce(ctx context.Context) (bool, error) { if err := tx.Commit(ctx); err != nil { return false, err } - params := &stripe.PaymentIntentCreateParams{ - Amount: stripe.Int64(amountMinor), Currency: stripe.String(currency), Customer: stripe.String(customerID), - PaymentMethod: stripe.String(paymentMethodID), Confirm: stripe.Bool(true), OffSession: stripe.Bool(true), - ErrorOnRequiresAction: stripe.Bool(true), Description: stripe.String("AIGW automatic prepaid balance top-up"), - Metadata: map[string]string{"aigw_action": "auto_topup", "aigw_tenant_id": tenantID, "aigw_topup_order_id": orderID}, - } - if strings.TrimSpace(customerEmail) != "" { - params.ReceiptEmail = stripe.String(strings.TrimSpace(customerEmail)) - } - params.SetIdempotencyKey("aigw_autotopup_" + orderID) + params := s.autoTopUpPaymentIntentParams(tenantID, orderID, customerID, customerEmail, paymentMethodID, currency, amountMinor) intent, callErr := s.createStripePaymentIntent(ctx, params) if callErr != nil { var stripeErr *stripe.Error @@ -387,6 +383,20 @@ func (s *Service) processAutoTopUpOnce(ctx context.Context) (bool, error) { return true, s.failAutoTopUp(ctx, tenantID, orderID, fmt.Errorf("automatic top-up PaymentIntent ended in status %s", intent.Status)) } +func (s *Service) autoTopUpPaymentIntentParams(tenantID, orderID, customerID, customerEmail, paymentMethodID, currency string, amountMinor int64) *stripe.PaymentIntentCreateParams { + params := &stripe.PaymentIntentCreateParams{ + Amount: stripe.Int64(amountMinor), Currency: stripe.String(currency), Customer: stripe.String(customerID), + PaymentMethod: stripe.String(paymentMethodID), Confirm: stripe.Bool(true), OffSession: stripe.Bool(true), + ErrorOnRequiresAction: stripe.Bool(true), Description: stripe.String("AIGW automatic prepaid balance top-up"), + Metadata: map[string]string{"aigw_action": "auto_topup", "aigw_tenant_id": tenantID, "aigw_topup_order_id": orderID}, + } + if strings.TrimSpace(customerEmail) != "" { + params.ReceiptEmail = stripe.String(strings.TrimSpace(customerEmail)) + } + params.SetIdempotencyKey("aigw_autotopup_" + orderID) + return params +} + func (s *Service) scheduleAutoTopUpRetry(ctx context.Context, tenantID, orderID string, cause error) error { message := truncateError(cause) _, err := s.db.Exec(ctx, `UPDATE topup_orders SET reconciliation_error=$2 WHERE id=$1 AND status='pending'`, orderID, message) diff --git a/internal/billing/auto_topup_test.go b/internal/billing/auto_topup_test.go index 85edb3d..55dc22d 100644 --- a/internal/billing/auto_topup_test.go +++ b/internal/billing/auto_topup_test.go @@ -25,6 +25,49 @@ func TestAutoTopUpReturnURL(t *testing.T) { } } +func TestAutoTopUpStripeContracts(t *testing.T) { + service := &Service{ + currency: "usd", + stripeSuccessURL: "https://console.example.test/billing?topup=success", + stripeCancelURL: "https://console.example.test/billing?topup=cancel", + integrationIdentifier: "aigw_balance_abcdefgh", + } + setup := service.autoTopUpSetupSessionParams(AutoTopUpSetupInput{ + TenantID: "tenant-123", CustomerEmail: " billing@example.test ", + }, "") + if setup.Mode == nil || *setup.Mode != string(stripe.CheckoutSessionModeSetup) || setup.Currency == nil || *setup.Currency != "usd" { + t.Fatalf("unexpected setup contract %+v", setup) + } + if len(setup.PaymentMethodTypes) != 0 || setup.CustomerCreation == nil || *setup.CustomerCreation != string(stripe.CheckoutSessionCustomerCreationAlways) { + t.Fatal("setup Checkout must create a customer and use Dashboard-managed payment methods") + } + if setup.CustomerEmail == nil || *setup.CustomerEmail != "billing@example.test" || setup.Metadata["aigw_action"] != autoTopUpAction { + t.Fatal("setup Checkout customer or metadata contract is incomplete") + } + if setup.IdempotencyKey == nil || !strings.HasPrefix(*setup.IdempotencyKey, "aigw_autotopup_setup_") { + t.Fatalf("setup idempotency key = %v", setup.IdempotencyKey) + } + + payment := service.autoTopUpPaymentIntentParams( + "tenant-123", "order-123", "cus_123", " billing@example.test ", "pm_123", "usd", 2000, + ) + if payment.Amount == nil || *payment.Amount != 2000 || payment.Currency == nil || *payment.Currency != "usd" || + payment.Customer == nil || *payment.Customer != "cus_123" || payment.PaymentMethod == nil || *payment.PaymentMethod != "pm_123" { + t.Fatalf("unexpected automatic top-up PaymentIntent %+v", payment) + } + if payment.Confirm == nil || !*payment.Confirm || payment.OffSession == nil || !*payment.OffSession || + payment.ErrorOnRequiresAction == nil || !*payment.ErrorOnRequiresAction { + t.Fatal("automatic top-up must be confirmed off-session and stop on required customer action") + } + if payment.ReceiptEmail == nil || *payment.ReceiptEmail != "billing@example.test" || + payment.Metadata["aigw_topup_order_id"] != "order-123" { + t.Fatal("automatic top-up receipt or reconciliation metadata is incomplete") + } + if payment.IdempotencyKey == nil || *payment.IdempotencyKey != "aigw_autotopup_order-123" { + t.Fatalf("payment idempotency key = %v", payment.IdempotencyKey) + } +} + func TestAutomaticTopUpSetupAndCreditAreIdempotentPostgres(t *testing.T) { databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL") if databaseURL == "" { diff --git a/internal/billing/ledger.go b/internal/billing/ledger.go index 7a00081..64f258c 100644 --- a/internal/billing/ledger.go +++ b/internal/billing/ledger.go @@ -190,6 +190,83 @@ func (s *Service) AdjustBalance(ctx context.Context, input AdjustmentInput) (Led return result, nil } +// ReleaseUnmeteredReservation is an audited operational escape hatch for a +// fail-closed success that cannot be reconciled. It never invents usage or +// changes wallet balance; it only returns the existing hold to availability. +func (s *Service) ReleaseUnmeteredReservation(ctx context.Context, tenantID, requestID string, input ReleaseReservationInput, actor ResolutionActor) (ReservationRelease, error) { + tenantID = strings.TrimSpace(tenantID) + requestID = strings.TrimSpace(requestID) + reason := normalizeDescription(input.Reason) + if tenantID == "" || requestID == "" || reason == "" || actor.ID == "" || actor.Type == "" { + return ReservationRelease{}, fmt.Errorf("%w: tenant, request, reason, and actor are required", ErrReservationNotReleasable) + } + tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted}) + if err != nil { + return ReservationRelease{}, fmt.Errorf("begin reservation release: %w", err) + } + defer tx.Rollback(ctx) + var result ReservationRelease + var projectID, currency, status string + if err := tx.QueryRow(ctx, `SELECT request_id,tenant_id::text,project_id::text,currency,reserved_micros,status + FROM billing_reservations WHERE request_id=$1 AND tenant_id=$2 FOR UPDATE`, requestID, tenantID).Scan( + &result.RequestID, &result.TenantID, &projectID, ¤cy, &result.ReservedMicros, &status); errors.Is(err, pgx.ErrNoRows) { + return ReservationRelease{}, ErrReservationNotReleasable + } else if err != nil { + return ReservationRelease{}, fmt.Errorf("lock reservation for release: %w", err) + } + if status == "released" { + var evidenceExists bool + if err := tx.QueryRow(ctx, `SELECT EXISTS(SELECT 1 FROM billing_ledger + WHERE source_type='unmetered_reservation' AND source_id=$1)`, requestID).Scan(&evidenceExists); err != nil { + return ReservationRelease{}, fmt.Errorf("read reservation release evidence: %w", err) + } + if !evidenceExists { + return ReservationRelease{}, fmt.Errorf("%w: reservation was released by normal settlement", ErrReservationNotReleasable) + } + if err := tx.QueryRow(ctx, `SELECT COALESCE(settled_at,created_at) FROM billing_reservations WHERE request_id=$1`, requestID).Scan(&result.ReleasedAt); err != nil { + return ReservationRelease{}, err + } + result.Status = status + return result, tx.Commit(ctx) + } + if status != "metering_failed" { + return ReservationRelease{}, fmt.Errorf("%w: reservation status is %s", ErrReservationNotReleasable, status) + } + var balance, held int64 + if err := tx.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1 FOR UPDATE`, tenantID).Scan(&balance, &held); err != nil { + return ReservationRelease{}, fmt.Errorf("lock wallet for reservation release: %w", err) + } + if result.ReservedMicros > held { + return ReservationRelease{}, errors.New("wallet reservation invariant violated during release") + } + if _, err := tx.Exec(ctx, `UPDATE tenant_wallets SET reserved_micros=reserved_micros-$2,updated_at=now() WHERE tenant_id=$1`, tenantID, result.ReservedMicros); err != nil { + return ReservationRelease{}, fmt.Errorf("release wallet hold: %w", err) + } + if err := tx.QueryRow(ctx, `UPDATE billing_reservations SET status='released',settled_at=now() + WHERE request_id=$1 RETURNING status,settled_at`, requestID).Scan(&result.Status, &result.ReleasedAt); err != nil { + return ReservationRelease{}, fmt.Errorf("mark reservation released: %w", err) + } + command, err := tx.Exec(ctx, `UPDATE usage_events SET metering_status='released_unmetered' + WHERE request_id=$1 AND tenant_id=$2 AND metering_status='missing' AND usage_reported=FALSE`, requestID, tenantID) + if err != nil { + return ReservationRelease{}, fmt.Errorf("mark unmetered usage resolved: %w", err) + } + if command.RowsAffected() != 1 { + return ReservationRelease{}, errors.New("unmetered usage invariant violated during release") + } + description := fmt.Sprintf("Unmetered reservation released by %s %s: %s", actor.Type, actor.ID, reason) + if _, err := tx.Exec(ctx, `INSERT INTO billing_ledger + (tenant_id,project_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description) + VALUES ($1,$2,$3,0,$4,'release','unmetered_reservation',$5,$6) + ON CONFLICT (source_type,source_id) DO NOTHING`, tenantID, projectID, currency, balance, requestID, description); err != nil { + return ReservationRelease{}, fmt.Errorf("write reservation release evidence: %w", err) + } + if err := tx.Commit(ctx); err != nil { + return ReservationRelease{}, fmt.Errorf("commit reservation release: %w", err) + } + return result, nil +} + func (s *Service) createTopUpOrder(ctx context.Context, input CheckoutInput) (string, int64, error) { if !s.stripeEnabled { return "", 0, ErrStripeDisabled diff --git a/internal/billing/operations.go b/internal/billing/operations.go index 461c59a..508ebd5 100644 --- a/internal/billing/operations.go +++ b/internal/billing/operations.go @@ -15,8 +15,10 @@ import ( "github.com/stripe/stripe-go/v86" ) +type stripePortalSessionCreator func(context.Context, *stripe.BillingPortalSessionCreateParams) (*stripe.BillingPortalSession, error) + func (s *Service) CreatePortalSession(ctx context.Context, tenantID string) (PortalResult, error) { - if !s.stripeEnabled || s.stripeClient == nil { + if !s.stripeEnabled || s.createStripePortalSession == nil { return PortalResult{}, ErrStripeDisabled } customerID, err := s.ensureStripeCustomer(ctx, tenantID) @@ -26,13 +28,13 @@ func (s *Service) CreatePortalSession(ctx context.Context, tenantID string) (Por if customerID == "" { return PortalResult{}, errors.New("no Stripe customer exists for this account") } - session, err := s.stripeClient.V1BillingPortalSessions.Create(ctx, &stripe.BillingPortalSessionCreateParams{ + session, err := s.createStripePortalSession(ctx, &stripe.BillingPortalSessionCreateParams{ Customer: stripe.String(customerID), ReturnURL: stripe.String(s.stripePortalReturnURL), }) if err != nil { return PortalResult{}, fmt.Errorf("create Stripe customer portal session: %w", err) } - if session.URL == "" { + if session == nil || session.URL == "" { return PortalResult{}, errors.New("Stripe returned an incomplete portal session") } return PortalResult{URL: session.URL}, nil diff --git a/internal/billing/operations_test.go b/internal/billing/operations_test.go index 46bedeb..4d90383 100644 --- a/internal/billing/operations_test.go +++ b/internal/billing/operations_test.go @@ -1,8 +1,15 @@ package billing import ( + "context" + "fmt" + "os" "testing" "time" + + "aigw/internal/controlplane" + + "github.com/stripe/stripe-go/v86" ) func TestOperationalStatusReadiness(t *testing.T) { @@ -34,3 +41,54 @@ func TestOperationalStatusReadiness(t *testing.T) { t.Fatal("unmetered success must fail readiness") } } + +func TestCustomerPortalSessionContractPostgres(t *testing.T) { + databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL") + if databaseURL == "" { + t.Skip("AIGW_TEST_DATABASE_URL is not set") + } + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + if err := controlplane.MigrateDatabase(ctx, databaseURL); err != nil { + t.Fatal(err) + } + service, err := New(ctx, Options{ + DatabaseURL: databaseURL, Currency: "usd", StripeEnabled: true, StripeAPIKey: "rk_test_placeholder", + StripePortalReturnURL: "https://console.example.test/billing", + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(service.Close) + + suffix := time.Now().UnixNano() + var tenantID string + if err := service.db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Portal contract') RETURNING id::text`, fmt.Sprintf("portal-%d", suffix)).Scan(&tenantID); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + if _, cleanupErr := service.db.Exec(context.Background(), `DELETE FROM tenants WHERE id=$1`, tenantID); cleanupErr != nil { + t.Errorf("cleanup portal contract tenant: %v", cleanupErr) + } + }) + customerID := fmt.Sprintf("cus_portal_%d", suffix) + if _, err := service.db.Exec(ctx, `INSERT INTO stripe_customers (tenant_id,stripe_customer_id,email) VALUES ($1,$2,'billing@example.test')`, tenantID, customerID); err != nil { + t.Fatal(err) + } + service.createStripePortalSession = func(_ context.Context, params *stripe.BillingPortalSessionCreateParams) (*stripe.BillingPortalSession, error) { + if params.Customer == nil || *params.Customer != customerID || params.ReturnURL == nil || *params.ReturnURL != service.stripePortalReturnURL { + t.Fatalf("unexpected Portal params %+v", params) + } + return &stripe.BillingPortalSession{URL: "https://billing.stripe.test/session"}, nil + } + result, err := service.CreatePortalSession(ctx, tenantID) + if err != nil || result.URL != "https://billing.stripe.test/session" { + t.Fatalf("CreatePortalSession result=%+v err=%v", result, err) + } + service.createStripePortalSession = func(context.Context, *stripe.BillingPortalSessionCreateParams) (*stripe.BillingPortalSession, error) { + return nil, nil + } + if _, err := service.CreatePortalSession(ctx, tenantID); err == nil { + t.Fatal("incomplete Stripe Portal response was accepted") + } +} diff --git a/internal/billing/service.go b/internal/billing/service.go index 8a91839..8a5f266 100644 --- a/internal/billing/service.go +++ b/internal/billing/service.go @@ -39,6 +39,7 @@ type Service struct { stripeProductTaxCode string integrationIdentifier string createStripeCheckout stripeCheckoutCreator + createStripePortalSession stripePortalSessionCreator createStripeCustomer stripeCustomerCreator updateStripeCustomer stripeCustomerUpdater retrieveStripeSetupIntent stripeSetupIntentRetriever @@ -73,6 +74,7 @@ func New(ctx context.Context, options Options) (*Service, error) { if options.StripeEnabled { service.stripeClient = stripe.NewClient(options.StripeAPIKey) service.createStripeCheckout = service.stripeClient.V1CheckoutSessions.Create + service.createStripePortalSession = service.stripeClient.V1BillingPortalSessions.Create service.createStripeCustomer = service.stripeClient.V1Customers.Create service.updateStripeCustomer = service.stripeClient.V1Customers.Update service.retrieveStripeSetupIntent = service.stripeClient.V1SetupIntents.Retrieve @@ -103,7 +105,7 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error { if input.Model.PriceCurrency != "" && input.Model.PriceCurrency != s.currency { return fmt.Errorf("model price currency %s does not match wallet currency %s", input.Model.PriceCurrency, s.currency) } - reserved, err := reservationCost(input.Model, input.Body, s.defaultMaxOutputTokens) + reserved, err := reservationCost(input.Model, input.Body, s.defaultMaxOutputTokens, input.Protocol) if err != nil { return err } @@ -156,6 +158,22 @@ func (s *Service) Authorize(ctx context.Context, input Authorization) error { return ErrQuotaExceeded } } + if input.Principal.DailySpendMicros > 0 { + now := time.Now().UTC() + period := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.UTC) + nextPeriod := period.AddDate(0, 0, 1) + var used, pending int64 + if err := tx.QueryRow(ctx, `SELECT + COALESCE((SELECT sum(cost_micros) FROM usage_events WHERE key_id=$1 AND started_at >= $2 AND started_at < $3),0), + COALESCE((SELECT sum(reserved_micros) FROM billing_reservations WHERE key_id=$1 AND status IN ('pending','metering_failed') AND created_at >= $2 AND created_at < $3),0)`, + input.Principal.KeyID, period, nextPeriod).Scan(&used, &pending); err != nil { + return fmt.Errorf("read API key daily spend quota: %w", err) + } + limit := input.Principal.DailySpendMicros + if reserved > limit || used > limit-reserved || pending > limit-used-reserved { + return ErrDailyQuotaExceeded + } + } if balance-held < reserved { return ErrInsufficientBalance } @@ -250,15 +268,15 @@ func (s *Service) settle(ctx context.Context, event domain.UsageEvent) error { } if _, err := tx.Exec(ctx, `INSERT INTO usage_events ( request_id,tenant_id,project_id,key_id,public_model,provider_id,upstream_model,protocol,stream, - status_code,success,error_type,attempts,started_at,duration_ms,input_tokens,output_tokens,total_tokens, + status_code,success,error_type,attempts,started_at,duration_ms,ttft_ms,input_tokens,output_tokens,total_tokens, cache_creation_input_tokens,cache_read_input_tokens,cost_micros,charged_micros,uncollected_micros, usage_reported,metering_status) VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,'usage_not_reported',$12,$13,$14, - $15,$16,$17,$18,$19,0,0,0,false,'missing') - ON CONFLICT (request_id) DO UPDATE SET error_type='usage_not_reported',usage_reported=false,metering_status='missing'`, + $15,$16,$17,$18,$19,$20,0,0,0,false,'missing') + ON CONFLICT (request_id) DO UPDATE SET error_type='usage_not_reported',ttft_ms=EXCLUDED.ttft_ms,usage_reported=false,metering_status='missing'`, event.RequestID, tenantID, projectID, keyID, modelID, event.ProviderID, event.UpstreamModel, string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.Attempts, event.StartedAt, - event.DurationMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, + event.DurationMS, event.TTFTMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens); err != nil { return fmt.Errorf("persist unmetered usage event: %w", err) } @@ -316,21 +334,21 @@ func (s *Service) settle(ctx context.Context, event domain.UsageEvent) error { } if usageAlreadyRecorded { if _, err := tx.Exec(ctx, `UPDATE usage_events SET cost_micros=$2, charged_micros=$3, uncollected_micros=$4, - usage_reported=$5,metering_status=$6 WHERE request_id=$1`, - event.RequestID, actualCost, charged, uncollected, event.UsageReported, meteringStatus(event)); err != nil { + ttft_ms=GREATEST(ttft_ms,$5),usage_reported=$6,metering_status=$7 WHERE request_id=$1`, + event.RequestID, actualCost, charged, uncollected, event.TTFTMS, event.UsageReported, meteringStatus(event)); err != nil { return fmt.Errorf("apply usage charge: %w", err) } } else if _, err := tx.Exec(ctx, ` INSERT INTO usage_events ( request_id, tenant_id, project_id, key_id, public_model, provider_id, upstream_model, protocol, stream, status_code, success, error_type, attempts, started_at, duration_ms, - input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, + ttft_ms, input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, cost_micros, charged_micros, uncollected_micros, usage_reported, metering_status) - VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25) + VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,$22,$23,$24,$25,$26) ON CONFLICT (request_id) DO NOTHING`, event.RequestID, tenantID, projectID, keyID, modelID, event.ProviderID, event.UpstreamModel, string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.ErrorType, event.Attempts, - event.StartedAt, event.DurationMS, event.Usage.InputTokens, event.Usage.OutputTokens, + event.StartedAt, event.DurationMS, event.TTFTMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens, actualCost, charged, uncollected, event.UsageReported, meteringStatus(event)); err != nil { return fmt.Errorf("persist usage event: %w", err) @@ -582,25 +600,32 @@ func meteringStatus(event domain.UsageEvent) string { return "missing" } -func reservationCost(model domain.Model, body []byte, defaultMaxOutput int64) (int64, error) { - maxOutput := defaultMaxOutput - if model.MaxOutputTokens > 0 && (maxOutput == 0 || model.MaxOutputTokens < maxOutput) { - maxOutput = model.MaxOutputTokens +func reservationCost(model domain.Model, body []byte, defaultMaxOutput int64, protocols ...domain.Protocol) (int64, error) { + protocol := domain.ProtocolOpenAI + if len(protocols) > 0 && protocols[0] != "" { + protocol = protocols[0] } + maxOutput := int64(0) var limits struct { MaxTokens int64 `json:"max_tokens"` MaxCompletionTokens int64 `json:"max_completion_tokens"` MaxOutputTokens int64 `json:"max_output_tokens"` } - if json.Unmarshal(body, &limits) == nil { - explicitMax := int64(0) - for _, value := range []int64{limits.MaxTokens, limits.MaxCompletionTokens, limits.MaxOutputTokens} { - if value > explicitMax { - explicitMax = value - } + if protocol != domain.ProtocolOpenAIEmbeddings { + maxOutput = defaultMaxOutput + if model.MaxOutputTokens > 0 && (maxOutput == 0 || model.MaxOutputTokens < maxOutput) { + maxOutput = model.MaxOutputTokens } - if explicitMax > 0 { - maxOutput = explicitMax + if json.Unmarshal(body, &limits) == nil { + explicitMax := int64(0) + for _, value := range []int64{limits.MaxTokens, limits.MaxCompletionTokens, limits.MaxOutputTokens} { + if value > explicitMax { + explicitMax = value + } + } + if explicitMax > 0 { + maxOutput = explicitMax + } } } cacheReservePrice := model.CacheReadPriceMicrosPerMillion @@ -618,21 +643,51 @@ func usageCost(usage domain.Usage, inputPrice, outputPrice, cacheReadPrice, cach } func calculateCost(input, output, cacheRead, cacheWrite, inputPrice, outputPrice, cacheReadPrice, cacheWritePrice int64) (int64, error) { - values := []int64{input, output, cacheRead, cacheWrite, inputPrice, outputPrice, cacheReadPrice, cacheWritePrice} - for _, value := range values { - if value < 0 { - return 0, errors.New("billing values cannot be negative") + return calculateMeteredCost([]meteredCharge{ + {Unit: domain.MeteringUnitToken, Quantity: input, PriceMicros: inputPrice, PerQuantity: microsPerUnit}, + {Unit: domain.MeteringUnitToken, Quantity: output, PriceMicros: outputPrice, PerQuantity: microsPerUnit}, + {Unit: domain.MeteringUnitToken, Quantity: cacheRead, PriceMicros: cacheReadPrice, PerQuantity: microsPerUnit}, + {Unit: domain.MeteringUnitToken, Quantity: cacheWrite, PriceMicros: cacheWritePrice, PerQuantity: microsPerUnit}, + }) +} + +type meteredCharge struct { + Unit domain.MeteringUnit + Quantity int64 + PriceMicros int64 + PerQuantity int64 +} + +// calculateMeteredCost is the common fixed-point primitive for token, image, +// and duration pricing. Token rates use PerQuantity=1_000_000; image and second +// rates can use PerQuantity=1 without changing wallet or ledger arithmetic. +func calculateMeteredCost(charges []meteredCharge) (int64, error) { + byScale := make(map[int64]*big.Int) + for _, charge := range charges { + if charge.Unit != domain.MeteringUnitToken && charge.Unit != domain.MeteringUnitImage && charge.Unit != domain.MeteringUnitSecond { + return 0, fmt.Errorf("unsupported metering unit %q", charge.Unit) + } + if charge.Quantity < 0 || charge.PriceMicros < 0 || charge.PerQuantity <= 0 { + return 0, errors.New("metering quantity, price, or scale is invalid") + } + if charge.Quantity == 0 || charge.PriceMicros == 0 { + continue + } + component := new(big.Int).Mul(big.NewInt(charge.Quantity), big.NewInt(charge.PriceMicros)) + if byScale[charge.PerQuantity] == nil { + byScale[charge.PerQuantity] = new(big.Int) } + byScale[charge.PerQuantity].Add(byScale[charge.PerQuantity], component) } total := new(big.Int) - for _, pair := range [][2]int64{{input, inputPrice}, {output, outputPrice}, {cacheRead, cacheReadPrice}, {cacheWrite, cacheWritePrice}} { - total.Add(total, new(big.Int).Mul(big.NewInt(pair[0]), big.NewInt(pair[1]))) + for scale, numerator := range byScale { + numerator.Add(numerator, big.NewInt(scale-1)) + numerator.Div(numerator, big.NewInt(scale)) + total.Add(total, numerator) } if total.Sign() == 0 { return 0, nil } - total.Add(total, big.NewInt(microsPerUnit-1)) - total.Div(total, big.NewInt(microsPerUnit)) if !total.IsInt64() { return 0, errors.New("calculated charge exceeds supported range") } diff --git a/internal/billing/service_test.go b/internal/billing/service_test.go index 5e21b20..a75e4e0 100644 --- a/internal/billing/service_test.go +++ b/internal/billing/service_test.go @@ -84,6 +84,16 @@ func TestAuthorizeEnforcesAPIKeyMonthlySpendCapPostgres(t *testing.T) { if err := service.Authorize(ctx, Authorization{RequestID: "req_key_budget_allowed", Principal: principal, Model: model, Body: []byte(`{"max_tokens":10}`)}); err != nil { t.Fatalf("Authorize at exact cap: %v", err) } + principal.MonthlySpendMicros = 0 + principal.DailySpendMicros = 19 + err = service.Authorize(ctx, Authorization{RequestID: "req_key_budget_daily_rejected", Principal: principal, Model: model, Body: []byte(`{"max_tokens":10}`)}) + if !errors.Is(err, ErrDailyQuotaExceeded) { + t.Fatalf("Authorize daily error = %v, want ErrDailyQuotaExceeded", err) + } + principal.DailySpendMicros = 20 + if err := service.Authorize(ctx, Authorization{RequestID: "req_key_budget_daily_allowed", Principal: principal, Model: model, Body: []byte(`{"max_tokens":10}`)}); err != nil { + t.Fatalf("Authorize at exact daily cap including pending reservation: %v", err) + } } func TestUsageCostUsesFixedPointAndRoundsOnce(t *testing.T) { @@ -111,6 +121,35 @@ func TestReservationUsesExplicitOutputLimit(t *testing.T) { } } +func TestEmbeddingReservationDoesNotReserveOutputTokens(t *testing.T) { + model := domain.Model{InputPriceMicrosPerMillion: 1_000_000, OutputPriceMicrosPerMillion: 50_000_000, MaxOutputTokens: 8192} + body := []byte(`{"model":"embedding","input":"hello","max_tokens":99999}`) + cost, err := reservationCost(model, body, 4096, domain.ProtocolOpenAIEmbeddings) + if err != nil { + t.Fatal(err) + } + if cost != int64(len(body)) { + t.Fatalf("embedding reservation = %d, want conservative input-only %d", cost, len(body)) + } +} + +func TestMeteredCostSupportsTokenImageAndSecondUnits(t *testing.T) { + cost, err := calculateMeteredCost([]meteredCharge{ + {Unit: domain.MeteringUnitToken, Quantity: 500_000, PriceMicros: 2_000_000, PerQuantity: 1_000_000}, + {Unit: domain.MeteringUnitImage, Quantity: 2, PriceMicros: 40_000, PerQuantity: 1}, + {Unit: domain.MeteringUnitSecond, Quantity: 3, PriceMicros: 500, PerQuantity: 1}, + }) + if err != nil { + t.Fatal(err) + } + if cost != 1_081_500 { + t.Fatalf("metered cost = %d, want 1081500", cost) + } + if _, err := calculateMeteredCost([]meteredCharge{{Unit: "byte", Quantity: 1, PriceMicros: 1, PerQuantity: 1}}); err == nil { + t.Fatal("unsupported metering unit was accepted") + } +} + func TestMinorToMicrosSupportsCurrencyExponents(t *testing.T) { tests := []struct { currency string @@ -153,6 +192,57 @@ func TestIntegrationIdentifierSuffixUsesLetters(t *testing.T) { } } +func TestStripeSDKVersionAndCheckoutContract(t *testing.T) { + if stripe.APIVersion != "2026-07-29.dahlia" { + t.Fatalf("Stripe API version = %q; review the integration before changing the pinned version", stripe.APIVersion) + } + service := &Service{ + currency: "usd", + stripeSuccessURL: "https://console.example.test/billing?topup=success", + stripeCancelURL: "https://console.example.test/billing?topup=cancel", + integrationIdentifier: "aigw_balance_abcdefgh", + } + params := service.checkoutSessionParams("order-123", CheckoutInput{TenantID: "tenant-123", AmountMinor: 2500}) + if params.Mode == nil || *params.Mode != string(stripe.CheckoutSessionModePayment) { + t.Fatalf("mode = %v", params.Mode) + } + if params.IntegrationIdentifier == nil || *params.IntegrationIdentifier != "aigw_balance_abcdefgh" { + t.Fatalf("integration identifier = %v", params.IntegrationIdentifier) + } + if len(params.PaymentMethodTypes) != 0 || len(params.ExcludedPaymentMethodTypes) != 0 { + t.Fatal("Checkout must use Dashboard-managed dynamic payment methods") + } + if params.AutomaticTax != nil || params.TaxIDCollection != nil { + t.Fatal("Stripe Tax must remain disabled unless registration is explicitly confirmed") + } + if params.InvoiceCreation == nil || params.InvoiceCreation.Enabled == nil || !*params.InvoiceCreation.Enabled { + t.Fatal("one-time prepaid top-up invoice creation is not enabled") + } + if len(params.LineItems) != 1 || params.LineItems[0].PriceData == nil || params.LineItems[0].PriceData.UnitAmount == nil || *params.LineItems[0].PriceData.UnitAmount != 2500 { + t.Fatalf("unexpected line item %+v", params.LineItems) + } + if params.IdempotencyKey == nil || *params.IdempotencyKey != "aigw_topup_order-123" { + t.Fatalf("idempotency key = %v", params.IdempotencyKey) + } +} + +func TestCheckoutContractEnablesTaxOnlyWhenExplicitlyConfigured(t *testing.T) { + service := &Service{ + currency: "usd", + stripeAutomaticTax: true, + stripeProductTaxCode: "txcd_10103000", + integrationIdentifier: "aigw_balance_abcdefgh", + } + params := service.checkoutSessionParams("order-tax", CheckoutInput{TenantID: "tenant-tax", AmountMinor: 1000}) + if params.AutomaticTax == nil || params.AutomaticTax.Enabled == nil || !*params.AutomaticTax.Enabled || + params.TaxIDCollection == nil || params.TaxIDCollection.Enabled == nil || !*params.TaxIDCollection.Enabled { + t.Fatal("explicit Stripe Tax configuration was not applied") + } + if params.LineItems[0].PriceData.ProductData.TaxCode == nil || *params.LineItems[0].PriceData.ProductData.TaxCode != "txcd_10103000" { + t.Fatal("canonical Stripe product tax code was not applied") + } +} + func TestCheckoutReturnURLPreservesCallbackAndSessionPlaceholder(t *testing.T) { success := checkoutReturnURL("https://console.example.test/admin/?topup=success", "order-123", true) if !strings.Contains(success, "topup=success") || !strings.Contains(success, "order_id=order-123") || @@ -196,6 +286,62 @@ func TestWebhookRejectsInvalidSignatureBeforeProcessing(t *testing.T) { } } +func TestIncompleteStripeCheckoutMarksOrderFailedPostgres(t *testing.T) { + databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL") + if databaseURL == "" { + t.Skip("AIGW_TEST_DATABASE_URL is not set") + } + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + if err := controlplane.MigrateDatabase(ctx, databaseURL); err != nil { + t.Fatal(err) + } + service, err := New(ctx, Options{ + DatabaseURL: databaseURL, Currency: "usd", MinTopUpMinor: 500, MaxTopUpMinor: 1_000_000, + StripeEnabled: true, StripeAPIKey: "rk_test_placeholder", + StripeSuccessURL: "https://console.example.test/billing?topup=success", + StripeCancelURL: "https://console.example.test/billing?topup=cancel", + }) + if err != nil { + t.Fatal(err) + } + t.Cleanup(service.Close) + + suffix := time.Now().UnixNano() + var tenantID string + if err := service.db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Incomplete Stripe Checkout') RETURNING id::text`, fmt.Sprintf("incomplete-checkout-%d", suffix)).Scan(&tenantID); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + if _, cleanupErr := service.db.Exec(context.Background(), `DELETE FROM topup_orders WHERE tenant_id=$1`, tenantID); cleanupErr != nil { + t.Errorf("cleanup incomplete Checkout orders: %v", cleanupErr) + } + if _, cleanupErr := service.db.Exec(context.Background(), `DELETE FROM tenants WHERE id=$1`, tenantID); cleanupErr != nil { + t.Errorf("cleanup incomplete Checkout tenant: %v", cleanupErr) + } + }) + if _, err := service.db.Exec(ctx, `INSERT INTO stripe_customers (tenant_id,stripe_customer_id,email) VALUES ($1,$2,'billing@example.test')`, tenantID, fmt.Sprintf("cus_incomplete_%d", suffix)); err != nil { + t.Fatal(err) + } + service.createStripeCheckout = func(context.Context, *stripe.CheckoutSessionCreateParams) (*stripe.CheckoutSession, error) { + return nil, nil + } + + _, err = service.CreateCheckout(ctx, CheckoutInput{TenantID: tenantID, AmountMinor: 500}) + if err == nil || !strings.Contains(err.Error(), "incomplete Checkout Session") { + t.Fatalf("CreateCheckout error = %v", err) + } + var status, reconciliationStatus, reconciliationError string + if err := service.db.QueryRow(ctx, `SELECT status,reconciliation_status,reconciliation_error + FROM topup_orders WHERE tenant_id=$1 ORDER BY created_at DESC LIMIT 1`, tenantID). + Scan(&status, &reconciliationStatus, &reconciliationError); err != nil { + t.Fatal(err) + } + if status != "failed" || reconciliationStatus != "unknown" || !strings.Contains(reconciliationError, "incomplete Checkout Session") { + t.Fatalf("order state status=%q reconciliation=%q error=%q", status, reconciliationStatus, reconciliationError) + } +} + func TestStripeWebhookCreditsPaidOrderExactlyOncePostgres(t *testing.T) { databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL") if databaseURL == "" { @@ -434,4 +580,32 @@ func TestSettlementWorkerPersistsUsageAndReleasesReservationPostgres(t *testing. if missingReservationStatus != "metering_failed" || missingJobStatus != "done" || balance != 990 || reserved != 30 || meteringStatus != "missing" { t.Fatalf("missing usage reservation=%s job=%s balance=%d reserved=%d metering=%s", missingReservationStatus, missingJobStatus, balance, reserved, meteringStatus) } + release, err := service.ReleaseUnmeteredReservation(ctx, tenantID, missingRequestID, + ReleaseReservationInput{Reason: "provider returned a non-meterable success"}, ResolutionActor{ID: "test-operator", Type: "integration"}) + if err != nil { + t.Fatal(err) + } + if release.Status != "released" || release.ReservedMicros != 30 { + t.Fatalf("unexpected reservation release: %+v", release) + } + if err := service.db.QueryRow(ctx, `SELECT balance_micros,reserved_micros FROM tenant_wallets WHERE tenant_id=$1`, tenantID).Scan(&balance, &reserved); err != nil { + t.Fatal(err) + } + var releaseLedgerCount int + if err := service.db.QueryRow(ctx, `SELECT count(*) FROM billing_ledger WHERE source_type='unmetered_reservation' AND source_id=$1 AND kind='release' AND amount_micros=0`, missingRequestID).Scan(&releaseLedgerCount); err != nil { + t.Fatal(err) + } + if balance != 990 || reserved != 0 || releaseLedgerCount != 1 { + t.Fatalf("released wallet balance=%d reserved=%d ledger=%d", balance, reserved, releaseLedgerCount) + } + if err := service.db.QueryRow(ctx, `SELECT metering_status FROM usage_events WHERE request_id=$1`, missingRequestID).Scan(&meteringStatus); err != nil { + t.Fatal(err) + } + if meteringStatus != "released_unmetered" { + t.Fatalf("released metering status=%s", meteringStatus) + } + if _, err := service.ReleaseUnmeteredReservation(ctx, tenantID, missingRequestID, + ReleaseReservationInput{Reason: "idempotent retry"}, ResolutionActor{ID: "test-operator", Type: "integration"}); err != nil { + t.Fatalf("idempotent release retry: %v", err) + } } diff --git a/internal/billing/stripe.go b/internal/billing/stripe.go index 032eb02..ea4b7ea 100644 --- a/internal/billing/stripe.go +++ b/internal/billing/stripe.go @@ -29,6 +29,35 @@ func (s *Service) CreateCheckout(ctx context.Context, input CheckoutInput) (Chec if err != nil { return CheckoutResult{}, err } + params := s.checkoutSessionParams(orderID, input) + customerID, err := s.ensureStripeCustomer(ctx, input.TenantID) + if err != nil { + return CheckoutResult{}, s.failCheckoutCreation(ctx, orderID, err) + } + if customerID != "" { + params.Customer = stripe.String(customerID) + } else { + params.CustomerCreation = stripe.String(string(stripe.CheckoutSessionCustomerCreationAlways)) + if strings.TrimSpace(input.CustomerEmail) != "" { + params.CustomerEmail = stripe.String(strings.TrimSpace(input.CustomerEmail)) + } + } + session, err := s.createStripeCheckout(ctx, params) + if err != nil { + return CheckoutResult{}, s.failCheckoutCreation(ctx, orderID, fmt.Errorf("create Stripe Checkout Session: %w", err)) + } + if session == nil || session.ID == "" || session.URL == "" { + return CheckoutResult{}, s.failCheckoutCreation(ctx, orderID, errors.New("Stripe returned an incomplete Checkout Session")) + } + if _, err := s.db.Exec(ctx, ` + UPDATE topup_orders SET stripe_session_id = $2, checkout_url = $3 + WHERE id = $1`, orderID, session.ID, session.URL); err != nil { + return CheckoutResult{}, s.failCheckoutCreation(ctx, orderID, fmt.Errorf("persist Stripe Checkout Session: %w", err)) + } + return CheckoutResult{OrderID: orderID, SessionID: session.ID, URL: session.URL}, nil +} + +func (s *Service) checkoutSessionParams(orderID string, input CheckoutInput) *stripe.CheckoutSessionCreateParams { params := &stripe.CheckoutSessionCreateParams{ Mode: stripe.String("payment"), ClientReferenceID: stripe.String(orderID), @@ -58,43 +87,23 @@ func (s *Service) CreateCheckout(ctx context.Context, input CheckoutInput) (Chec }, }}, } - customerID, err := s.ensureStripeCustomer(ctx, input.TenantID) - if err != nil { - if _, updateErr := s.db.Exec(ctx, `UPDATE topup_orders SET status = 'failed' WHERE id = $1 AND status = 'pending'`, orderID); updateErr != nil { - return CheckoutResult{}, errors.Join(err, fmt.Errorf("mark top-up order failed: %w", updateErr)) - } - return CheckoutResult{}, err - } - if customerID != "" { - params.Customer = stripe.String(customerID) - } else { - params.CustomerCreation = stripe.String(string(stripe.CheckoutSessionCustomerCreationAlways)) - if strings.TrimSpace(input.CustomerEmail) != "" { - params.CustomerEmail = stripe.String(strings.TrimSpace(input.CustomerEmail)) - } - } if s.stripeAutomaticTax { params.AutomaticTax = &stripe.CheckoutSessionCreateAutomaticTaxParams{Enabled: stripe.Bool(true)} params.TaxIDCollection = &stripe.CheckoutSessionCreateTaxIDCollectionParams{Enabled: stripe.Bool(true)} params.LineItems[0].PriceData.ProductData.TaxCode = stripe.String(s.stripeProductTaxCode) } params.SetIdempotencyKey("aigw_topup_" + orderID) - session, err := s.createStripeCheckout(ctx, params) - if err != nil { - if _, updateErr := s.db.Exec(ctx, `UPDATE topup_orders SET status = 'failed' WHERE id = $1 AND status = 'pending'`, orderID); updateErr != nil { - return CheckoutResult{}, errors.Join(fmt.Errorf("create Stripe Checkout Session: %w", err), fmt.Errorf("mark top-up order failed: %w", updateErr)) - } - return CheckoutResult{}, fmt.Errorf("create Stripe Checkout Session: %w", err) - } - if session.ID == "" || session.URL == "" { - return CheckoutResult{}, errors.New("Stripe returned an incomplete Checkout Session") - } - if _, err := s.db.Exec(ctx, ` - UPDATE topup_orders SET stripe_session_id = $2, checkout_url = $3 - WHERE id = $1`, orderID, session.ID, session.URL); err != nil { - return CheckoutResult{}, fmt.Errorf("persist Stripe Checkout Session: %w", err) - } - return CheckoutResult{OrderID: orderID, SessionID: session.ID, URL: session.URL}, nil + return params +} + +func (s *Service) failCheckoutCreation(ctx context.Context, orderID string, cause error) error { + _, updateErr := s.db.Exec(ctx, `UPDATE topup_orders + SET status='failed', reconciliation_status='unknown', reconciliation_error=$2 + WHERE id=$1 AND status='pending'`, orderID, truncateError(cause)) + if updateErr != nil { + return errors.Join(cause, fmt.Errorf("mark top-up order failed: %w", updateErr)) + } + return cause } func checkoutReturnURL(raw, orderID string, includeStripeSession bool) string { diff --git a/internal/billing/stripe_preflight.go b/internal/billing/stripe_preflight.go new file mode 100644 index 0000000..4f9250c --- /dev/null +++ b/internal/billing/stripe_preflight.go @@ -0,0 +1,124 @@ +package billing + +import ( + "context" + "encoding/json" + "errors" + "fmt" + "io" + "net/http" + "net/url" + "strings" + + "github.com/stripe/stripe-go/v86" +) + +const ( + stripeAPIBaseURL = "https://api.stripe.com" + maxStripeErrorBodySize = 64 << 10 +) + +var ErrLiveStripeKey = errors.New("Stripe permission preflight only accepts test-mode keys") + +type StripePermissionCheck struct { + Name string `json:"name"` + OK bool `json:"ok"` + StatusCode int `json:"status_code"` + ErrorCode string `json:"error_code,omitempty"` + ErrorType string `json:"error_type,omitempty"` +} + +type StripePreflightResult struct { + APIVersion string `json:"api_version"` + TestMode bool `json:"test_mode"` + Ready bool `json:"ready"` + Checks []StripePermissionCheck `json:"checks"` +} + +type stripeReadCheck struct { + name string + path string +} + +var stripeRequiredReadChecks = []stripeReadCheck{ + {name: "customers_read", path: "/v1/customers"}, + {name: "checkout_sessions_read", path: "/v1/checkout/sessions"}, + {name: "setup_intents_read", path: "/v1/setup_intents"}, + {name: "payment_intents_read", path: "/v1/payment_intents"}, + {name: "refunds_read", path: "/v1/refunds"}, + {name: "charges_read", path: "/v1/charges"}, + {name: "disputes_read", path: "/v1/disputes"}, + {name: "invoices_read", path: "/v1/invoices"}, + {name: "billing_portal_configurations_read", path: "/v1/billing_portal/configurations"}, +} + +// CheckStripePermissions validates the read side of the restricted-key contract +// without creating Stripe objects. Write permissions are exercised by the +// sandbox Checkout, Portal, automatic top-up, refund, and reconciliation flows. +func CheckStripePermissions(ctx context.Context, apiKey string) (StripePreflightResult, error) { + return checkStripePermissions(ctx, apiKey, stripeAPIBaseURL, http.DefaultClient) +} + +func checkStripePermissions(ctx context.Context, apiKey, baseURL string, client *http.Client) (StripePreflightResult, error) { + apiKey = strings.TrimSpace(apiKey) + result := StripePreflightResult{ + APIVersion: stripe.APIVersion, + TestMode: isStripeTestKey(apiKey), + Checks: make([]StripePermissionCheck, 0, len(stripeRequiredReadChecks)), + } + if !result.TestMode { + return result, ErrLiveStripeKey + } + result.Ready = true + if client == nil { + client = http.DefaultClient + } + for _, check := range stripeRequiredReadChecks { + item := runStripeReadCheck(ctx, client, apiKey, baseURL, check) + result.Checks = append(result.Checks, item) + result.Ready = result.Ready && item.OK + } + return result, nil +} + +func runStripeReadCheck(ctx context.Context, client *http.Client, apiKey, baseURL string, check stripeReadCheck) StripePermissionCheck { + endpoint, err := url.JoinPath(baseURL, check.path) + if err != nil { + return StripePermissionCheck{Name: check.name, ErrorType: "configuration_error"} + } + query := url.Values{"limit": []string{"1"}} + request, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint+"?"+query.Encode(), nil) + if err != nil { + return StripePermissionCheck{Name: check.name, ErrorType: "configuration_error"} + } + request.Header.Set("Authorization", "Bearer "+apiKey) + request.Header.Set("Stripe-Version", stripe.APIVersion) + response, err := client.Do(request) + if err != nil { + return StripePermissionCheck{Name: check.name, ErrorType: "network_error"} + } + defer response.Body.Close() + item := StripePermissionCheck{Name: check.name, OK: response.StatusCode >= 200 && response.StatusCode < 300, StatusCode: response.StatusCode} + if item.OK { + _, _ = io.Copy(io.Discard, response.Body) + return item + } + var envelope struct { + Error struct { + Code string `json:"code"` + Type string `json:"type"` + } `json:"error"` + } + if err := json.NewDecoder(io.LimitReader(response.Body, maxStripeErrorBodySize)).Decode(&envelope); err == nil { + item.ErrorCode = envelope.Error.Code + item.ErrorType = envelope.Error.Type + } + if item.ErrorType == "" { + item.ErrorType = fmt.Sprintf("http_%d", response.StatusCode) + } + return item +} + +func isStripeTestKey(value string) bool { + return strings.HasPrefix(value, "rk_test_") || strings.HasPrefix(value, "sk_test_") +} diff --git a/internal/billing/stripe_preflight_test.go b/internal/billing/stripe_preflight_test.go new file mode 100644 index 0000000..e8f1841 --- /dev/null +++ b/internal/billing/stripe_preflight_test.go @@ -0,0 +1,90 @@ +package billing + +import ( + "context" + "io" + "net/http" + "strings" + "testing" + + "github.com/stripe/stripe-go/v86" +) + +func TestStripePermissionPreflightChecksRequiredResourcesWithoutLeakingKey(t *testing.T) { + const key = "rk_test_do_not_log_this_value" + seen := make(map[string]bool) + client := &http.Client{Transport: stripeRoundTripFunc(func(r *http.Request) (*http.Response, error) { + if r.Method != http.MethodGet || r.URL.Query().Get("limit") != "1" { + t.Errorf("unexpected request %s %s", r.Method, r.URL.String()) + } + if r.Header.Get("Authorization") != "Bearer "+key { + t.Errorf("missing Stripe bearer authentication") + } + if r.Header.Get("Stripe-Version") != stripe.APIVersion { + t.Errorf("Stripe-Version = %q", r.Header.Get("Stripe-Version")) + } + seen[r.URL.Path] = true + return stripeTestResponse(http.StatusOK, `{"object":"list","data":[]}`), nil + })} + + result, err := checkStripePermissions(context.Background(), key, "https://stripe.test", client) + if err != nil { + t.Fatal(err) + } + if !result.Ready || !result.TestMode || result.APIVersion != stripe.APIVersion { + t.Fatalf("unexpected result %+v", result) + } + if len(result.Checks) != len(stripeRequiredReadChecks) { + t.Fatalf("checks = %d, want %d", len(result.Checks), len(stripeRequiredReadChecks)) + } + for _, check := range stripeRequiredReadChecks { + if !seen[check.path] { + t.Errorf("endpoint %s was not checked", check.path) + } + } +} + +func TestStripePermissionPreflightReportsSanitizedStripeError(t *testing.T) { + client := &http.Client{Transport: stripeRoundTripFunc(func(*http.Request) (*http.Response, error) { + return stripeTestResponse(http.StatusForbidden, `{"error":{"type":"invalid_request_error","code":"permission_denied","message":"secret details"}}`), nil + })} + + result, err := checkStripePermissions(context.Background(), "rk_test_placeholder", "https://stripe.test", client) + if err != nil { + t.Fatal(err) + } + if result.Ready || len(result.Checks) == 0 { + t.Fatalf("unexpected result %+v", result) + } + for _, check := range result.Checks { + if check.OK || check.StatusCode != http.StatusForbidden || check.ErrorCode != "permission_denied" || check.ErrorType != "invalid_request_error" { + t.Fatalf("unexpected check %+v", check) + } + } +} + +type stripeRoundTripFunc func(*http.Request) (*http.Response, error) + +func (fn stripeRoundTripFunc) RoundTrip(request *http.Request) (*http.Response, error) { + return fn(request) +} + +func stripeTestResponse(status int, body string) *http.Response { + return &http.Response{ + StatusCode: status, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader(body)), + } +} + +func TestStripePermissionPreflightRejectsLiveAndMalformedKeys(t *testing.T) { + for _, key := range []string{"", "rk_live_forbidden", "not-a-stripe-key"} { + result, err := checkStripePermissions(context.Background(), key, "http://unused", nil) + if err != ErrLiveStripeKey { + t.Fatalf("key %q: error = %v, want ErrLiveStripeKey", key, err) + } + if result.Ready || result.Checks == nil || len(result.Checks) != 0 { + t.Fatalf("key %q: unexpected rejected result %+v", key, result) + } + } +} diff --git a/internal/billing/types.go b/internal/billing/types.go index 1633ecc..338eeb2 100644 --- a/internal/billing/types.go +++ b/internal/billing/types.go @@ -9,18 +9,20 @@ import ( ) var ( - ErrInsufficientBalance = errors.New("insufficient balance") - ErrStripeDisabled = errors.New("Stripe top-ups are disabled") - ErrInvalidAmount = errors.New("invalid amount") - ErrQuotaExceeded = errors.New("monthly spend quota exceeded") - ErrTopUpOrderNotFound = errors.New("top-up order not found") - ErrUsageNotReported = errors.New("billable successful response did not report usage") - ErrCannotResolveTopUp = errors.New("top-up order cannot be resolved as missing") - ErrBillingAccountNotFound = errors.New("billing account not found") - ErrPaymentMethodRequired = errors.New("a saved payment method is required") - ErrAutoTopUpNeedsAttention = errors.New("automatic top-up payment method requires attention") - ErrInvalidBillingProfile = errors.New("invalid billing profile") - ErrBillingProfileSync = errors.New("billing profile Stripe synchronization failed") + ErrInsufficientBalance = errors.New("insufficient balance") + ErrStripeDisabled = errors.New("Stripe top-ups are disabled") + ErrInvalidAmount = errors.New("invalid amount") + ErrQuotaExceeded = errors.New("monthly spend quota exceeded") + ErrDailyQuotaExceeded = errors.New("daily spend quota exceeded") + ErrTopUpOrderNotFound = errors.New("top-up order not found") + ErrUsageNotReported = errors.New("billable successful response did not report usage") + ErrCannotResolveTopUp = errors.New("top-up order cannot be resolved as missing") + ErrBillingAccountNotFound = errors.New("billing account not found") + ErrPaymentMethodRequired = errors.New("a saved payment method is required") + ErrAutoTopUpNeedsAttention = errors.New("automatic top-up payment method requires attention") + ErrInvalidBillingProfile = errors.New("invalid billing profile") + ErrBillingProfileSync = errors.New("billing profile Stripe synchronization failed") + ErrReservationNotReleasable = errors.New("billing reservation is not releasable") ) type Meter interface { @@ -32,6 +34,7 @@ type Authorization struct { RequestID string Principal domain.Principal Model domain.Model + Protocol domain.Protocol Body []byte Policy domain.LimitPolicy } @@ -127,6 +130,18 @@ type AdjustmentInput struct { Description string `json:"description"` } +type ReleaseReservationInput struct { + Reason string `json:"reason"` +} + +type ReservationRelease struct { + RequestID string `json:"request_id"` + TenantID string `json:"tenant_id"` + ReservedMicros int64 `json:"reserved_micros"` + Status string `json:"status"` + ReleasedAt time.Time `json:"released_at"` +} + type CheckoutInput struct { TenantID string `json:"tenant_id"` AmountMinor int64 `json:"amount_minor"` diff --git a/internal/catalog/catalog.go b/internal/catalog/catalog.go index 362d5b0..32024aa 100644 --- a/internal/catalog/catalog.go +++ b/internal/catalog/catalog.go @@ -42,6 +42,7 @@ func New(cfg config.Config) *Catalog { model := domain.Model{ ID: modelCfg.ID, OwnedBy: modelCfg.OwnedBy, + Capabilities: append([]string(nil), modelCfg.Capabilities...), InputPriceMicrosPerMillion: modelCfg.InputPriceMicrosPerMillion, OutputPriceMicrosPerMillion: modelCfg.OutputPriceMicrosPerMillion, CacheReadPriceMicrosPerMillion: modelCfg.CacheReadPriceMicrosPerMillion, @@ -139,12 +140,30 @@ func (c *Catalog) Models(protocol domain.Protocol) []domain.Model { return result } +// AllModels returns a detached view of the current atomic catalog snapshot. +// It is intended for control-loop work such as active provider probes; request +// routing should continue to use ModelForPrincipal and ModelsFor. +func (c *Catalog) AllModels() []domain.Model { + current := c.state.Load() + if current == nil { + return nil + } + result := make([]domain.Model, len(current.list)) + for index, model := range current.list { + result[index] = model + result[index].Routes = append([]domain.Route(nil), model.Routes...) + } + return result +} + func protocolCompatible(provider domain.Provider, requestProtocol domain.Protocol) bool { switch requestProtocol { case domain.ProtocolOpenAI: return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "chat_completions" case domain.ProtocolOpenAIResponses: return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "responses" + case domain.ProtocolOpenAIEmbeddings: + return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "embeddings" case domain.ProtocolAnthropic: return provider.Protocol == domain.ProtocolAnthropic && provider.EffectiveWireAPI() == "messages" default: diff --git a/internal/config/config.go b/internal/config/config.go index 047f372..2df74cf 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -15,15 +15,16 @@ import ( ) type Config struct { - Server ServerConfig `json:"server"` - Auth AuthConfig `json:"auth"` - ControlPlane ControlPlaneConfig `json:"control_plane"` - Admin AdminConfig `json:"admin"` - UpstreamHTTP UpstreamHTTPConfig `json:"upstream_http"` - Providers []ProviderConfig `json:"providers"` - Models []ModelConfig `json:"models"` - Billing BillingConfig `json:"billing"` - Observability ObservabilityConfig `json:"observability"` + Server ServerConfig `json:"server"` + Auth AuthConfig `json:"auth"` + ControlPlane ControlPlaneConfig `json:"control_plane"` + Admin AdminConfig `json:"admin"` + UpstreamHTTP UpstreamHTTPConfig `json:"upstream_http"` + ProviderHealth ProviderHealthConfig `json:"provider_health"` + Providers []ProviderConfig `json:"providers"` + Models []ModelConfig `json:"models"` + Billing BillingConfig `json:"billing"` + Observability ObservabilityConfig `json:"observability"` } type ServerConfig struct { @@ -125,6 +126,18 @@ type UpstreamHTTPConfig struct { ResponseHeaderTimeoutSecs int `json:"response_header_timeout_seconds"` } +type ProviderHealthConfig struct { + ActiveProbesEnabledEnv string `json:"active_probes_enabled_env"` + ActiveProbesEnabled bool `json:"-"` + SharedHistoryEnabledEnv string `json:"shared_history_enabled_env"` + SharedHistoryEnabled bool `json:"-"` + SharedHistoryStream string `json:"shared_history_stream"` + SharedHistoryTTLSeconds int `json:"shared_history_ttl_seconds"` + SharedHistoryMaxEvents int64 `json:"shared_history_max_events"` + ProbeIntervalSeconds int `json:"probe_interval_seconds"` + ProbeTimeoutSeconds int `json:"probe_timeout_seconds"` +} + type ProviderConfig struct { ID string `json:"id"` Slug string `json:"slug"` @@ -139,6 +152,7 @@ type ProviderConfig struct { type ModelConfig struct { ID string `json:"id"` OwnedBy string `json:"owned_by"` + Capabilities []string `json:"capabilities"` InputPriceMicrosPerMillion int64 `json:"input_price_micros_per_million"` OutputPriceMicrosPerMillion int64 `json:"output_price_micros_per_million"` CacheReadPriceMicrosPerMillion int64 `json:"cache_read_price_micros_per_million"` @@ -359,6 +373,27 @@ func applyDefaults(cfg *Config) { if cfg.UpstreamHTTP.ResponseHeaderTimeoutSecs == 0 { cfg.UpstreamHTTP.ResponseHeaderTimeoutSecs = 60 } + if cfg.ProviderHealth.ActiveProbesEnabledEnv == "" { + cfg.ProviderHealth.ActiveProbesEnabledEnv = "AIGW_PROVIDER_ACTIVE_PROBES_ENABLED" + } + if cfg.ProviderHealth.SharedHistoryEnabledEnv == "" { + cfg.ProviderHealth.SharedHistoryEnabledEnv = "AIGW_PROVIDER_SHARED_HISTORY_ENABLED" + } + if cfg.ProviderHealth.SharedHistoryStream == "" { + cfg.ProviderHealth.SharedHistoryStream = "aigw:provider-health:events" + } + if cfg.ProviderHealth.SharedHistoryTTLSeconds == 0 { + cfg.ProviderHealth.SharedHistoryTTLSeconds = 900 + } + if cfg.ProviderHealth.SharedHistoryMaxEvents == 0 { + cfg.ProviderHealth.SharedHistoryMaxEvents = 20000 + } + if cfg.ProviderHealth.ProbeIntervalSeconds == 0 { + cfg.ProviderHealth.ProbeIntervalSeconds = 30 + } + if cfg.ProviderHealth.ProbeTimeoutSeconds == 0 { + cfg.ProviderHealth.ProbeTimeoutSeconds = 5 + } if cfg.Observability.UsageBuffer == 0 { cfg.Observability.UsageBuffer = 8192 } @@ -451,6 +486,14 @@ func resolveSecrets(cfg *Config) error { return err } cfg.Server.DeploymentRegion = strings.ToLower(strings.TrimSpace(os.Getenv(cfg.Server.DeploymentRegionEnv))) + cfg.ProviderHealth.ActiveProbesEnabled, err = envBool(cfg.ProviderHealth.ActiveProbesEnabledEnv) + if err != nil { + return err + } + cfg.ProviderHealth.SharedHistoryEnabled, err = envBool(cfg.ProviderHealth.SharedHistoryEnabledEnv) + if err != nil { + return err + } if cfg.ControlPlane.Enabled { cfg.ControlPlane.DatabaseURL = os.Getenv(cfg.ControlPlane.DatabaseURLEnv) cfg.ControlPlane.RedisURL = os.Getenv(cfg.ControlPlane.RedisURLEnv) @@ -571,6 +614,15 @@ func Validate(cfg Config) error { if cfg.Observability.UsageBuffer < 1 { return errors.New("observability.usage_buffer must be positive") } + if cfg.ProviderHealth.ProbeIntervalSeconds < 5 || cfg.ProviderHealth.ProbeIntervalSeconds > 3600 || + cfg.ProviderHealth.ProbeTimeoutSeconds < 1 || cfg.ProviderHealth.ProbeTimeoutSeconds >= cfg.ProviderHealth.ProbeIntervalSeconds { + return errors.New("provider_health probe interval must be 5-3600 seconds and timeout must be shorter than the interval") + } + if cfg.ProviderHealth.SharedHistoryTTLSeconds < 60 || cfg.ProviderHealth.SharedHistoryTTLSeconds > 86400 || + cfg.ProviderHealth.SharedHistoryMaxEvents < 100 || cfg.ProviderHealth.SharedHistoryMaxEvents > 1_000_000 || + strings.TrimSpace(cfg.ProviderHealth.SharedHistoryStream) == "" { + return errors.New("provider_health shared history requires a stream, TTL of 60-86400 seconds, and 100-1000000 events") + } if cfg.Server.SplitListeners { seen := map[string]string{} for name, address := range map[string]string{"public": cfg.Server.PublicAddress, "admin": cfg.Server.AdminAddress, "webhook": cfg.Server.WebhookAddress, "operations": cfg.Server.OperationsAddress} { @@ -716,8 +768,8 @@ func Validate(cfg Config) error { if provider.Protocol != domain.ProtocolOpenAI && provider.Protocol != domain.ProtocolAnthropic { return fmt.Errorf("provider %q: unsupported protocol %q", provider.ID, provider.Protocol) } - if provider.Protocol == domain.ProtocolOpenAI && provider.WireAPI != "chat_completions" && provider.WireAPI != "responses" { - return fmt.Errorf("provider %q: wire_api must be chat_completions or responses", provider.ID) + if provider.Protocol == domain.ProtocolOpenAI && provider.WireAPI != "chat_completions" && provider.WireAPI != "responses" && provider.WireAPI != "embeddings" { + return fmt.Errorf("provider %q: wire_api must be chat_completions, responses, or embeddings", provider.ID) } if provider.Protocol == domain.ProtocolAnthropic && provider.WireAPI != "messages" { return fmt.Errorf("provider %q: wire_api must be messages", provider.ID) diff --git a/internal/config/config_test.go b/internal/config/config_test.go index bb3e814..ff824ca 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -38,6 +38,73 @@ func TestLoadAppliesDefaultsAndResolvesSecrets(t *testing.T) { } } +func TestLoadResolvesActiveProviderProbeConfiguration(t *testing.T) { + t.Setenv("TEST_UPSTREAM_KEY", "secret") + t.Setenv("TEST_UPSTREAM_URL", "https://example.com/v1") + t.Setenv("TEST_ACTIVE_PROBES", "true") + path := writeConfig(t, `{ + "provider_health":{"active_probes_enabled_env":"TEST_ACTIVE_PROBES","probe_interval_seconds":15,"probe_timeout_seconds":2}, + "providers":[{"id":"primary","protocol":"openai","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"}], + "models":[{"id":"example/model","routes":[{"provider":"primary","upstream_model":"model"}]}] +}`) + cfg, err := Load(path) + if err != nil { + t.Fatal(err) + } + if !cfg.ProviderHealth.ActiveProbesEnabled || cfg.ProviderHealth.ProbeIntervalSeconds != 15 || cfg.ProviderHealth.ProbeTimeoutSeconds != 2 { + t.Fatalf("unexpected provider health config: %+v", cfg.ProviderHealth) + } +} + +func TestLoadResolvesSharedProviderHealthConfiguration(t *testing.T) { + t.Setenv("AIGW_DATABASE_URL", "postgres://aigw:aigw@postgres/aigw") + t.Setenv("AIGW_REDIS_URL", "redis://redis:6379/0") + t.Setenv("AIGW_CREDENTIAL_KEY", "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=") + t.Setenv("TEST_SHARED_HEALTH", "true") + path := writeConfig(t, `{ + "control_plane":{"enabled":true}, + "provider_health":{"shared_history_enabled_env":"TEST_SHARED_HEALTH","shared_history_stream":"test:health","shared_history_ttl_seconds":120,"shared_history_max_events":500} +}`) + cfg, err := Load(path) + if err != nil { + t.Fatal(err) + } + if !cfg.ProviderHealth.SharedHistoryEnabled || cfg.ProviderHealth.SharedHistoryStream != "test:health" || + cfg.ProviderHealth.SharedHistoryTTLSeconds != 120 || cfg.ProviderHealth.SharedHistoryMaxEvents != 500 { + t.Fatalf("unexpected shared provider health config: %+v", cfg.ProviderHealth) + } +} + +func TestLoadAllowsSharedProviderHealthWithoutRedis(t *testing.T) { + t.Setenv("AIGW_DATABASE_URL", "postgres://aigw:aigw@postgres/aigw") + t.Setenv("AIGW_CREDENTIAL_KEY", "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA=") + t.Setenv("TEST_SHARED_HEALTH", "true") + path := writeConfig(t, `{ + "control_plane":{"enabled":true}, + "provider_health":{"shared_history_enabled_env":"TEST_SHARED_HEALTH"} +}`) + cfg, err := Load(path) + if err != nil { + t.Fatal(err) + } + if !cfg.ProviderHealth.SharedHistoryEnabled || cfg.ControlPlane.RedisURL != "" { + t.Fatalf("unexpected degraded shared provider health config: %+v", cfg.ProviderHealth) + } +} + +func TestLoadRejectsInvalidActiveProviderProbeConfiguration(t *testing.T) { + t.Setenv("TEST_UPSTREAM_KEY", "secret") + t.Setenv("TEST_UPSTREAM_URL", "https://example.com/v1") + path := writeConfig(t, `{ + "provider_health":{"probe_interval_seconds":5,"probe_timeout_seconds":5}, + "providers":[{"id":"primary","protocol":"openai","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"}], + "models":[{"id":"example/model","routes":[{"provider":"primary","upstream_model":"model"}]}] +}`) + if _, err := Load(path); err == nil { + t.Fatal("expected invalid provider probe timeout to be rejected") + } +} + func TestLoadValidatesPublicProviderSlugs(t *testing.T) { t.Setenv("TEST_UPSTREAM_KEY", "secret") t.Setenv("TEST_UPSTREAM_URL", "https://example.com") @@ -68,6 +135,22 @@ func TestLoadAcceptsOpenAIResponsesWireAPI(t *testing.T) { } } +func TestLoadAcceptsOpenAIEmbeddingsWireAPI(t *testing.T) { + t.Setenv("TEST_UPSTREAM_KEY", "secret") + t.Setenv("TEST_UPSTREAM_URL", "https://example.com/v1") + path := writeConfig(t, `{ + "providers": [{"id":"embeddings","protocol":"openai","wire_api":"embeddings","base_url_env":"TEST_UPSTREAM_URL","api_key_env":"TEST_UPSTREAM_KEY"}], + "models": [{"id":"example/embedding","capabilities":["embeddings"],"routes":[{"provider":"embeddings","upstream_model":"embedding-model"}]}] +}`) + cfg, err := Load(path) + if err != nil { + t.Fatal(err) + } + if cfg.Providers[0].WireAPI != "embeddings" || len(cfg.Models[0].Capabilities) != 1 || cfg.Models[0].Capabilities[0] != "embeddings" { + t.Fatalf("unexpected embeddings config: provider=%+v model=%+v", cfg.Providers[0], cfg.Models[0]) + } +} + func TestLoadRejectsIncompatibleWireAPI(t *testing.T) { t.Setenv("TEST_UPSTREAM_KEY", "secret") t.Setenv("TEST_UPSTREAM_URL", "https://example.com") diff --git a/internal/controlplane/api_key_test.go b/internal/controlplane/api_key_test.go new file mode 100644 index 0000000..51a3a7d --- /dev/null +++ b/internal/controlplane/api_key_test.go @@ -0,0 +1,26 @@ +package controlplane + +import ( + "crypto/sha256" + "strings" + "testing" +) + +func TestGenerateAPIKeySecretReturnsOnlyDisplayFragments(t *testing.T) { + raw, prefix, suffix, hash, err := generateAPIKeySecret() + if err != nil { + t.Fatal(err) + } + if !strings.HasPrefix(raw, "sk-aigw-") || !strings.HasPrefix(raw, strings.TrimSuffix(prefix, "...")) { + t.Fatalf("prefix %q does not identify the generated key", prefix) + } + if len(suffix) != 6 || !strings.HasSuffix(raw, suffix) { + t.Fatalf("suffix %q does not identify the generated key", suffix) + } + if len(prefix)+len(suffix) >= len(raw) { + t.Fatal("display fragments reveal the complete key") + } + if hash != sha256.Sum256([]byte(raw)) { + t.Fatal("generated digest does not authenticate the raw key") + } +} diff --git a/internal/controlplane/mail_operations.go b/internal/controlplane/mail_operations.go index d42a328..fc7b81b 100644 --- a/internal/controlplane/mail_operations.go +++ b/internal/controlplane/mail_operations.go @@ -149,21 +149,25 @@ func (s *Store) queueBillingNotifications(ctx context.Context, config MailNotifi FROM billing_ledger GROUP BY tenant_id) SELECT w.tenant_id::text,w.currency,w.balance_micros-w.reserved_micros,u.email,u.display_name, COALESCE(spend.today,0),COALESCE(spend.baseline,0), - COALESCE(pref.low_balance_enabled,TRUE),COALESCE(pref.low_balance_threshold_micros,$1) + COALESCE(pref.low_balance_enabled,TRUE),COALESCE(pref.low_balance_threshold_micros,$1), + COALESCE(pref.spend_anomaly_enabled,TRUE),COALESCE(pref.spend_anomaly_multiplier,$2), + COALESCE(pref.spend_anomaly_min_micros,$3) FROM tenant_wallets w JOIN console_users u ON u.tenant_id=w.tenant_id LEFT JOIN spend ON spend.tenant_id=w.tenant_id LEFT JOIN tenant_preferences pref ON pref.tenant_id=w.tenant_id WHERE u.status='active' AND u.email_verified_at IS NOT NULL AND u.role IN ('tenant_admin','tenant_billing') - AND NOT EXISTS (SELECT 1 FROM mail_suppressions s WHERE lower(s.recipient)=lower(u.email))`, config.LowBalanceMicros) + AND NOT EXISTS (SELECT 1 FROM mail_suppressions s WHERE lower(s.recipient)=lower(u.email))`, + config.LowBalanceMicros, config.SpendAnomalyMultiplier, config.SpendAnomalyMinMicros) if err != nil { return err } defer rows.Close() for rows.Next() { var tenantID, currency, email, name string - var available, today, baseline, lowBalanceThreshold int64 - var lowBalanceEnabled bool - if err := rows.Scan(&tenantID, ¤cy, &available, &email, &name, &today, &baseline, &lowBalanceEnabled, &lowBalanceThreshold); err != nil { + var available, today, baseline, lowBalanceThreshold, anomalyMultiplier, anomalyMinimum int64 + var lowBalanceEnabled, anomalyEnabled bool + if err := rows.Scan(&tenantID, ¤cy, &available, &email, &name, &today, &baseline, + &lowBalanceEnabled, &lowBalanceThreshold, &anomalyEnabled, &anomalyMultiplier, &anomalyMinimum); err != nil { return err } day := time.Now().UTC().Format("2006-01-02") @@ -173,7 +177,7 @@ func (s *Store) queueBillingNotifications(ctx context.Context, config MailNotifi return err } } - if baseline > 0 && today >= config.SpendAnomalyMinMicros && today >= baseline*config.SpendAnomalyMultiplier { + if anomalyEnabled && baseline > 0 && today >= anomalyMinimum && today >= baseline*anomalyMultiplier { body := fmt.Sprintf("Hi %s,\n\nAIGW detected unusual API spend today: %.6f %s versus a seven-day daily baseline of %.6f %s. Review API keys and usage in the console.\n", displayName(name), float64(today)/1_000_000, strings.ToUpper(currency), float64(baseline)/1_000_000, strings.ToUpper(currency)) if err := s.queueNotification(ctx, tenantID, email, "spend_anomaly", day, "Unusual AIGW API spend detected", body); err != nil { return err diff --git a/internal/controlplane/mail_operations_integration_test.go b/internal/controlplane/mail_operations_integration_test.go new file mode 100644 index 0000000..912ab7f --- /dev/null +++ b/internal/controlplane/mail_operations_integration_test.go @@ -0,0 +1,144 @@ +package controlplane + +import ( + "context" + "encoding/base64" + "fmt" + "net/url" + "os" + "strings" + "testing" + "time" + + "github.com/jackc/pgx/v5/pgxpool" +) + +func TestBillingNotificationsUseTenantPreferencesLedgerAndEncryptedOutboxPostgres(t *testing.T) { + databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL") + if databaseURL == "" { + t.Skip("AIGW_TEST_DATABASE_URL is not set") + } + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + rootDB, err := pgxpool.New(ctx, databaseURL) + if err != nil { + t.Fatal(err) + } + defer rootDB.Close() + schema := fmt.Sprintf("mail_notifications_%d", time.Now().UnixNano()) + if _, err := rootDB.Exec(ctx, "CREATE SCHEMA "+schema); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _, _ = rootDB.Exec(context.Background(), "DROP SCHEMA "+schema+" CASCADE") }) + parsed, err := url.Parse(databaseURL) + if err != nil { + t.Fatal(err) + } + query := parsed.Query() + query.Set("search_path", schema) + parsed.RawQuery = query.Encode() + isolatedURL := parsed.String() + if err := MigrateDatabase(ctx, isolatedURL); err != nil { + t.Fatal(err) + } + credentialKey := base64.StdEncoding.EncodeToString([]byte("01234567890123456789012345678901")) + store, err := NewStore(ctx, Options{DatabaseURL: isolatedURL, CredentialKey: credentialKey}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = store.Close() }) + + tenant, _, err := store.CreateTenant(ctx, CreateTenantInput{Slug: "mail-alert-test", Name: "Mail Alert Test"}) + if err != nil { + t.Fatal(err) + } + if _, err := store.db.Exec(ctx, `INSERT INTO tenant_wallets (tenant_id,currency,balance_micros) + VALUES ($1,'usd',1000000)`, tenant.ID); err != nil { + t.Fatal(err) + } + enabled := true + threshold := int64(2_000_000) + anomalyMultiplier := int64(4) + anomalyMinimum := int64(400_000) + if _, err := store.SetBillingPreferences(ctx, SetBillingPreferencesInput{TenantID: tenant.ID, + LowBalanceEnabled: &enabled, LowBalanceThresholdMicros: &threshold, SpendAnomalyEnabled: &enabled, + SpendAnomalyMultiplier: &anomalyMultiplier, SpendAnomalyMinMicros: &anomalyMinimum}, + BillingPreferenceDefaults{LowBalanceThresholdMicros: 5_000_000, SpendAnomalyMultiplier: 10, SpendAnomalyMinMicros: 900_000}); err != nil { + t.Fatal(err) + } + if _, err := store.db.Exec(ctx, `INSERT INTO console_users + (tenant_id,email,display_name,role,status,email_verified_at) VALUES + ($1,'billing-alert@example.test','Billing Owner','tenant_billing','active',now()), + ($1,'developer-no-alert@example.test','Developer','tenant_developer','active',now()), + ($1,'unverified-no-alert@example.test','Unverified Billing','tenant_billing','active',NULL)`, tenant.ID); err != nil { + t.Fatal(err) + } + if _, err := store.db.Exec(ctx, `INSERT INTO billing_ledger + (tenant_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description,created_at) + SELECT $1,'usd',-100000,1000000,'usage','request','historical-'||day::text,'Historical usage', + date_trunc('day',now())-make_interval(days=>day) + FROM generate_series(1,7) AS day`, tenant.ID); err != nil { + t.Fatal(err) + } + if _, err := store.db.Exec(ctx, `INSERT INTO billing_ledger + (tenant_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description,created_at) + VALUES ($1,'usd',-500000,1000000,'usage','request','today-spend','Today usage',now())`, tenant.ID); err != nil { + t.Fatal(err) + } + + config := MailNotificationConfig{LowBalanceMicros: 5_000_000, SpendAnomalyMultiplier: 10, SpendAnomalyMinMicros: 900_000} + if err := store.queueBillingNotifications(ctx, config); err != nil { + t.Fatal(err) + } + if err := store.queueBillingNotifications(ctx, config); err != nil { + t.Fatal(err) + } + var billingMessages, otherMessages, notificationEvents int + if err := store.db.QueryRow(ctx, `SELECT + count(*) FILTER (WHERE recipient='billing-alert@example.test'), + count(*) FILTER (WHERE recipient<>'billing-alert@example.test') + FROM console_mail_outbox`).Scan(&billingMessages, &otherMessages); err != nil { + t.Fatal(err) + } + if err := store.db.QueryRow(ctx, `SELECT count(*) FROM mail_notification_events`).Scan(¬ificationEvents); err != nil { + t.Fatal(err) + } + if billingMessages != 2 || otherMessages != 0 || notificationEvents != 2 { + t.Fatalf("notification dedupe or recipient filtering failed: billing=%d other=%d events=%d", billingMessages, otherMessages, notificationEvents) + } + var plaintextLeaks int + if err := store.db.QueryRow(ctx, `SELECT count(*) FROM console_mail_outbox + WHERE convert_from(body_ciphertext,'UTF8') LIKE '%1.000000 USD%'`).Scan(&plaintextLeaks); err == nil { + if plaintextLeaks != 0 { + t.Fatal("notification body was stored as plaintext") + } + } else { + // Authenticated encryption output is arbitrary bytes and usually is not valid UTF-8. + var containsPlaintext bool + if scanErr := store.db.QueryRow(ctx, `SELECT bool_or(position(convert_to('1.000000 USD','UTF8') in body_ciphertext)>0) + FROM console_mail_outbox`).Scan(&containsPlaintext); scanErr != nil { + t.Fatal(scanErr) + } + if containsPlaintext { + t.Fatal("notification body was stored as plaintext") + } + } + + bodies := make([]string, 0, 2) + for range 2 { + message, ok, err := store.ClaimMail(ctx) + if err != nil { + t.Fatal(err) + } + if !ok || message.Recipient != "billing-alert@example.test" { + t.Fatalf("unexpected claimed notification: ok=%v message=%+v", ok, message) + } + bodies = append(bodies, message.Body) + } + joined := strings.Join(bodies, "\n") + for _, expected := range []string{"1.000000 USD", "0.500000 USD", "0.100000 USD"} { + if !strings.Contains(joined, expected) { + t.Fatalf("decrypted notifications do not contain %q: %s", expected, joined) + } + } +} diff --git a/internal/controlplane/manager.go b/internal/controlplane/manager.go index 212963b..58c31fe 100644 --- a/internal/controlplane/manager.go +++ b/internal/controlplane/manager.go @@ -15,11 +15,14 @@ import ( const broadcastQueueSize = 128 +const redisHealthCheckTimeout = 500 * time.Millisecond + type managerStore interface { LoadSnapshot(context.Context) (Snapshot, error) DatabaseGeneration(context.Context) (int64, error) PublishChange(context.Context, ChangeEvent) error Subscribe(context.Context) (<-chan ChangeMessage, func() error, error) + PingRedis(context.Context) error RedisEnabled() bool } @@ -129,7 +132,7 @@ func (m *Manager) RedisConnected() bool { func (m *Manager) Run(ctx context.Context) { var workers sync.WaitGroup if m.store.RedisEnabled() { - workers.Add(2) + workers.Add(3) go func() { defer workers.Done() m.runSubscriptions(ctx) @@ -138,6 +141,10 @@ func (m *Manager) Run(ctx context.Context) { defer workers.Done() m.runBroadcasts(ctx) }() + go func() { + defer workers.Done() + m.runRedisHealth(ctx) + }() } else { m.logger.Info("control_plane_redis_disabled", "fallback", "postgres_polling") } @@ -145,6 +152,35 @@ func (m *Manager) Run(ctx context.Context) { workers.Wait() } +func (m *Manager) runRedisHealth(ctx context.Context) { + interval := m.pollInterval + if interval > time.Second { + interval = time.Second + } + if interval <= 0 { + interval = time.Second + } + ticker := time.NewTicker(interval) + defer ticker.Stop() + for { + pingContext, cancel := context.WithTimeout(ctx, redisHealthCheckTimeout) + err := m.store.PingRedis(pingContext) + cancel() + if err != nil { + if m.redisConnected.Swap(false) { + m.logger.Warn("control_plane_redis_unavailable", "error", err, "fallback", "postgres_polling") + } + } else if !m.redisConnected.Swap(true) { + m.logger.Info("control_plane_redis_recovered") + } + select { + case <-ctx.Done(): + return + case <-ticker.C: + } + } +} + func (m *Manager) runPolling(ctx context.Context) { ticker := time.NewTicker(m.pollInterval) defer ticker.Stop() diff --git a/internal/controlplane/manager_test.go b/internal/controlplane/manager_test.go index b292060..78a0fdb 100644 --- a/internal/controlplane/manager_test.go +++ b/internal/controlplane/manager_test.go @@ -25,6 +25,7 @@ type fakeManagerStore struct { publishErr error published chan ChangeEvent subscribe func(context.Context, int64) (<-chan ChangeMessage, func() error, error) + pingRedis func(context.Context) error } type capturePolicies struct{ values []domain.LimitPolicy } @@ -84,6 +85,13 @@ func (s *fakeManagerStore) RedisEnabled() bool { return s.redisEnabled } +func (s *fakeManagerStore) PingRedis(ctx context.Context) error { + if s.pingRedis != nil { + return s.pingRedis(ctx) + } + return nil +} + func newTestManager(store managerStore, logger *slog.Logger, interval time.Duration) *Manager { return NewManager(store, catalog.NewModels(nil), auth.NewDynamic(nil, false), logger, interval) } @@ -196,6 +204,14 @@ func TestSubscriptionMessageRestoresConnectedStateAfterPublishFailure(t *testing store := newFakeManagerStore(1) store.redisEnabled = true store.publishErr = errors.New("redis unavailable") + var redisAvailable atomic.Bool + redisAvailable.Store(true) + store.pingRedis = func(context.Context) error { + if !redisAvailable.Load() { + return errors.New("redis unavailable") + } + return nil + } messages := make(chan ChangeMessage, 1) store.subscribe = func(_ context.Context, _ int64) (<-chan ChangeMessage, func() error, error) { return messages, func() error { return nil }, nil @@ -212,12 +228,14 @@ func TestSubscriptionMessageRestoresConnectedStateAfterPublishFailure(t *testing close(done) }() waitUntil(t, time.Second, manager.RedisConnected) + redisAvailable.Store(false) if err := manager.AfterMutation(context.Background(), 1, "model", "model-1"); err != nil { t.Fatal(err) } waitUntil(t, time.Second, func() bool { return !manager.RedisConnected() }) store.snapshot.Store(Snapshot{Generation: 2}) + redisAvailable.Store(true) messages <- ChangeMessage{Payload: `{"generation":2,"resource":"model"}`} waitUntil(t, time.Second, func() bool { return manager.RedisConnected() && manager.Generation() == 2 }) @@ -229,6 +247,37 @@ func TestSubscriptionMessageRestoresConnectedStateAfterPublishFailure(t *testing } } +func TestRedisHealthCheckReportsFailureAndRecovery(t *testing.T) { + store := newFakeManagerStore(1) + store.redisEnabled = true + var redisAvailable atomic.Bool + redisAvailable.Store(true) + store.pingRedis = func(context.Context) error { + if !redisAvailable.Load() { + return errors.New("redis unavailable") + } + return nil + } + manager := newTestManager(store, slog.New(slog.NewTextHandler(&safeLogBuffer{}, nil)), 10*time.Millisecond) + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { + manager.Run(ctx) + close(done) + }() + waitUntil(t, time.Second, manager.RedisConnected) + redisAvailable.Store(false) + waitUntil(t, time.Second, func() bool { return !manager.RedisConnected() }) + redisAvailable.Store(true) + waitUntil(t, time.Second, manager.RedisConnected) + cancel() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("manager did not stop") + } +} + func TestRedisCanBeDisabled(t *testing.T) { store := newFakeManagerStore(3) manager := newTestManager(store, slog.New(slog.NewTextHandler(&bytes.Buffer{}, nil)), 10*time.Millisecond) diff --git a/internal/controlplane/mutations.go b/internal/controlplane/mutations.go index 9f81d6e..aba4d36 100644 --- a/internal/controlplane/mutations.go +++ b/internal/controlplane/mutations.go @@ -85,8 +85,8 @@ func (s *Store) CreateAPIKey(ctx context.Context, input CreateAPIKeyInput) (Crea if input.TenantID == "" || input.ProjectID == "" || input.Name == "" { return CreatedAPIKey{}, 0, errors.New("API key requires tenant_id, project_id, and name") } - if len(input.Name) > 120 || input.MonthlySpendMicros < 0 { - return CreatedAPIKey{}, 0, errors.New("API key name or monthly spend limit is invalid") + if len(input.Name) > 120 || input.MonthlySpendMicros < 0 || input.DailySpendMicros < 0 || input.RequestsPerMinute < 0 || input.TokensPerMinute < 0 { + return CreatedAPIKey{}, 0, errors.New("API key name, spend limits, or rate limits are invalid") } if input.ExpiresAt != nil && !input.ExpiresAt.After(time.Now()) { return CreatedAPIKey{}, 0, errors.New("API key expiry must be in the future") @@ -107,13 +107,10 @@ func (s *Store) CreateAPIKey(ctx context.Context, input CreateAPIKeyInput) (Crea } scopesJSON, _ := json.Marshal(scopes) tagsJSON, _ := json.Marshal(tags) - random := make([]byte, 32) - if _, err := rand.Read(random); err != nil { - return CreatedAPIKey{}, 0, fmt.Errorf("generate API key: %w", err) + rawKey, prefix, suffix, hash, err := generateAPIKeySecret() + if err != nil { + return CreatedAPIKey{}, 0, err } - rawKey := "sk-aigw-" + base64.RawURLEncoding.EncodeToString(random) - hash := sha256.Sum256([]byte(rawKey)) - prefix := rawKey[:min(18, len(rawKey))] + "..." tx, err := s.db.Begin(ctx) if err != nil { @@ -122,14 +119,17 @@ func (s *Store) CreateAPIKey(ctx context.Context, input CreateAPIKeyInput) (Crea defer tx.Rollback(ctx) var result CreatedAPIKey err = tx.QueryRow(ctx, ` - INSERT INTO api_keys (tenant_id, project_id, name, key_prefix, key_hash, scopes, tags, monthly_spend_micros, expires_at) - VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9) - RETURNING id::text, tenant_id::text, project_id::text, name, key_prefix, scopes, tags, - monthly_spend_micros, status, expires_at, last_used_at, created_at`, - input.TenantID, input.ProjectID, input.Name, prefix, hash[:], scopesJSON, tagsJSON, - input.MonthlySpendMicros, input.ExpiresAt, - ).Scan(&result.ID, &result.TenantID, &result.ProjectID, &result.Name, &result.KeyPrefix, - &scopesJSON, &tagsJSON, &result.MonthlySpendMicros, &result.Status, &result.ExpiresAt, + INSERT INTO api_keys (tenant_id, project_id, name, key_prefix, key_suffix, key_hash, scopes, tags, monthly_spend_micros, + daily_spend_micros, requests_per_minute, tokens_per_minute, expires_at) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13) + RETURNING id::text, tenant_id::text, project_id::text, name, key_prefix, key_suffix, scopes, tags, + monthly_spend_micros, daily_spend_micros, requests_per_minute, tokens_per_minute, + status, expires_at, last_used_at, created_at`, + input.TenantID, input.ProjectID, input.Name, prefix, suffix, hash[:], scopesJSON, tagsJSON, + input.MonthlySpendMicros, input.DailySpendMicros, input.RequestsPerMinute, input.TokensPerMinute, input.ExpiresAt, + ).Scan(&result.ID, &result.TenantID, &result.ProjectID, &result.Name, &result.KeyPrefix, &result.KeySuffix, + &scopesJSON, &tagsJSON, &result.MonthlySpendMicros, &result.DailySpendMicros, + &result.RequestsPerMinute, &result.TokensPerMinute, &result.Status, &result.ExpiresAt, &result.LastUsedAt, &result.CreatedAt) if err != nil { return CreatedAPIKey{}, 0, fmt.Errorf("create API key: %w", err) @@ -163,6 +163,90 @@ func (s *Store) RevokeAPIKey(ctx context.Context, id string) (int64, error) { return s.toggle(ctx, `UPDATE api_keys SET status = 'revoked', revoked_at = now() WHERE id = $1 AND status <> 'revoked'`, id) } +func (s *Store) DisableAPIKey(ctx context.Context, id string) (int64, error) { + return s.toggle(ctx, `UPDATE api_keys SET status='disabled', disabled_at=now() WHERE id=$1 AND status='active'`, id) +} + +func (s *Store) EnableAPIKey(ctx context.Context, id string) (int64, error) { + return s.toggle(ctx, `UPDATE api_keys SET status='active', disabled_at=NULL WHERE id=$1 AND status='disabled' AND (expires_at IS NULL OR expires_at > now())`, id) +} + +func (s *Store) RotateAPIKey(ctx context.Context, id string) (CreatedAPIKey, int64, error) { + tx, err := s.db.Begin(ctx) + if err != nil { + return CreatedAPIKey{}, 0, err + } + defer tx.Rollback(ctx) + + var tenantID, projectID, name, status string + var scopesJSON, tagsJSON []byte + var monthlySpendMicros, dailySpendMicros, requestsPerMinute, tokensPerMinute int64 + var expiresAt *time.Time + err = tx.QueryRow(ctx, `SELECT tenant_id::text,project_id::text,name,scopes,tags,monthly_spend_micros, + daily_spend_micros,requests_per_minute,tokens_per_minute,expires_at,status + FROM api_keys WHERE id=$1 FOR UPDATE`, id).Scan(&tenantID, &projectID, &name, &scopesJSON, &tagsJSON, + &monthlySpendMicros, &dailySpendMicros, &requestsPerMinute, &tokensPerMinute, &expiresAt, &status) + if errors.Is(err, pgx.ErrNoRows) { + return CreatedAPIKey{}, 0, ErrNotFound + } + if err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("lock API key for rotation: %w", err) + } + if status == "revoked" { + return CreatedAPIKey{}, 0, errors.New("revoked API key cannot be rotated") + } + if expiresAt != nil && !expiresAt.After(time.Now()) { + return CreatedAPIKey{}, 0, errors.New("expired API key cannot be rotated") + } + rawKey, prefix, suffix, hash, err := generateAPIKeySecret() + if err != nil { + return CreatedAPIKey{}, 0, err + } + var result CreatedAPIKey + err = tx.QueryRow(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_suffix,key_hash,scopes,tags, + monthly_spend_micros,daily_spend_micros,requests_per_minute,tokens_per_minute,expires_at,rotated_from_id) + VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14) + RETURNING id::text,tenant_id::text,project_id::text,name,key_prefix,key_suffix,scopes,tags,monthly_spend_micros, + daily_spend_micros,requests_per_minute,tokens_per_minute,status,expires_at,last_used_at,created_at`, + tenantID, projectID, name, prefix, suffix, hash[:], scopesJSON, tagsJSON, monthlySpendMicros, dailySpendMicros, + requestsPerMinute, tokensPerMinute, expiresAt, id).Scan(&result.ID, &result.TenantID, &result.ProjectID, + &result.Name, &result.KeyPrefix, &result.KeySuffix, &scopesJSON, &tagsJSON, &result.MonthlySpendMicros, &result.DailySpendMicros, + &result.RequestsPerMinute, &result.TokensPerMinute, &result.Status, &result.ExpiresAt, &result.LastUsedAt, &result.CreatedAt) + if err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("create rotated API key: %w", err) + } + if _, err := tx.Exec(ctx, `INSERT INTO api_key_model_restrictions (api_key_id,model_id) + SELECT $1,model_id FROM api_key_model_restrictions WHERE api_key_id=$2`, result.ID, id); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("copy rotated API key model restrictions: %w", err) + } + var allowedModelsJSON []byte + if err := tx.QueryRow(ctx, `SELECT COALESCE(jsonb_agg(m.public_id ORDER BY m.public_id),'[]'::jsonb) + FROM api_key_model_restrictions r JOIN models m ON m.id=r.model_id WHERE r.api_key_id=$1`, result.ID).Scan(&allowedModelsJSON); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("read rotated API key model restrictions: %w", err) + } + if _, err := tx.Exec(ctx, `UPDATE api_keys SET status='revoked',revoked_at=now() WHERE id=$1`, id); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("revoke rotated API key: %w", err) + } + if err := json.Unmarshal(scopesJSON, &result.Scopes); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("decode rotated API key scopes: %w", err) + } + if err := json.Unmarshal(tagsJSON, &result.Tags); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("decode rotated API key tags: %w", err) + } + if err := json.Unmarshal(allowedModelsJSON, &result.AllowedModels); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("decode rotated API key model restrictions: %w", err) + } + result.Key = rawKey + generation, err := bumpGeneration(ctx, tx) + if err != nil { + return CreatedAPIKey{}, 0, err + } + if err := tx.Commit(ctx); err != nil { + return CreatedAPIKey{}, 0, err + } + return result, generation, nil +} + func (s *Store) CreateProvider(ctx context.Context, input CreateProviderInput) (Provider, int64, error) { input.Slug = strings.ToLower(strings.TrimSpace(input.Slug)) input.Name = strings.TrimSpace(input.Name) @@ -184,7 +268,7 @@ func (s *Store) CreateProvider(ctx context.Context, input CreateProviderInput) ( if !slugPattern.MatchString(input.Slug) || input.Name == "" || input.APIKey == "" || (input.Protocol != "openai" && input.Protocol != "anthropic") { return Provider{}, 0, errors.New("provider requires a unique 3-64 character lowercase slug, name, protocol openai|anthropic, base_url, and api_key") } - if (input.Protocol == "openai" && input.WireAPI != "chat_completions" && input.WireAPI != "responses") || + if (input.Protocol == "openai" && input.WireAPI != "chat_completions" && input.WireAPI != "responses" && input.WireAPI != "embeddings") || (input.Protocol == "anthropic" && input.WireAPI != "messages") { return Provider{}, 0, errors.New("provider wire_api is incompatible with protocol") } @@ -474,3 +558,15 @@ func uniqueStrings(values []string) []string { } return result } + +func generateAPIKeySecret() (string, string, string, [sha256.Size]byte, error) { + random := make([]byte, 32) + if _, err := rand.Read(random); err != nil { + return "", "", "", [sha256.Size]byte{}, fmt.Errorf("generate API key: %w", err) + } + rawKey := "sk-aigw-" + base64.RawURLEncoding.EncodeToString(random) + hash := sha256.Sum256([]byte(rawKey)) + prefix := rawKey[:min(18, len(rawKey))] + "..." + suffix := rawKey[max(0, len(rawKey)-6):] + return rawKey, prefix, suffix, hash, nil +} diff --git a/internal/controlplane/preferences.go b/internal/controlplane/preferences.go index 73bcce6..2f80ad5 100644 --- a/internal/controlplane/preferences.go +++ b/internal/controlplane/preferences.go @@ -10,22 +10,37 @@ import ( "github.com/jackc/pgx/v5" ) -const maxLowBalanceThresholdMicros int64 = 1_000_000_000_000_000 +const maxAlertThresholdMicros int64 = 1_000_000_000_000_000 -func (s *Store) GetTenantPreferences(ctx context.Context, tenantID string, defaultThresholdMicros int64) (TenantPreferences, error) { - if defaultThresholdMicros <= 0 { - defaultThresholdMicros = 5_000_000 +func normalizeBillingPreferenceDefaults(defaults BillingPreferenceDefaults) BillingPreferenceDefaults { + if defaults.LowBalanceThresholdMicros <= 0 { + defaults.LowBalanceThresholdMicros = 5_000_000 } - result := TenantPreferences{TenantID: strings.TrimSpace(tenantID), LowBalanceEnabled: true, LowBalanceThresholdMicros: defaultThresholdMicros} + if defaults.SpendAnomalyMultiplier < 2 { + defaults.SpendAnomalyMultiplier = 3 + } + if defaults.SpendAnomalyMinMicros < 0 { + defaults.SpendAnomalyMinMicros = 10_000_000 + } + return defaults +} + +func (s *Store) GetTenantPreferences(ctx context.Context, tenantID string, defaults BillingPreferenceDefaults) (TenantPreferences, error) { + defaults = normalizeBillingPreferenceDefaults(defaults) + result := TenantPreferences{TenantID: strings.TrimSpace(tenantID), LowBalanceEnabled: true, + LowBalanceThresholdMicros: defaults.LowBalanceThresholdMicros, SpendAnomalyEnabled: true, + SpendAnomalyMultiplier: defaults.SpendAnomalyMultiplier, SpendAnomalyMinMicros: defaults.SpendAnomalyMinMicros} if result.TenantID == "" { return result, nil } var defaultModel, fallbackModel *string var updatedAt time.Time err := s.db.QueryRow(ctx, ` - SELECT default_model, fallback_model, low_balance_enabled, low_balance_threshold_micros, updated_at - FROM tenant_preferences WHERE tenant_id=$1`, result.TenantID).Scan( - &defaultModel, &fallbackModel, &result.LowBalanceEnabled, &result.LowBalanceThresholdMicros, &updatedAt) + SELECT default_model, fallback_model, low_balance_enabled, low_balance_threshold_micros, + spend_anomaly_enabled, COALESCE(spend_anomaly_multiplier,$2), COALESCE(spend_anomaly_min_micros,$3), updated_at + FROM tenant_preferences WHERE tenant_id=$1`, result.TenantID, defaults.SpendAnomalyMultiplier, defaults.SpendAnomalyMinMicros).Scan( + &defaultModel, &fallbackModel, &result.LowBalanceEnabled, &result.LowBalanceThresholdMicros, + &result.SpendAnomalyEnabled, &result.SpendAnomalyMultiplier, &result.SpendAnomalyMinMicros, &updatedAt) if errors.Is(err, pgx.ErrNoRows) { return result, nil } @@ -42,10 +57,11 @@ func (s *Store) GetTenantPreferences(ctx context.Context, tenantID string, defau return result, nil } -func (s *Store) SetDeveloperPreferences(ctx context.Context, input SetDeveloperPreferencesInput) (TenantPreferences, error) { +func (s *Store) SetDeveloperPreferences(ctx context.Context, input SetDeveloperPreferencesInput, defaults BillingPreferenceDefaults) (TenantPreferences, error) { input.TenantID = strings.TrimSpace(input.TenantID) input.DefaultModel = strings.TrimSpace(input.DefaultModel) input.FallbackModel = strings.TrimSpace(input.FallbackModel) + defaults = normalizeBillingPreferenceDefaults(defaults) if input.TenantID == "" { return TenantPreferences{}, errors.New("tenant_id is required") } @@ -87,9 +103,12 @@ func (s *Store) SetDeveloperPreferences(ctx context.Context, input SetDeveloperP ON CONFLICT (tenant_id) DO UPDATE SET default_model=EXCLUDED.default_model, fallback_model=EXCLUDED.fallback_model, updated_at=now() RETURNING tenant_id::text, default_model, fallback_model, low_balance_enabled, - low_balance_threshold_micros, updated_at`, input.TenantID, input.DefaultModel, input.FallbackModel).Scan( + low_balance_threshold_micros, spend_anomaly_enabled, COALESCE(spend_anomaly_multiplier,$4), + COALESCE(spend_anomaly_min_micros,$5), updated_at`, input.TenantID, input.DefaultModel, input.FallbackModel, + defaults.SpendAnomalyMultiplier, defaults.SpendAnomalyMinMicros).Scan( &result.TenantID, &defaultModel, &fallbackModel, &result.LowBalanceEnabled, - &result.LowBalanceThresholdMicros, &updatedAt); err != nil { + &result.LowBalanceThresholdMicros, &result.SpendAnomalyEnabled, &result.SpendAnomalyMultiplier, + &result.SpendAnomalyMinMicros, &updatedAt); err != nil { return TenantPreferences{}, fmt.Errorf("save developer preferences: %w", err) } if defaultModel != nil { @@ -105,17 +124,21 @@ func (s *Store) SetDeveloperPreferences(ctx context.Context, input SetDeveloperP return result, nil } -func (s *Store) SetBillingPreferences(ctx context.Context, input SetBillingPreferencesInput, defaultThresholdMicros int64) (TenantPreferences, error) { +func (s *Store) SetBillingPreferences(ctx context.Context, input SetBillingPreferencesInput, defaults BillingPreferenceDefaults) (TenantPreferences, error) { input.TenantID = strings.TrimSpace(input.TenantID) if input.TenantID == "" { return TenantPreferences{}, errors.New("tenant_id is required") } - if defaultThresholdMicros <= 0 { - defaultThresholdMicros = 5_000_000 - } - if input.LowBalanceThresholdMicros != nil && (*input.LowBalanceThresholdMicros < 0 || *input.LowBalanceThresholdMicros > maxLowBalanceThresholdMicros) { + defaults = normalizeBillingPreferenceDefaults(defaults) + if input.LowBalanceThresholdMicros != nil && (*input.LowBalanceThresholdMicros < 0 || *input.LowBalanceThresholdMicros > maxAlertThresholdMicros) { return TenantPreferences{}, errors.New("low balance threshold is outside the supported range") } + if input.SpendAnomalyMultiplier != nil && (*input.SpendAnomalyMultiplier < 2 || *input.SpendAnomalyMultiplier > 1000) { + return TenantPreferences{}, errors.New("spend anomaly multiplier must be between 2 and 1000") + } + if input.SpendAnomalyMinMicros != nil && (*input.SpendAnomalyMinMicros < 0 || *input.SpendAnomalyMinMicros > maxAlertThresholdMicros) { + return TenantPreferences{}, errors.New("spend anomaly minimum is outside the supported range") + } tx, err := s.db.Begin(ctx) if err != nil { return TenantPreferences{}, err @@ -131,15 +154,25 @@ func (s *Store) SetBillingPreferences(ctx context.Context, input SetBillingPrefe var defaultModel, fallbackModel *string var updatedAt time.Time if err := tx.QueryRow(ctx, ` - INSERT INTO tenant_preferences (tenant_id, low_balance_enabled, low_balance_threshold_micros) - VALUES ($1, COALESCE($2::boolean, TRUE), COALESCE($3::bigint, $4::bigint)) + INSERT INTO tenant_preferences (tenant_id, low_balance_enabled, low_balance_threshold_micros, + spend_anomaly_enabled, spend_anomaly_multiplier, spend_anomaly_min_micros) + VALUES ($1, COALESCE($2::boolean, TRUE), COALESCE($3::bigint, $7::bigint), + COALESCE($4::boolean, TRUE), $5::bigint, $6::bigint) ON CONFLICT (tenant_id) DO UPDATE SET low_balance_enabled=COALESCE($2::boolean, tenant_preferences.low_balance_enabled), - low_balance_threshold_micros=COALESCE($3::bigint, tenant_preferences.low_balance_threshold_micros), updated_at=now() + low_balance_threshold_micros=COALESCE($3::bigint, tenant_preferences.low_balance_threshold_micros), + spend_anomaly_enabled=COALESCE($4::boolean, tenant_preferences.spend_anomaly_enabled), + spend_anomaly_multiplier=COALESCE($5::bigint, tenant_preferences.spend_anomaly_multiplier), + spend_anomaly_min_micros=COALESCE($6::bigint, tenant_preferences.spend_anomaly_min_micros), updated_at=now() RETURNING tenant_id::text, default_model, fallback_model, low_balance_enabled, - low_balance_threshold_micros, updated_at`, input.TenantID, input.LowBalanceEnabled, input.LowBalanceThresholdMicros, defaultThresholdMicros).Scan( + low_balance_threshold_micros, spend_anomaly_enabled, + COALESCE(spend_anomaly_multiplier,$8), COALESCE(spend_anomaly_min_micros,$9), updated_at`, + input.TenantID, input.LowBalanceEnabled, input.LowBalanceThresholdMicros, input.SpendAnomalyEnabled, + input.SpendAnomalyMultiplier, input.SpendAnomalyMinMicros, defaults.LowBalanceThresholdMicros, + defaults.SpendAnomalyMultiplier, defaults.SpendAnomalyMinMicros).Scan( &result.TenantID, &defaultModel, &fallbackModel, &result.LowBalanceEnabled, - &result.LowBalanceThresholdMicros, &updatedAt); err != nil { + &result.LowBalanceThresholdMicros, &result.SpendAnomalyEnabled, &result.SpendAnomalyMultiplier, + &result.SpendAnomalyMinMicros, &updatedAt); err != nil { return TenantPreferences{}, fmt.Errorf("save billing preferences: %w", err) } if defaultModel != nil { diff --git a/internal/controlplane/preferences_test.go b/internal/controlplane/preferences_test.go index 88a3d50..2d390d2 100644 --- a/internal/controlplane/preferences_test.go +++ b/internal/controlplane/preferences_test.go @@ -6,11 +6,13 @@ import ( ) func TestGetTenantPreferencesWithoutTenantUsesConfiguredDefault(t *testing.T) { - result, err := (&Store{}).GetTenantPreferences(context.Background(), "", 12_500_000) + defaults := BillingPreferenceDefaults{LowBalanceThresholdMicros: 12_500_000, SpendAnomalyMultiplier: 7, SpendAnomalyMinMicros: 8_500_000} + result, err := (&Store{}).GetTenantPreferences(context.Background(), "", defaults) if err != nil { t.Fatal(err) } - if !result.LowBalanceEnabled || result.LowBalanceThresholdMicros != 12_500_000 { + if !result.LowBalanceEnabled || result.LowBalanceThresholdMicros != 12_500_000 || !result.SpendAnomalyEnabled || + result.SpendAnomalyMultiplier != 7 || result.SpendAnomalyMinMicros != 8_500_000 { t.Fatalf("unexpected defaults: %+v", result) } } @@ -19,13 +21,19 @@ func TestPreferenceValidationRejectsUnsafeValuesBeforeDatabaseAccess(t *testing. store := &Store{} if _, err := store.SetDeveloperPreferences(context.Background(), SetDeveloperPreferencesInput{ TenantID: "tenant", DefaultModel: "same", FallbackModel: "same", - }); err == nil { + }, BillingPreferenceDefaults{}); err == nil { t.Fatal("expected identical default and fallback models to fail") } - threshold := maxLowBalanceThresholdMicros + 1 + threshold := maxAlertThresholdMicros + 1 if _, err := store.SetBillingPreferences(context.Background(), SetBillingPreferencesInput{ TenantID: "tenant", LowBalanceThresholdMicros: &threshold, - }, 5_000_000); err == nil { + }, BillingPreferenceDefaults{}); err == nil { t.Fatal("expected excessive low balance threshold to fail") } + multiplier := int64(1) + if _, err := store.SetBillingPreferences(context.Background(), SetBillingPreferencesInput{ + TenantID: "tenant", SpendAnomalyMultiplier: &multiplier, + }, BillingPreferenceDefaults{}); err == nil { + t.Fatal("expected invalid spend anomaly multiplier to fail") + } } diff --git a/internal/controlplane/queries.go b/internal/controlplane/queries.go index 48610c7..6a5ff0b 100644 --- a/internal/controlplane/queries.go +++ b/internal/controlplane/queries.go @@ -85,12 +85,17 @@ func (s *Store) ListAPIKeys(ctx context.Context) ([]APIKey, error) { } func (s *Store) ListAPIKeysFor(ctx context.Context, tenantID string) ([]APIKey, error) { - periodStart := time.Date(time.Now().UTC().Year(), time.Now().UTC().Month(), 1, 0, 0, 0, 0, time.UTC) + now := time.Now().UTC() + periodStart := time.Date(now.Year(), now.Month(), 1, 0, 0, 0, 0, time.UTC) periodEnd := periodStart.AddDate(0, 1, 0) + dayStart := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.UTC) + dayEnd := dayStart.AddDate(0, 0, 1) query := ` - SELECT k.id::text, k.tenant_id::text, k.project_id::text, k.name, k.key_prefix, k.scopes, - k.tags, k.monthly_spend_micros, k.status, k.expires_at, k.last_used_at, k.created_at, + SELECT k.id::text, k.tenant_id::text, k.project_id::text, k.name, k.key_prefix, k.key_suffix, k.scopes, + k.tags, k.monthly_spend_micros, k.daily_spend_micros, k.requests_per_minute, k.tokens_per_minute, + k.status, k.expires_at, k.last_used_at, k.created_at, usage.month_spend, usage.month_requests, pending.month_reserved, + daily.day_spend, daily.day_requests, daily_pending.day_reserved, COALESCE(( SELECT jsonb_agg(m.public_id ORDER BY m.public_id) FROM api_key_model_restrictions r @@ -107,10 +112,20 @@ func (s *Store) ListAPIKeysFor(ctx context.Context, tenantID string) ([]APIKey, FROM billing_reservations b WHERE b.key_id = k.id AND b.status IN ('pending', 'metering_failed') AND b.created_at >= $1 AND b.created_at < $2 - ) pending` - args := []any{periodStart, periodEnd} + ) pending + CROSS JOIN LATERAL ( + SELECT COALESCE(SUM(u.cost_micros), 0)::bigint AS day_spend, COUNT(*)::bigint AS day_requests + FROM usage_events u WHERE u.key_id = k.id AND u.started_at >= $3 AND u.started_at < $4 + ) daily + CROSS JOIN LATERAL ( + SELECT COALESCE(SUM(b.reserved_micros), 0)::bigint AS day_reserved + FROM billing_reservations b + WHERE b.key_id = k.id AND b.status IN ('pending', 'metering_failed') + AND b.created_at >= $3 AND b.created_at < $4 + ) daily_pending` + args := []any{periodStart, periodEnd, dayStart, dayEnd} if tenantID != "" { - query += ` WHERE k.tenant_id=$3` + query += ` WHERE k.tenant_id=$5` args = append(args, tenantID) } query += ` ORDER BY k.created_at DESC` @@ -123,10 +138,12 @@ func (s *Store) ListAPIKeysFor(ctx context.Context, tenantID string) ([]APIKey, for rows.Next() { var item APIKey var scopesJSON, tagsJSON, allowedModelsJSON []byte - if err := rows.Scan(&item.ID, &item.TenantID, &item.ProjectID, &item.Name, &item.KeyPrefix, - &scopesJSON, &tagsJSON, &item.MonthlySpendMicros, &item.Status, &item.ExpiresAt, + if err := rows.Scan(&item.ID, &item.TenantID, &item.ProjectID, &item.Name, &item.KeyPrefix, &item.KeySuffix, + &scopesJSON, &tagsJSON, &item.MonthlySpendMicros, &item.DailySpendMicros, &item.RequestsPerMinute, + &item.TokensPerMinute, &item.Status, &item.ExpiresAt, &item.LastUsedAt, &item.CreatedAt, &item.CurrentMonthSpendMicros, &item.CurrentMonthRequests, - &item.CurrentMonthReservedMicros, &allowedModelsJSON); err != nil { + &item.CurrentMonthReservedMicros, &item.CurrentDaySpendMicros, &item.CurrentDayRequests, + &item.CurrentDayReservedMicros, &allowedModelsJSON); err != nil { return nil, fmt.Errorf("scan API key: %w", err) } if err := json.Unmarshal(scopesJSON, &item.Scopes); err != nil { diff --git a/internal/controlplane/schema.sql b/internal/controlplane/schema.sql index ef1ccdd..d71f37f 100644 --- a/internal/controlplane/schema.sql +++ b/internal/controlplane/schema.sql @@ -41,24 +41,35 @@ CREATE TABLE IF NOT EXISTS api_keys ( project_id UUID NOT NULL, name TEXT NOT NULL, key_prefix TEXT NOT NULL, + key_suffix TEXT NOT NULL DEFAULT '', key_hash BYTEA NOT NULL UNIQUE CHECK (octet_length(key_hash) = 32), scopes JSONB NOT NULL DEFAULT '["inference"]'::jsonb, - status TEXT NOT NULL DEFAULT 'active' CHECK (status IN ('active', 'revoked')), + status TEXT NOT NULL DEFAULT 'active' CHECK (status IN ('active', 'disabled', 'revoked')), last_used_at TIMESTAMPTZ, created_at TIMESTAMPTZ NOT NULL DEFAULT now(), revoked_at TIMESTAMPTZ, FOREIGN KEY (project_id, tenant_id) REFERENCES projects(id, tenant_id) ON DELETE CASCADE ); ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS monthly_spend_micros BIGINT NOT NULL DEFAULT 0 CHECK (monthly_spend_micros >= 0); +ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS daily_spend_micros BIGINT NOT NULL DEFAULT 0 CHECK (daily_spend_micros >= 0); +ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS requests_per_minute BIGINT NOT NULL DEFAULT 0 CHECK (requests_per_minute >= 0); +ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS tokens_per_minute BIGINT NOT NULL DEFAULT 0 CHECK (tokens_per_minute >= 0); ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS expires_at TIMESTAMPTZ; ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS tags JSONB NOT NULL DEFAULT '[]'::jsonb; +ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS disabled_at TIMESTAMPTZ; +ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS rotated_from_id UUID REFERENCES api_keys(id) ON DELETE SET NULL; +ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS key_suffix TEXT NOT NULL DEFAULT ''; +ALTER TABLE api_keys DROP CONSTRAINT IF EXISTS api_keys_key_suffix_check; +ALTER TABLE api_keys ADD CONSTRAINT api_keys_key_suffix_check CHECK (key_suffix = '' OR key_suffix ~ '^[A-Za-z0-9_-]{6}$'); +ALTER TABLE api_keys DROP CONSTRAINT IF EXISTS api_keys_status_check; +ALTER TABLE api_keys ADD CONSTRAINT api_keys_status_check CHECK (status IN ('active', 'disabled', 'revoked')); CREATE TABLE IF NOT EXISTS providers ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), slug TEXT NOT NULL UNIQUE CHECK (slug ~ '^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$'), name TEXT NOT NULL UNIQUE, protocol TEXT NOT NULL CHECK (protocol IN ('openai', 'anthropic')), - wire_api TEXT NOT NULL DEFAULT 'chat_completions' CHECK (wire_api IN ('chat_completions', 'responses', 'messages')), + wire_api TEXT NOT NULL DEFAULT 'chat_completions' CHECK (wire_api IN ('chat_completions', 'responses', 'embeddings', 'messages')), base_url TEXT NOT NULL, api_key_ciphertext BYTEA NOT NULL, enabled BOOLEAN NOT NULL DEFAULT TRUE, @@ -90,10 +101,10 @@ ALTER TABLE providers ADD CONSTRAINT providers_slug_check CHECK (slug ~ '^[a-z0- ALTER TABLE providers ADD COLUMN IF NOT EXISTS wire_api TEXT NOT NULL DEFAULT 'chat_completions'; UPDATE providers SET wire_api = 'messages' WHERE protocol = 'anthropic' AND wire_api = 'chat_completions'; ALTER TABLE providers DROP CONSTRAINT IF EXISTS providers_wire_api_check; -ALTER TABLE providers ADD CONSTRAINT providers_wire_api_check CHECK (wire_api IN ('chat_completions', 'responses', 'messages')); +ALTER TABLE providers ADD CONSTRAINT providers_wire_api_check CHECK (wire_api IN ('chat_completions', 'responses', 'embeddings', 'messages')); ALTER TABLE providers DROP CONSTRAINT IF EXISTS providers_protocol_wire_api_check; ALTER TABLE providers ADD CONSTRAINT providers_protocol_wire_api_check CHECK ( - (protocol = 'openai' AND wire_api IN ('chat_completions', 'responses')) OR + (protocol = 'openai' AND wire_api IN ('chat_completions', 'responses', 'embeddings')) OR (protocol = 'anthropic' AND wire_api = 'messages') ); @@ -204,6 +215,7 @@ CREATE TABLE IF NOT EXISTS model_routes ( ); CREATE INDEX IF NOT EXISTS api_keys_active_hash_idx ON api_keys (key_hash) WHERE status = 'active'; +CREATE INDEX IF NOT EXISTS api_keys_rotated_from_idx ON api_keys (rotated_from_id) WHERE rotated_from_id IS NOT NULL; CREATE INDEX IF NOT EXISTS projects_tenant_idx ON projects (tenant_id); CREATE INDEX IF NOT EXISTS model_routes_model_idx ON model_routes (model_id) WHERE enabled; CREATE INDEX IF NOT EXISTS model_routes_provider_idx ON model_routes (provider_id) WHERE enabled; @@ -224,9 +236,21 @@ CREATE TABLE IF NOT EXISTS tenant_preferences ( fallback_model TEXT REFERENCES models(public_id) ON DELETE SET NULL, low_balance_enabled BOOLEAN NOT NULL DEFAULT TRUE, low_balance_threshold_micros BIGINT NOT NULL DEFAULT 5000000 CHECK (low_balance_threshold_micros >= 0), + spend_anomaly_enabled BOOLEAN NOT NULL DEFAULT TRUE, + spend_anomaly_multiplier BIGINT, + spend_anomaly_min_micros BIGINT, updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), CHECK (default_model IS NULL OR fallback_model IS NULL OR default_model <> fallback_model) ); +ALTER TABLE tenant_preferences ADD COLUMN IF NOT EXISTS spend_anomaly_enabled BOOLEAN NOT NULL DEFAULT TRUE; +ALTER TABLE tenant_preferences ADD COLUMN IF NOT EXISTS spend_anomaly_multiplier BIGINT; +ALTER TABLE tenant_preferences ADD COLUMN IF NOT EXISTS spend_anomaly_min_micros BIGINT; +ALTER TABLE tenant_preferences DROP CONSTRAINT IF EXISTS tenant_preferences_spend_anomaly_multiplier_check; +ALTER TABLE tenant_preferences ADD CONSTRAINT tenant_preferences_spend_anomaly_multiplier_check + CHECK (spend_anomaly_multiplier IS NULL OR spend_anomaly_multiplier BETWEEN 2 AND 1000); +ALTER TABLE tenant_preferences DROP CONSTRAINT IF EXISTS tenant_preferences_spend_anomaly_min_check; +ALTER TABLE tenant_preferences ADD CONSTRAINT tenant_preferences_spend_anomaly_min_check + CHECK (spend_anomaly_min_micros IS NULL OR spend_anomaly_min_micros BETWEEN 0 AND 1000000000000000); CREATE TABLE IF NOT EXISTS billing_reservations ( request_id TEXT PRIMARY KEY, @@ -289,6 +313,7 @@ CREATE TABLE IF NOT EXISTS usage_events ( attempts INTEGER NOT NULL DEFAULT 0, started_at TIMESTAMPTZ NOT NULL, duration_ms BIGINT NOT NULL DEFAULT 0, + ttft_ms BIGINT NOT NULL DEFAULT 0 CHECK (ttft_ms >= 0), input_tokens BIGINT NOT NULL DEFAULT 0, output_tokens BIGINT NOT NULL DEFAULT 0, total_tokens BIGINT NOT NULL DEFAULT 0, @@ -299,15 +324,18 @@ CREATE TABLE IF NOT EXISTS usage_events ( uncollected_micros BIGINT NOT NULL DEFAULT 0, usage_reported BOOLEAN NOT NULL DEFAULT FALSE, metering_status TEXT NOT NULL DEFAULT 'not_billable' - CHECK (metering_status IN ('not_billable','reported','missing','upstream_failed')), + CHECK (metering_status IN ('not_billable','reported','missing','released_unmetered','upstream_failed')), created_at TIMESTAMPTZ NOT NULL DEFAULT now(), FOREIGN KEY (project_id, tenant_id) REFERENCES projects(id, tenant_id) ON DELETE RESTRICT ); ALTER TABLE usage_events ADD COLUMN IF NOT EXISTS usage_reported BOOLEAN NOT NULL DEFAULT FALSE; ALTER TABLE usage_events ADD COLUMN IF NOT EXISTS metering_status TEXT NOT NULL DEFAULT 'not_billable'; +ALTER TABLE usage_events ADD COLUMN IF NOT EXISTS ttft_ms BIGINT NOT NULL DEFAULT 0; +ALTER TABLE usage_events DROP CONSTRAINT IF EXISTS usage_events_ttft_ms_check; +ALTER TABLE usage_events ADD CONSTRAINT usage_events_ttft_ms_check CHECK (ttft_ms >= 0); ALTER TABLE usage_events DROP CONSTRAINT IF EXISTS usage_events_metering_status_check; ALTER TABLE usage_events ADD CONSTRAINT usage_events_metering_status_check - CHECK (metering_status IN ('not_billable','reported','missing','upstream_failed')); + CHECK (metering_status IN ('not_billable','reported','missing','released_unmetered','upstream_failed')); -- Usage persistence is independent from billing. Older installations created this -- foreign key, which prevented recording requests when prepaid billing was disabled. diff --git a/internal/controlplane/snapshot.go b/internal/controlplane/snapshot.go index 75a3a80..48be436 100644 --- a/internal/controlplane/snapshot.go +++ b/internal/controlplane/snapshot.go @@ -246,7 +246,7 @@ func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Prov func loadAPIKeys(ctx context.Context, tx pgx.Tx) ([]auth.HashedKeyRecord, error) { rows, err := tx.Query(ctx, ` SELECT k.id::text, k.key_hash, k.tenant_id::text, k.project_id::text, k.scopes, - k.monthly_spend_micros, k.expires_at, + k.monthly_spend_micros, k.daily_spend_micros, k.requests_per_minute, k.tokens_per_minute, k.expires_at, COALESCE(( SELECT jsonb_agg(m.public_id ORDER BY m.public_id) FROM api_key_model_restrictions r @@ -265,10 +265,10 @@ func loadAPIKeys(ctx context.Context, tx pgx.Tx) ([]auth.HashedKeyRecord, error) for rows.Next() { var keyID, tenantID, projectID string var hashBytes, scopesJSON, allowedModelsJSON []byte - var monthlySpendMicros int64 + var monthlySpendMicros, dailySpendMicros, requestsPerMinute, tokensPerMinute int64 var expiresAt *time.Time if err := rows.Scan(&keyID, &hashBytes, &tenantID, &projectID, &scopesJSON, - &monthlySpendMicros, &expiresAt, &allowedModelsJSON); err != nil { + &monthlySpendMicros, &dailySpendMicros, &requestsPerMinute, &tokensPerMinute, &expiresAt, &allowedModelsJSON); err != nil { return nil, fmt.Errorf("scan API key: %w", err) } if len(hashBytes) != sha256.Size { @@ -290,7 +290,8 @@ func loadAPIKeys(ctx context.Context, tx pgx.Tx) ([]auth.HashedKeyRecord, error) } records = append(records, auth.HashedKeyRecord{Hash: hash, Principal: domain.Principal{ KeyID: keyID, TenantID: tenantID, ProjectID: projectID, Scopes: scopes, - AllowedModels: allowedModels, MonthlySpendMicros: monthlySpendMicros, ExpiresAt: expiresAt, + AllowedModels: allowedModels, MonthlySpendMicros: monthlySpendMicros, DailySpendMicros: dailySpendMicros, + RequestsPerMinute: requestsPerMinute, TokensPerMinute: tokensPerMinute, ExpiresAt: expiresAt, }}) } if err := rows.Err(); err != nil { diff --git a/internal/controlplane/store.go b/internal/controlplane/store.go index c4d7016..82c87ae 100644 --- a/internal/controlplane/store.go +++ b/internal/controlplane/store.go @@ -23,7 +23,10 @@ var schemaSQL string var ErrRedisDisabled = errors.New("Redis propagation is disabled") -const migrationVersion int64 = 2026080605 +const ( + migrationVersion int64 = 2026080610 + migrationLockID int64 = 0x41494757 // "AIGW"; stable across migration versions. +) type Options struct { DatabaseURL string @@ -82,6 +85,13 @@ func (s *Store) RedisEnabled() bool { return s.redis != nil } +func (s *Store) PingRedis(ctx context.Context) error { + if s.redis == nil { + return ErrRedisDisabled + } + return s.redis.Ping(ctx).Err() +} + func (s *Store) Ping(ctx context.Context) error { return s.db.Ping(ctx) } func (s *Store) Migrate(ctx context.Context) error { @@ -112,28 +122,40 @@ func MigrationStatusDatabase(ctx context.Context, databaseURL string) (Migration } func applySchema(ctx context.Context, db *pgxpool.Pool) error { + hash := sha256.Sum256([]byte(schemaSQL)) + checksum := hex.EncodeToString(hash[:]) tx, err := db.Begin(ctx) if err != nil { return fmt.Errorf("begin migration: %w", err) } defer tx.Rollback(ctx) - if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock($1)`, migrationVersion); err != nil { + if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock($1)`, migrationLockID); err != nil { return err } - if _, err := tx.Exec(ctx, schemaSQL); err != nil { - return fmt.Errorf("apply control-plane schema: %w", err) + var migrationsExist bool + if err := tx.QueryRow(ctx, `SELECT to_regclass('schema_migrations') IS NOT NULL`).Scan(&migrationsExist); err != nil { + return fmt.Errorf("inspect migration table: %w", err) } - hash := sha256.Sum256([]byte(schemaSQL)) - checksum := hex.EncodeToString(hash[:]) - var existing string - err = tx.QueryRow(ctx, `SELECT checksum FROM schema_migrations WHERE version=$1`, migrationVersion).Scan(&existing) - if err == nil && existing != checksum { - return fmt.Errorf("migration %d checksum changed; deploy an explicit new migration version", migrationVersion) + if migrationsExist { + var existing string + err = tx.QueryRow(ctx, `SELECT checksum FROM schema_migrations WHERE version=$1`, migrationVersion).Scan(&existing) + if err == nil { + if existing != checksum { + return fmt.Errorf("migration %d checksum changed; deploy an explicit new migration version", migrationVersion) + } + if err := tx.Commit(ctx); err != nil { + return fmt.Errorf("commit migration check: %w", err) + } + return nil + } + if !errors.Is(err, pgx.ErrNoRows) { + return fmt.Errorf("read migration checksum: %w", err) + } } - if !errors.Is(err, pgx.ErrNoRows) && err != nil { - return err + if _, err := tx.Exec(ctx, schemaSQL); err != nil { + return fmt.Errorf("apply control-plane schema: %w", err) } - if _, err := tx.Exec(ctx, `INSERT INTO schema_migrations(version,name,checksum) VALUES ($1,$2,$3) ON CONFLICT DO NOTHING`, migrationVersion, "tenant-billing-profiles", checksum); err != nil { + if _, err := tx.Exec(ctx, `INSERT INTO schema_migrations(version,name,checksum) VALUES ($1,$2,$3) ON CONFLICT DO NOTHING`, migrationVersion, "tenant-spend-alert-preferences", checksum); err != nil { return err } if err := tx.Commit(ctx); err != nil { diff --git a/internal/controlplane/store_integration_test.go b/internal/controlplane/store_integration_test.go index 9bf093c..b0038c1 100644 --- a/internal/controlplane/store_integration_test.go +++ b/internal/controlplane/store_integration_test.go @@ -6,6 +6,7 @@ import ( "fmt" "net/url" "os" + "strings" "testing" "time" @@ -54,25 +55,30 @@ func TestAPIKeyRestrictionsRoundTripIntoRuntimeSnapshotPostgres(t *testing.T) { if err != nil { t.Fatal(err) } - provider, _, err := store.CreateProvider(ctx, CreateProviderInput{Name: "key-control-provider", Protocol: "openai", WireAPI: "responses", BaseURL: "https://example.invalid/v1", APIKey: "provider-secret"}) + provider, _, err := store.CreateProvider(ctx, CreateProviderInput{Name: "key-control-provider", Protocol: "openai", WireAPI: "embeddings", BaseURL: "https://example.invalid/v1", APIKey: "provider-secret"}) if err != nil { t.Fatal(err) } if provider.Slug != "key-control-provider" { t.Fatalf("derived provider slug = %q, want key-control-provider", provider.Slug) } - model, _, err := store.CreateModel(ctx, CreateModelInput{PublicID: "model/key-control", DisplayName: "Key Control", PriceCurrency: "usd", InputPriceMicrosPerMillion: 100_000, OutputPriceMicrosPerMillion: 200_000, Routes: []RouteInput{{ProviderID: provider.ID, UpstreamModel: "upstream-key-control", Weight: 1}}}) + if _, _, err := store.CreateProvider(ctx, CreateProviderInput{Name: "invalid-anthropic-embeddings", Protocol: "anthropic", WireAPI: "embeddings", BaseURL: "https://example.invalid/v1", APIKey: "provider-secret"}); err == nil { + t.Fatal("Anthropic provider accepted the OpenAI Embeddings wire API") + } + model, _, err := store.CreateModel(ctx, CreateModelInput{PublicID: "model/key-control", DisplayName: "Key Control", Capabilities: []string{"embeddings"}, PriceCurrency: "usd", InputPriceMicrosPerMillion: 100_000, OutputPriceMicrosPerMillion: 200_000, Routes: []RouteInput{{ProviderID: provider.ID, UpstreamModel: "upstream-key-control", Weight: 1}}}) if err != nil { t.Fatal(err) } expiresAt := time.Now().Add(24 * time.Hour).UTC().Truncate(time.Microsecond) created, _, err := store.CreateAPIKey(ctx, CreateAPIKeyInput{TenantID: tenant.ID, ProjectID: project.ID, Name: "production backend", Scopes: []string{"inference"}, Tags: []string{"production", "backend"}, - AllowedModels: []string{model.PublicID}, MonthlySpendMicros: 25_000_000, ExpiresAt: &expiresAt}) + AllowedModels: []string{model.PublicID}, MonthlySpendMicros: 25_000_000, DailySpendMicros: 5_000_000, + RequestsPerMinute: 12, TokensPerMinute: 34_000, ExpiresAt: &expiresAt}) if err != nil { t.Fatal(err) } - if created.Key == "" || created.MonthlySpendMicros != 25_000_000 || len(created.AllowedModels) != 1 || len(created.Tags) != 2 { + if created.Key == "" || len(created.KeySuffix) != 6 || !strings.HasSuffix(created.Key, created.KeySuffix) || + created.MonthlySpendMicros != 25_000_000 || created.DailySpendMicros != 5_000_000 || created.RequestsPerMinute != 12 || created.TokensPerMinute != 34_000 || len(created.AllowedModels) != 1 || len(created.Tags) != 2 { t.Fatalf("unexpected created key: %+v", created.APIKey) } if _, err := store.db.Exec(ctx, `INSERT INTO usage_events @@ -93,12 +99,15 @@ func TestAPIKeyRestrictionsRoundTripIntoRuntimeSnapshotPostgres(t *testing.T) { if err != nil { t.Fatal(err) } - if len(keys) != 1 || keys[0].AllowedModels[0] != model.PublicID || keys[0].ExpiresAt == nil || !keys[0].ExpiresAt.Equal(expiresAt) { + if len(keys) != 1 || keys[0].KeySuffix != created.KeySuffix || keys[0].AllowedModels[0] != model.PublicID || keys[0].ExpiresAt == nil || !keys[0].ExpiresAt.Equal(expiresAt) { t.Fatalf("key restrictions did not round trip: %+v", keys) } if keys[0].CurrentMonthSpendMicros != 42_000 || keys[0].CurrentMonthReservedMicros != 9_000 || keys[0].CurrentMonthRequests != 1 { t.Fatalf("key month activity is incorrect: %+v", keys[0]) } + if keys[0].CurrentDaySpendMicros != 42_000 || keys[0].CurrentDayReservedMicros != 9_000 || keys[0].CurrentDayRequests != 1 { + t.Fatalf("key daily activity is incorrect: %+v", keys[0]) + } snapshot, err := store.LoadSnapshot(ctx) if err != nil { t.Fatal(err) @@ -106,8 +115,11 @@ func TestAPIKeyRestrictionsRoundTripIntoRuntimeSnapshotPostgres(t *testing.T) { if len(snapshot.APIKeys) != 1 { t.Fatalf("snapshot API keys = %d, want 1", len(snapshot.APIKeys)) } + if len(snapshot.Models) != 1 || len(snapshot.Models[0].Routes) != 1 || snapshot.Models[0].Routes[0].Provider.EffectiveWireAPI() != "embeddings" || len(snapshot.Models[0].Capabilities) != 1 || snapshot.Models[0].Capabilities[0] != "embeddings" { + t.Fatalf("Embeddings model did not round trip into runtime snapshot: %+v", snapshot.Models) + } principal := snapshot.APIKeys[0].Principal - if principal.MonthlySpendMicros != 25_000_000 || principal.ExpiresAt == nil { + if principal.MonthlySpendMicros != 25_000_000 || principal.DailySpendMicros != 5_000_000 || principal.RequestsPerMinute != 12 || principal.TokensPerMinute != 34_000 || principal.ExpiresAt == nil { t.Fatalf("snapshot lost API key controls: %+v", principal) } if _, ok := principal.AllowedModels[model.PublicID]; !ok { @@ -132,6 +144,38 @@ func TestAPIKeyRestrictionsRoundTripIntoRuntimeSnapshotPostgres(t *testing.T) { if err != nil || len(keys) != 1 { t.Fatalf("failed key transaction leaked a row: keys=%d err=%v", len(keys), err) } + if _, err := store.DisableAPIKey(ctx, created.ID); err != nil { + t.Fatal(err) + } + snapshot, err = store.LoadSnapshot(ctx) + if err != nil { + t.Fatal(err) + } + if len(snapshot.APIKeys) != 0 { + t.Fatalf("disabled key remained in runtime snapshot: %+v", snapshot.APIKeys) + } + if _, err := store.EnableAPIKey(ctx, created.ID); err != nil { + t.Fatal(err) + } + rotated, _, err := store.RotateAPIKey(ctx, created.ID) + if err != nil { + t.Fatal(err) + } + if rotated.ID == created.ID || rotated.Key == "" || rotated.Key == created.Key || len(rotated.KeySuffix) != 6 || + !strings.HasSuffix(rotated.Key, rotated.KeySuffix) || rotated.KeySuffix == created.KeySuffix || rotated.DailySpendMicros != created.DailySpendMicros || rotated.RequestsPerMinute != created.RequestsPerMinute || len(rotated.AllowedModels) != 1 || rotated.AllowedModels[0] != model.PublicID { + t.Fatalf("rotated key did not preserve controls: old=%+v new=%+v", created.APIKey, rotated.APIKey) + } + snapshot, err = store.LoadSnapshot(ctx) + if err != nil { + t.Fatal(err) + } + if len(snapshot.APIKeys) != 1 || snapshot.APIKeys[0].Principal.KeyID != rotated.ID { + t.Fatalf("rotation snapshot = %+v, want only new key %s", snapshot.APIKeys, rotated.ID) + } + keys, err = store.ListAPIKeysFor(ctx, tenant.ID) + if err != nil || len(keys) != 2 || keys[0].Status != "active" || keys[1].Status != "revoked" { + t.Fatalf("rotation lifecycle rows are incorrect: keys=%+v err=%v", keys, err) + } } func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) { @@ -192,10 +236,21 @@ func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) { t.Fatal(err) } defer scopedDB.Close() - var tenantID, providerID, modelID string + var tenantID, projectID, providerID, modelID string if err := scopedDB.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ('preference-test','Preference Test') RETURNING id::text`).Scan(&tenantID); err != nil { t.Fatal(err) } + if err := scopedDB.QueryRow(ctx, `INSERT INTO projects (tenant_id,slug,name) VALUES ($1,'default','Default') RETURNING id::text`, tenantID).Scan(&projectID); err != nil { + t.Fatal(err) + } + if _, err := scopedDB.Exec(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_hash) + VALUES ($1,$2,'Legacy key','sk-aigw-legacy...',decode(repeat('01',32),'hex'))`, tenantID, projectID); err != nil { + t.Fatalf("legacy key without suffix did not retain its empty default: %v", err) + } + if _, err := scopedDB.Exec(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_suffix,key_hash) + VALUES ($1,$2,'Invalid display suffix','sk-aigw-invalid...','too-long',decode(repeat('02',32),'hex'))`, tenantID, projectID); err == nil { + t.Fatal("invalid API key display suffix was accepted") + } if err := scopedDB.QueryRow(ctx, `INSERT INTO providers (slug,name,protocol,wire_api,base_url,api_key_ciphertext) VALUES ('preference-provider','Preference provider','openai','responses','https://example.invalid',decode('00','hex')) RETURNING id::text`).Scan(&providerID); err != nil { t.Fatal(err) @@ -210,7 +265,8 @@ func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) { t.Fatal(err) } store := &Store{db: scopedDB} - prefs, err := store.SetDeveloperPreferences(ctx, SetDeveloperPreferencesInput{TenantID: tenantID, DefaultModel: "model/preference-test"}) + preferenceDefaults := BillingPreferenceDefaults{LowBalanceThresholdMicros: 5_000_000, SpendAnomalyMultiplier: 3, SpendAnomalyMinMicros: 10_000_000} + prefs, err := store.SetDeveloperPreferences(ctx, SetDeveloperPreferencesInput{TenantID: tenantID, DefaultModel: "model/preference-test"}, preferenceDefaults) if err != nil { t.Fatal(err) } @@ -219,15 +275,21 @@ func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) { } enabled := false threshold := int64(9_750_000) + anomalyEnabled := false + anomalyMultiplier := int64(7) + anomalyMinimum := int64(8_250_000) if _, err := store.SetBillingPreferences(ctx, SetBillingPreferencesInput{TenantID: tenantID, - LowBalanceEnabled: &enabled, LowBalanceThresholdMicros: &threshold}, 5_000_000); err != nil { + LowBalanceEnabled: &enabled, LowBalanceThresholdMicros: &threshold, SpendAnomalyEnabled: &anomalyEnabled, + SpendAnomalyMultiplier: &anomalyMultiplier, SpendAnomalyMinMicros: &anomalyMinimum}, preferenceDefaults); err != nil { t.Fatal(err) } - prefs, err = store.GetTenantPreferences(ctx, tenantID, 5_000_000) + prefs, err = store.GetTenantPreferences(ctx, tenantID, preferenceDefaults) if err != nil { t.Fatal(err) } - if prefs.LowBalanceEnabled || prefs.LowBalanceThresholdMicros != threshold || prefs.DefaultModel != "model/preference-test" { + if prefs.LowBalanceEnabled || prefs.LowBalanceThresholdMicros != threshold || prefs.SpendAnomalyEnabled || + prefs.SpendAnomalyMultiplier != anomalyMultiplier || prefs.SpendAnomalyMinMicros != anomalyMinimum || + prefs.DefaultModel != "model/preference-test" { t.Fatalf("preferences did not round trip: %+v", prefs) } } diff --git a/internal/controlplane/types.go b/internal/controlplane/types.go index f6402d8..995fc8d 100644 --- a/internal/controlplane/types.go +++ b/internal/controlplane/types.go @@ -36,13 +36,20 @@ type APIKey struct { ProjectID string `json:"project_id"` Name string `json:"name"` KeyPrefix string `json:"key_prefix"` + KeySuffix string `json:"key_suffix"` Scopes []string `json:"scopes"` Tags []string `json:"tags"` AllowedModels []string `json:"allowed_models"` MonthlySpendMicros int64 `json:"monthly_spend_micros"` + DailySpendMicros int64 `json:"daily_spend_micros"` + RequestsPerMinute int64 `json:"requests_per_minute"` + TokensPerMinute int64 `json:"tokens_per_minute"` CurrentMonthSpendMicros int64 `json:"current_month_spend_micros"` CurrentMonthReservedMicros int64 `json:"current_month_reserved_micros"` CurrentMonthRequests int64 `json:"current_month_requests"` + CurrentDaySpendMicros int64 `json:"current_day_spend_micros"` + CurrentDayReservedMicros int64 `json:"current_day_reserved_micros"` + CurrentDayRequests int64 `json:"current_day_requests"` Status string `json:"status"` ExpiresAt *time.Time `json:"expires_at,omitempty"` LastUsedAt *time.Time `json:"last_used_at,omitempty"` @@ -117,11 +124,17 @@ type DeveloperProviderHealth struct { WireAPI string `json:"wire_api"` State string `json:"state"` Attempts uint64 `json:"attempts"` + ActiveProbes uint64 `json:"active_probes"` RecentSamples int `json:"recent_samples"` AvailabilityPercent float64 `json:"availability_percent"` HeaderLatencyEWMA int64 `json:"header_latency_ewma_ms"` + TTFTSamples uint64 `json:"ttft_samples"` + TTFTEWMA int64 `json:"ttft_ewma_ms"` + SharedAttempts uint64 `json:"shared_attempts"` + SharedTTFTSamples uint64 `json:"shared_ttft_samples"` ConsecutiveFailures uint64 `json:"consecutive_failures"` LastObservedAt *time.Time `json:"last_observed_at,omitempty"` + LastProbeAt *time.Time `json:"last_probe_at,omitempty"` CircuitOpenUntil *time.Time `json:"circuit_open_until,omitempty"` } @@ -160,9 +173,18 @@ type TenantPreferences struct { FallbackModel string `json:"fallback_model,omitempty"` LowBalanceEnabled bool `json:"low_balance_enabled"` LowBalanceThresholdMicros int64 `json:"low_balance_threshold_micros"` + SpendAnomalyEnabled bool `json:"spend_anomaly_enabled"` + SpendAnomalyMultiplier int64 `json:"spend_anomaly_multiplier"` + SpendAnomalyMinMicros int64 `json:"spend_anomaly_min_micros"` UpdatedAt *time.Time `json:"updated_at,omitempty"` } +type BillingPreferenceDefaults struct { + LowBalanceThresholdMicros int64 + SpendAnomalyMultiplier int64 + SpendAnomalyMinMicros int64 +} + type SetDeveloperPreferencesInput struct { TenantID string `json:"tenant_id"` DefaultModel string `json:"default_model"` @@ -173,6 +195,9 @@ type SetBillingPreferencesInput struct { TenantID string `json:"tenant_id"` LowBalanceEnabled *bool `json:"low_balance_enabled"` LowBalanceThresholdMicros *int64 `json:"low_balance_threshold_micros"` + SpendAnomalyEnabled *bool `json:"spend_anomaly_enabled"` + SpendAnomalyMultiplier *int64 `json:"spend_anomaly_multiplier"` + SpendAnomalyMinMicros *int64 `json:"spend_anomaly_min_micros"` } type Model struct { @@ -260,6 +285,9 @@ type CreateAPIKeyInput struct { Tags []string `json:"tags"` AllowedModels []string `json:"allowed_models"` MonthlySpendMicros int64 `json:"monthly_spend_micros"` + DailySpendMicros int64 `json:"daily_spend_micros"` + RequestsPerMinute int64 `json:"requests_per_minute"` + TokensPerMinute int64 `json:"tokens_per_minute"` ExpiresAt *time.Time `json:"expires_at"` } @@ -469,6 +497,7 @@ type UsageRecord struct { Attempts int `json:"attempts"` StartedAt time.Time `json:"started_at"` DurationMS int64 `json:"duration_ms"` + TTFTMS int64 `json:"ttft_ms"` InputTokens int64 `json:"input_tokens"` OutputTokens int64 `json:"output_tokens"` TotalTokens int64 `json:"total_tokens"` @@ -481,6 +510,11 @@ type UsageRecord struct { MeteringStatus string `json:"metering_status"` } +type UsagePage struct { + Data []UsageRecord `json:"data"` + NextCursor string `json:"next_cursor,omitempty"` +} + type UsageDailyPoint struct { Day time.Time `json:"day"` RequestCount int64 `json:"request_count"` @@ -491,7 +525,10 @@ type UsageDailyPoint struct { ChargedMicros int64 `json:"charged_micros"` UncollectedMicros int64 `json:"uncollected_micros"` AverageDurationMS int64 `json:"average_duration_ms"` + P50DurationMS int64 `json:"p50_duration_ms"` P95DurationMS int64 `json:"p95_duration_ms"` + P50TTFTMS int64 `json:"p50_ttft_ms"` + P95TTFTMS int64 `json:"p95_ttft_ms"` } // UsageAnalytics is a persisted-ledger aggregation used by the developer @@ -501,6 +538,7 @@ type UsageAnalytics struct { RangeEnd time.Time `json:"range_end"` Models []UsageModelAnalytics `json:"models"` Providers []UsageProviderAnalytics `json:"providers"` + Keys []UsageKeyAnalytics `json:"keys"` } type UsageModelAnalytics struct { @@ -518,7 +556,10 @@ type UsageModelAnalytics struct { UncollectedMicros int64 `json:"uncollected_micros"` MissingUsageRequests int64 `json:"missing_usage_requests"` AverageDurationMS int64 `json:"average_duration_ms"` + P50DurationMS int64 `json:"p50_duration_ms"` P95DurationMS int64 `json:"p95_duration_ms"` + P50TTFTMS int64 `json:"p50_ttft_ms"` + P95TTFTMS int64 `json:"p95_ttft_ms"` PreviousChargedMicros int64 `json:"previous_charged_micros"` ChargeChangePercent *float64 `json:"charge_change_percent,omitempty"` } @@ -540,11 +581,34 @@ type UsageProviderAnalytics struct { UncollectedMicros int64 `json:"uncollected_micros"` MissingUsageRequests int64 `json:"missing_usage_requests"` AverageDurationMS int64 `json:"average_duration_ms"` + P50DurationMS int64 `json:"p50_duration_ms"` P95DurationMS int64 `json:"p95_duration_ms"` + P50TTFTMS int64 `json:"p50_ttft_ms"` + P95TTFTMS int64 `json:"p95_ttft_ms"` PreviousChargedMicros int64 `json:"previous_charged_micros"` ChargeChangePercent *float64 `json:"charge_change_percent,omitempty"` } +type UsageKeyAnalytics struct { + KeyID string `json:"key_id"` + KeyName string `json:"key_name"` + RequestCount int64 `json:"request_count"` + SuccessfulRequests int64 `json:"successful_requests"` + ErrorCount int64 `json:"error_count"` + ModelCount int64 `json:"model_count"` + InputTokens int64 `json:"input_tokens"` + OutputTokens int64 `json:"output_tokens"` + TotalTokens int64 `json:"total_tokens"` + ChargedMicros int64 `json:"charged_micros"` + UncollectedMicros int64 `json:"uncollected_micros"` + MissingUsageRequests int64 `json:"missing_usage_requests"` + AverageDurationMS int64 `json:"average_duration_ms"` + P50DurationMS int64 `json:"p50_duration_ms"` + P95DurationMS int64 `json:"p95_duration_ms"` + P50TTFTMS int64 `json:"p50_ttft_ms"` + P95TTFTMS int64 `json:"p95_ttft_ms"` +} + type UsageSummary struct { PeriodStart time.Time `json:"period_start"` TenantID string `json:"tenant_id"` diff --git a/internal/controlplane/usage.go b/internal/controlplane/usage.go index ebabaf6..05bc155 100644 --- a/internal/controlplane/usage.go +++ b/internal/controlplane/usage.go @@ -2,6 +2,9 @@ package controlplane import ( "context" + "encoding/base64" + "encoding/json" + "errors" "fmt" "strings" "time" @@ -25,6 +28,14 @@ type UsageQuery struct { From time.Time To time.Time Limit int + Cursor string +} + +var ErrInvalidUsageCursor = errors.New("invalid usage cursor") + +type usageCursor struct { + StartedAt time.Time `json:"started_at"` + RequestID string `json:"request_id"` } func (s *Store) RecordUsage(ctx context.Context, event domain.UsageEvent) error { @@ -37,13 +48,13 @@ func (s *Store) RecordUsage(ctx context.Context, event domain.UsageEvent) error INSERT INTO usage_events ( request_id, tenant_id, project_id, key_id, public_model, provider_id, upstream_model, protocol, stream, status_code, success, error_type, attempts, started_at, duration_ms, - input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, + ttft_ms, input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, cost_micros, charged_micros, uncollected_micros, usage_reported, metering_status) - VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,0,0,0,$21,$22) + VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,0,0,0,$22,$23) ON CONFLICT (request_id) DO NOTHING`, event.RequestID, event.TenantID, event.ProjectID, event.KeyID, event.PublicModel, event.ProviderID, event.UpstreamModel, string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.ErrorType, event.Attempts, event.StartedAt, event.DurationMS, - event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, + event.TTFTMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens, event.UsageReported, usageMeteringStatus(event)) if err != nil { return fmt.Errorf("persist usage event: %w", err) @@ -104,13 +115,13 @@ func boolInt(value bool) int { return 0 } -func (s *Store) ListUsage(ctx context.Context, query UsageQuery) ([]UsageRecord, error) { +func (s *Store) ListUsage(ctx context.Context, query UsageQuery) (UsagePage, error) { limit := query.Limit if limit < 1 || limit > 1000 { limit = 200 } where := []string{"1=1"} - args := make([]any, 0, 13) + args := make([]any, 0, 16) index := 1 for _, item := range []struct{ value, clause string }{{query.TenantID, "tenant_id=$"}, {query.ProjectID, "project_id=$"}, {query.KeyID, "key_id=$"}, {query.Model, "public_model=$"}, {query.RequestID, "request_id=$"}} { if strings.TrimSpace(item.value) != "" { @@ -151,31 +162,68 @@ func (s *Store) ListUsage(ctx context.Context, query UsageQuery) ([]UsageRecord, args = append(args, query.To) index++ } - args = append(args, limit) + if strings.TrimSpace(query.Cursor) != "" { + cursor, err := decodeUsageCursor(query.Cursor) + if err != nil { + return UsagePage{}, err + } + where = append(where, "(started_at, request_id) < ($"+fmt.Sprint(index)+",$"+fmt.Sprint(index+1)+")") + args = append(args, cursor.StartedAt, cursor.RequestID) + index += 2 + } + args = append(args, limit+1) rows, err := s.db.Query(ctx, `SELECT request_id, tenant_id::text, project_id::text, COALESCE((SELECT name FROM projects p WHERE p.id=usage_events.project_id),''), key_id::text, COALESCE((SELECT name FROM api_keys k WHERE k.id=usage_events.key_id),''), public_model, COALESCE(provider_id,''), COALESCE((SELECT name FROM providers p WHERE p.id::text=usage_events.provider_id),''), COALESCE(upstream_model,''), protocol, stream, status_code, success, error_type, - attempts, started_at, duration_ms, input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, + attempts, started_at, duration_ms, ttft_ms, input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, cost_micros, charged_micros, uncollected_micros, usage_reported, metering_status FROM usage_events WHERE `+ - strings.Join(where, " AND ")+` ORDER BY created_at DESC LIMIT $`+fmt.Sprint(index), args...) + strings.Join(where, " AND ")+` ORDER BY started_at DESC, request_id DESC LIMIT $`+fmt.Sprint(index), args...) if err != nil { - return nil, fmt.Errorf("query usage events: %w", err) + return UsagePage{}, fmt.Errorf("query usage events: %w", err) } defer rows.Close() - result := make([]UsageRecord, 0) + result := make([]UsageRecord, 0, limit+1) for rows.Next() { var item UsageRecord if err := rows.Scan(&item.RequestID, &item.TenantID, &item.ProjectID, &item.ProjectName, &item.KeyID, &item.KeyName, &item.PublicModel, &item.ProviderID, &item.ProviderName, &item.UpstreamModel, &item.Protocol, &item.Stream, &item.StatusCode, &item.Success, &item.ErrorType, &item.Attempts, &item.StartedAt, &item.DurationMS, - &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.CacheCreationInputTokens, &item.CacheReadInputTokens, + &item.TTFTMS, &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.CacheCreationInputTokens, &item.CacheReadInputTokens, &item.CostMicros, &item.ChargedMicros, &item.UncollectedMicros, &item.UsageReported, &item.MeteringStatus); err != nil { - return nil, fmt.Errorf("scan usage event: %w", err) + return UsagePage{}, fmt.Errorf("scan usage event: %w", err) } result = append(result, item) } - return result, rows.Err() + if err := rows.Err(); err != nil { + return UsagePage{}, err + } + page := UsagePage{Data: result} + if len(result) > limit { + page.Data = result[:limit] + page.NextCursor = encodeUsageCursor(page.Data[len(page.Data)-1]) + } + return page, nil +} + +func encodeUsageCursor(record UsageRecord) string { + payload, _ := json.Marshal(usageCursor{StartedAt: record.StartedAt.UTC(), RequestID: record.RequestID}) + return base64.RawURLEncoding.EncodeToString(payload) +} + +func decodeUsageCursor(raw string) (usageCursor, error) { + if len(raw) > 2048 { + return usageCursor{}, ErrInvalidUsageCursor + } + payload, err := base64.RawURLEncoding.DecodeString(raw) + if err != nil { + return usageCursor{}, ErrInvalidUsageCursor + } + var cursor usageCursor + if json.Unmarshal(payload, &cursor) != nil || cursor.StartedAt.IsZero() || strings.TrimSpace(cursor.RequestID) == "" { + return usageCursor{}, ErrInvalidUsageCursor + } + return cursor, nil } func (s *Store) UsageDaily(ctx context.Context, query UsageQuery) ([]UsageDailyPoint, error) { @@ -225,7 +273,10 @@ func (s *Store) UsageDaily(ctx context.Context, query UsageQuery) ([]UsageDailyP count(*) FILTER (WHERE success), COALESCE(sum(input_tokens),0), COALESCE(sum(output_tokens),0), COALESCE(sum(total_tokens),0), COALESCE(sum(charged_micros),0), COALESCE(sum(uncollected_micros),0), COALESCE(round(avg(duration_ms)),0)::bigint, - COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY duration_ms)),0)::bigint + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY ttft_ms) FILTER (WHERE ttft_ms > 0)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY ttft_ms) FILTER (WHERE ttft_ms > 0)),0)::bigint FROM usage_events WHERE `+strings.Join(where, " AND ")+` GROUP BY 1 ORDER BY 1`, args...) if err != nil { return nil, fmt.Errorf("query daily usage: %w", err) @@ -235,7 +286,8 @@ func (s *Store) UsageDaily(ctx context.Context, query UsageQuery) ([]UsageDailyP for rows.Next() { var item UsageDailyPoint if err := rows.Scan(&item.Day, &item.RequestCount, &item.SuccessfulRequests, &item.InputTokens, &item.OutputTokens, - &item.TotalTokens, &item.ChargedMicros, &item.UncollectedMicros, &item.AverageDurationMS, &item.P95DurationMS); err != nil { + &item.TotalTokens, &item.ChargedMicros, &item.UncollectedMicros, &item.AverageDurationMS, &item.P50DurationMS, + &item.P95DurationMS, &item.P50TTFTMS, &item.P95TTFTMS); err != nil { return nil, fmt.Errorf("scan daily usage: %w", err) } result = append(result, item) diff --git a/internal/controlplane/usage_analytics.go b/internal/controlplane/usage_analytics.go index 7cc042f..1bc6228 100644 --- a/internal/controlplane/usage_analytics.go +++ b/internal/controlplane/usage_analytics.go @@ -22,7 +22,13 @@ func (s *Store) UsageAnalytics(ctx context.Context, query UsageQuery) (UsageAnal if !to.After(from) { return UsageAnalytics{}, fmt.Errorf("usage analytics range must be positive") } - result := UsageAnalytics{RangeStart: from, RangeEnd: to, Models: make([]UsageModelAnalytics, 0), Providers: make([]UsageProviderAnalytics, 0)} + result := UsageAnalytics{ + RangeStart: from, + RangeEnd: to, + Models: make([]UsageModelAnalytics, 0), + Providers: make([]UsageProviderAnalytics, 0), + Keys: make([]UsageKeyAnalytics, 0), + } modelPrevious, err := s.usageModelCharges(ctx, query, from.Add(-to.Sub(from)), from) if err != nil { @@ -49,8 +55,13 @@ func (s *Store) UsageAnalytics(ctx context.Context, query UsageQuery) (UsageAnal providers[index].PreviousChargedMicros = providerPrevious[providers[index].ProviderID] providers[index].ChargeChangePercent = chargeChange(providers[index].ChargedMicros, providers[index].PreviousChargedMicros) } + keys, err := s.usageKeyAnalytics(ctx, query, from, to) + if err != nil { + return UsageAnalytics{}, err + } result.Models = models result.Providers = providers + result.Keys = keys return result, nil } @@ -128,7 +139,11 @@ func (s *Store) usageModelAnalytics(ctx context.Context, query UsageQuery, from, count(DISTINCT NULLIF(e.provider_id,'')), COALESCE(sum(e.input_tokens),0), COALESCE(sum(e.output_tokens),0), COALESCE(sum(e.total_tokens),0), COALESCE(sum(e.cache_read_input_tokens),0), COALESCE(sum(e.cache_creation_input_tokens),0), COALESCE(sum(e.charged_micros),0), COALESCE(sum(e.uncollected_micros),0), count(*) FILTER (WHERE e.metering_status='missing'), - COALESCE(round(avg(e.duration_ms)),0)::bigint, COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint + COALESCE(round(avg(e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY e.ttft_ms) FILTER (WHERE e.ttft_ms > 0)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.ttft_ms) FILTER (WHERE e.ttft_ms > 0)),0)::bigint FROM usage_events e WHERE `+where+` GROUP BY e.public_model ORDER BY sum(e.charged_micros) DESC, e.public_model`, args...) if err != nil { return nil, fmt.Errorf("query usage model analytics: %w", err) @@ -139,7 +154,8 @@ func (s *Store) usageModelAnalytics(ctx context.Context, query UsageQuery, from, var item UsageModelAnalytics if err := rows.Scan(&item.PublicModel, &item.RequestCount, &item.SuccessfulRequests, &item.ErrorCount, &item.ProviderCount, &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.CacheReadInputTokens, &item.CacheCreationInputTokens, - &item.ChargedMicros, &item.UncollectedMicros, &item.MissingUsageRequests, &item.AverageDurationMS, &item.P95DurationMS); err != nil { + &item.ChargedMicros, &item.UncollectedMicros, &item.MissingUsageRequests, &item.AverageDurationMS, + &item.P50DurationMS, &item.P95DurationMS, &item.P50TTFTMS, &item.P95TTFTMS); err != nil { return nil, fmt.Errorf("scan usage model analytics: %w", err) } result = append(result, item) @@ -174,7 +190,10 @@ func (s *Store) usageProviderAnalytics(ctx context.Context, query UsageQuery, fr COALESCE(sum(e.input_tokens),0), COALESCE(sum(e.output_tokens),0), COALESCE(sum(e.total_tokens),0), COALESCE(sum(e.cache_read_input_tokens),0), COALESCE(sum(e.cache_creation_input_tokens),0), COALESCE(sum(e.charged_micros),0), COALESCE(sum(e.uncollected_micros),0), count(*) FILTER (WHERE e.metering_status='missing'), COALESCE(round(avg(e.duration_ms)),0)::bigint, - COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY e.ttft_ms) FILTER (WHERE e.ttft_ms > 0)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.ttft_ms) FILTER (WHERE e.ttft_ms > 0)),0)::bigint FROM usage_events e LEFT JOIN providers p ON p.id::text=e.provider_id WHERE `+where+` GROUP BY e.provider_id, p.name, p.wire_api ORDER BY sum(e.charged_micros) DESC, COALESCE(NULLIF(p.name,''),'Unassigned')`, args...) if err != nil { @@ -186,7 +205,8 @@ func (s *Store) usageProviderAnalytics(ctx context.Context, query UsageQuery, fr var item UsageProviderAnalytics if err := rows.Scan(&item.ProviderID, &item.ProviderName, &item.WireAPI, &item.RequestCount, &item.SuccessfulRequests, &item.ErrorCount, &item.ModelCount, &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.CacheReadInputTokens, &item.CacheCreationInputTokens, - &item.ChargedMicros, &item.UncollectedMicros, &item.MissingUsageRequests, &item.AverageDurationMS, &item.P95DurationMS); err != nil { + &item.ChargedMicros, &item.UncollectedMicros, &item.MissingUsageRequests, &item.AverageDurationMS, + &item.P50DurationMS, &item.P95DurationMS, &item.P50TTFTMS, &item.P95TTFTMS); err != nil { return nil, fmt.Errorf("scan usage provider analytics: %w", err) } result = append(result, item) @@ -194,6 +214,37 @@ func (s *Store) usageProviderAnalytics(ctx context.Context, query UsageQuery, fr return result, rows.Err() } +func (s *Store) usageKeyAnalytics(ctx context.Context, query UsageQuery, from, to time.Time) ([]UsageKeyAnalytics, error) { + where, args := analyticsUsageWhere(query, from, to) + rows, err := s.db.Query(ctx, `SELECT e.key_id::text, COALESCE(NULLIF(k.name,''),'Deleted key'), + count(*), count(*) FILTER (WHERE e.success), count(*) FILTER (WHERE NOT e.success), count(DISTINCT e.public_model), + COALESCE(sum(e.input_tokens),0), COALESCE(sum(e.output_tokens),0), COALESCE(sum(e.total_tokens),0), + COALESCE(sum(e.charged_micros),0), COALESCE(sum(e.uncollected_micros),0), count(*) FILTER (WHERE e.metering_status='missing'), + COALESCE(round(avg(e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY e.ttft_ms) FILTER (WHERE e.ttft_ms > 0)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.ttft_ms) FILTER (WHERE e.ttft_ms > 0)),0)::bigint + FROM usage_events e LEFT JOIN api_keys k ON k.id=e.key_id WHERE `+where+` + GROUP BY e.key_id, k.name ORDER BY sum(e.charged_micros) DESC, COALESCE(NULLIF(k.name,''),'Deleted key')`, args...) + if err != nil { + return nil, fmt.Errorf("query usage API key analytics: %w", err) + } + defer rows.Close() + result := make([]UsageKeyAnalytics, 0) + for rows.Next() { + var item UsageKeyAnalytics + if err := rows.Scan(&item.KeyID, &item.KeyName, &item.RequestCount, &item.SuccessfulRequests, &item.ErrorCount, + &item.ModelCount, &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.ChargedMicros, + &item.UncollectedMicros, &item.MissingUsageRequests, &item.AverageDurationMS, &item.P50DurationMS, + &item.P95DurationMS, &item.P50TTFTMS, &item.P95TTFTMS); err != nil { + return nil, fmt.Errorf("scan usage API key analytics: %w", err) + } + result = append(result, item) + } + return result, rows.Err() +} + func (s *Store) usageProviderCharges(ctx context.Context, query UsageQuery, from, to time.Time) (map[string]int64, error) { where, args := analyticsUsageWhere(query, from, to) rows, err := s.db.Query(ctx, `SELECT COALESCE(e.provider_id,''), COALESCE(sum(e.charged_micros),0) diff --git a/internal/controlplane/usage_integration_test.go b/internal/controlplane/usage_integration_test.go index 2545c1d..10305bf 100644 --- a/internal/controlplane/usage_integration_test.go +++ b/internal/controlplane/usage_integration_test.go @@ -78,7 +78,7 @@ func TestUsageFiltersAndDailyAggregationPostgres(t *testing.T) { started := time.Now().UTC().Add(-time.Hour).Truncate(time.Second) events := []domain.UsageEvent{ - {RequestID: fmt.Sprintf("req_usage_ok_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started, DurationMS: 120, Usage: domain.Usage{InputTokens: 10, OutputTokens: 4, TotalTokens: 14, CacheReadInputTokens: 5}}, + {RequestID: fmt.Sprintf("req_usage_ok_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started, DurationMS: 120, TTFTMS: 47, Usage: domain.Usage{InputTokens: 10, OutputTokens: 4, TotalTokens: 14, CacheReadInputTokens: 5}}, {RequestID: fmt.Sprintf("req_usage_error_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAIResponses, Stream: true, StatusCode: 502, Success: false, ErrorType: "provider_error", StartedAt: started.Add(time.Minute), DurationMS: 350}, } for _, event := range events { @@ -94,7 +94,7 @@ func TestUsageFiltersAndDailyAggregationPostgres(t *testing.T) { if err != nil { t.Fatal(err) } - if len(records) != 1 || records[0].RequestID != events[0].RequestID || records[0].ProjectName != "Production" || records[0].KeyName != "Production key" || records[0].MeteringStatus != "reported" { + if len(records.Data) != 1 || records.Data[0].RequestID != events[0].RequestID || records.Data[0].ProjectName != "Production" || records.Data[0].KeyName != "Production key" || records.Data[0].MeteringStatus != "reported" || records.Data[0].TTFTMS != 47 { t.Fatalf("unexpected filtered usage: %+v", records) } streaming := true @@ -103,7 +103,7 @@ func TestUsageFiltersAndDailyAggregationPostgres(t *testing.T) { if err != nil { t.Fatal(err) } - if len(records) != 1 || records[0].RequestID != events[1].RequestID { + if len(records.Data) != 1 || records.Data[0].RequestID != events[1].RequestID { t.Fatalf("unexpected provider/protocol/transport usage filter: %+v", records) } filteredPoints, err := store.UsageDaily(ctx, UsageQuery{TenantID: tenantID, Provider: providerSlug, Protocol: string(domain.ProtocolOpenAIResponses), ErrorType: "provider_error", Stream: &streaming, From: started.Add(-time.Minute), To: started.Add(time.Hour)}) @@ -117,11 +117,11 @@ func TestUsageFiltersAndDailyAggregationPostgres(t *testing.T) { if err != nil { t.Fatal(err) } - if len(points) != 1 || points[0].RequestCount != 2 || points[0].SuccessfulRequests != 1 || points[0].TotalTokens != 14 || points[0].ChargedMicros != 125 || points[0].P95DurationMS < 120 { + if len(points) != 1 || points[0].RequestCount != 2 || points[0].SuccessfulRequests != 1 || points[0].TotalTokens != 14 || points[0].ChargedMicros != 125 || points[0].P50DurationMS != 235 || points[0].P95DurationMS < 120 || points[0].P50TTFTMS != 47 || points[0].P95TTFTMS != 47 { t.Fatalf("unexpected daily usage: %+v", points) } - previous := domain.UsageEvent{RequestID: fmt.Sprintf("req_usage_previous_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started.Add(-5 * time.Minute), DurationMS: 80, Usage: domain.Usage{InputTokens: 3, OutputTokens: 2, TotalTokens: 5}} + previous := domain.UsageEvent{RequestID: fmt.Sprintf("req_usage_previous_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started.Add(-5 * time.Minute), DurationMS: 80, TTFTMS: 30, Usage: domain.Usage{InputTokens: 3, OutputTokens: 2, TotalTokens: 5}} missing := domain.UsageEvent{RequestID: fmt.Sprintf("req_usage_missing_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, StartedAt: started.Add(2 * time.Minute), DurationMS: 200} otherTenantEvent := domain.UsageEvent{RequestID: fmt.Sprintf("req_usage_other_%d", suffix), TenantID: otherTenantID, ProjectID: otherProjectID, KeyID: otherKeyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started.Add(3 * time.Minute), DurationMS: 900, Usage: domain.Usage{InputTokens: 1000, OutputTokens: 1000, TotalTokens: 2000}} for _, event := range []domain.UsageEvent{previous, missing, otherTenantEvent} { @@ -129,6 +129,34 @@ func TestUsageFiltersAndDailyAggregationPostgres(t *testing.T) { t.Fatal(err) } } + requestDetail, err := store.ListUsage(ctx, UsageQuery{TenantID: tenantID, RequestID: events[0].RequestID, Limit: 1}) + if err != nil { + t.Fatal(err) + } + if len(requestDetail.Data) != 1 || requestDetail.Data[0].RequestID != events[0].RequestID { + t.Fatalf("tenant could not retrieve its ledger-linked request: %+v", requestDetail) + } + crossTenantDetail, err := store.ListUsage(ctx, UsageQuery{TenantID: tenantID, RequestID: otherTenantEvent.RequestID, Limit: 1}) + if err != nil { + t.Fatal(err) + } + if len(crossTenantDetail.Data) != 0 { + t.Fatalf("tenant retrieved another tenant's request by request id: %+v", crossTenantDetail) + } + firstPage, err := store.ListUsage(ctx, UsageQuery{TenantID: tenantID, From: started.Add(-time.Hour), To: started.Add(10 * time.Minute), Limit: 1}) + if err != nil { + t.Fatal(err) + } + if len(firstPage.Data) != 1 || firstPage.NextCursor == "" { + t.Fatalf("expected a cursor for the first usage page: %+v", firstPage) + } + secondPage, err := store.ListUsage(ctx, UsageQuery{TenantID: tenantID, From: started.Add(-time.Hour), To: started.Add(10 * time.Minute), Limit: 1, Cursor: firstPage.NextCursor}) + if err != nil { + t.Fatal(err) + } + if len(secondPage.Data) != 1 || secondPage.Data[0].RequestID == firstPage.Data[0].RequestID { + t.Fatalf("cursor did not advance usage page: first=%+v second=%+v", firstPage, secondPage) + } if _, err := db.Exec(ctx, `UPDATE usage_events SET charged_micros=25,cost_micros=25 WHERE request_id=$1`, previous.RequestID); err != nil { t.Fatal(err) } @@ -136,9 +164,12 @@ func TestUsageFiltersAndDailyAggregationPostgres(t *testing.T) { if err != nil { t.Fatal(err) } - if len(analytics.Models) != 1 || analytics.Models[0].RequestCount != 3 || analytics.Models[0].SuccessfulRequests != 2 || analytics.Models[0].ProviderCount != 1 || analytics.Models[0].CacheReadInputTokens != 5 || analytics.Models[0].MissingUsageRequests != 1 || analytics.Models[0].ChargedMicros != 125 || analytics.Models[0].PreviousChargedMicros != 25 || analytics.Models[0].ChargeChangePercent == nil || *analytics.Models[0].ChargeChangePercent != 400 { + if len(analytics.Models) != 1 || analytics.Models[0].RequestCount != 3 || analytics.Models[0].SuccessfulRequests != 2 || analytics.Models[0].ProviderCount != 1 || analytics.Models[0].CacheReadInputTokens != 5 || analytics.Models[0].MissingUsageRequests != 1 || analytics.Models[0].ChargedMicros != 125 || analytics.Models[0].P50DurationMS != 200 || analytics.Models[0].P95DurationMS < 200 || analytics.Models[0].P50TTFTMS != 47 || analytics.Models[0].P95TTFTMS != 47 || analytics.Models[0].PreviousChargedMicros != 25 || analytics.Models[0].ChargeChangePercent == nil || *analytics.Models[0].ChargeChangePercent != 400 { t.Fatalf("unexpected model analytics: %+v", analytics.Models) } + if len(analytics.Keys) != 1 || analytics.Keys[0].KeyID != keyID || analytics.Keys[0].RequestCount != 3 || analytics.Keys[0].ChargedMicros != 125 { + t.Fatalf("unexpected key analytics: %+v", analytics.Keys) + } if len(analytics.Providers) != 1 || analytics.Providers[0].ProviderID != providerID || analytics.Providers[0].WireAPI != "responses" || analytics.Providers[0].RequestCount != 3 || analytics.Providers[0].MissingUsageRequests != 1 || analytics.Providers[0].P95DurationMS < 200 { t.Fatalf("unexpected provider analytics: %+v", analytics.Providers) } diff --git a/internal/controlplane/usage_test.go b/internal/controlplane/usage_test.go new file mode 100644 index 0000000..d47ce16 --- /dev/null +++ b/internal/controlplane/usage_test.go @@ -0,0 +1,26 @@ +package controlplane + +import ( + "errors" + "testing" + "time" +) + +func TestUsageCursorRoundTrip(t *testing.T) { + record := UsageRecord{RequestID: "req_cursor_test", StartedAt: time.Date(2026, 8, 6, 2, 3, 4, 567, time.UTC)} + cursor, err := decodeUsageCursor(encodeUsageCursor(record)) + if err != nil { + t.Fatal(err) + } + if cursor.RequestID != record.RequestID || !cursor.StartedAt.Equal(record.StartedAt) { + t.Fatalf("cursor = %+v, want request %s at %v", cursor, record.RequestID, record.StartedAt) + } +} + +func TestUsageCursorRejectsInvalidInput(t *testing.T) { + for _, raw := range []string{"not-base64!", "e30", string(make([]byte, 2049))} { + if _, err := decodeUsageCursor(raw); !errors.Is(err, ErrInvalidUsageCursor) { + t.Fatalf("decodeUsageCursor(%q) error = %v", raw, err) + } + } +} diff --git a/internal/domain/types.go b/internal/domain/types.go index 81da8f2..c700123 100644 --- a/internal/domain/types.go +++ b/internal/domain/types.go @@ -5,9 +5,21 @@ import "time" type Protocol string const ( - ProtocolOpenAI Protocol = "openai" - ProtocolOpenAIResponses Protocol = "openai_responses" - ProtocolAnthropic Protocol = "anthropic" + ProtocolOpenAI Protocol = "openai" + ProtocolOpenAIResponses Protocol = "openai_responses" + ProtocolOpenAIEmbeddings Protocol = "openai_embeddings" + ProtocolAnthropic Protocol = "anthropic" +) + +// MeteringUnit defines the physical quantity a price applies to. Token-priced +// endpoints are implemented today; image and second keep future media pricing +// explicit instead of overloading token counters. +type MeteringUnit string + +const ( + MeteringUnitToken MeteringUnit = "token" + MeteringUnitImage MeteringUnit = "image" + MeteringUnitSecond MeteringUnit = "second" ) type Principal struct { @@ -17,6 +29,9 @@ type Principal struct { Scopes []string AllowedModels map[string]struct{} MonthlySpendMicros int64 + DailySpendMicros int64 + RequestsPerMinute int64 + TokensPerMinute int64 ExpiresAt *time.Time } @@ -126,6 +141,7 @@ type UsageEvent struct { Attempts int `json:"attempts"` StartedAt time.Time `json:"started_at"` DurationMS int64 `json:"duration_ms"` + TTFTMS int64 `json:"ttft_ms,omitempty"` Usage Usage `json:"usage"` UsageReported bool `json:"usage_reported"` } 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 } diff --git a/internal/httpapi/api_test.go b/internal/httpapi/api_test.go index cbb3967..906e136 100644 --- a/internal/httpapi/api_test.go +++ b/internal/httpapi/api_test.go @@ -74,10 +74,14 @@ func TestInferenceBrowserOriginCORS(t *testing.T) { type fakeBillingMeter struct { authorizeErr error + authorized chan billing.Authorization settled chan domain.UsageEvent } -func (m *fakeBillingMeter) Authorize(context.Context, billing.Authorization) error { +func (m *fakeBillingMeter) Authorize(_ context.Context, authorization billing.Authorization) error { + if m.authorized != nil { + m.authorized <- authorization + } return m.authorizeErr } @@ -132,7 +136,7 @@ func TestOpenAIProxyRewritesModelAndEmitsUsage(t *testing.T) { select { case event := <-sink.events: - if event.PublicModel != "public/model" || event.UpstreamModel != "upstream-model" || event.Usage.TotalTokens != 5 || !event.Success { + if event.PublicModel != "public/model" || event.UpstreamModel != "upstream-model" || event.Usage.TotalTokens != 5 || event.TTFTMS < 1 || !event.Success { t.Fatalf("unexpected usage event: %+v", event) } case <-time.After(time.Second): @@ -178,6 +182,53 @@ func TestProxyFailsOverBeforeWritingResponse(t *testing.T) { } } +func TestGatewayLearnsTTFTAndPrefersFasterProvider(t *testing.T) { + var fastCalls atomic.Int64 + fast := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + fastCalls.Add(1) + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) + })) + defer fast.Close() + var slowCalls atomic.Int64 + slow := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + slowCalls.Add(1) + time.Sleep(25 * time.Millisecond) + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"choices":[],"usage":{"prompt_tokens":1,"completion_tokens":1,"total_tokens":2}}`) + })) + defer slow.Close() + + gateway, sink := newTestGateway(t, + []config.ProviderConfig{ + {ID: "fast", Protocol: domain.ProtocolOpenAI, BaseURL: fast.URL + "/v1", APIKey: "one"}, + {ID: "slow", Protocol: domain.ProtocolOpenAI, BaseURL: slow.URL + "/v1", APIKey: "two"}, + }, + []config.RouteConfig{ + {Provider: "fast", UpstreamModel: "model", Weight: 1}, + {Provider: "slow", UpstreamModel: "model", Weight: 1}, + }, + ) + defer gateway.Close() + + for range 20 { + response := postOpenAI(t, gateway.URL, false) + _, _ = io.Copy(io.Discard, response.Body) + _ = response.Body.Close() + if response.StatusCode != http.StatusOK { + t.Fatalf("unexpected gateway status: %d", response.StatusCode) + } + select { + case <-sink.events: + case <-time.After(time.Second): + t.Fatal("usage event was not emitted") + } + } + if fastCalls.Load() != 16 || slowCalls.Load() != 4 { + t.Fatalf("TTFT feedback was not applied with bounded exploration: fast=%d slow=%d", fastCalls.Load(), slowCalls.Load()) + } +} + func TestProviderSelectorPinsRouteAndKeepsCanonicalUsageModel(t *testing.T) { var primaryCalls atomic.Int64 primary := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { @@ -465,11 +516,106 @@ func TestOpenAIResponsesProxyRewritesModelAndEmitsUsage(t *testing.T) { t.Fatalf("unexpected status %d: %s", response.StatusCode, payload) } event := <-sink.events - if event.Protocol != domain.ProtocolOpenAIResponses || event.UpstreamModel != "gpt-upstream" || event.Usage.TotalTokens != 13 || !event.UsageReported { + if event.Protocol != domain.ProtocolOpenAIResponses || event.UpstreamModel != "gpt-upstream" || event.Usage.TotalTokens != 13 || event.TTFTMS < 1 || !event.UsageReported { t.Fatalf("unexpected Responses usage event: %+v", event) } } +func TestOpenAIEmbeddingsProxyRewritesModelAndMetersInputTokens(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/v1/embeddings" { + t.Errorf("unexpected path: %s", r.URL.Path) + } + if r.Header.Get("Authorization") != "Bearer upstream-secret" { + t.Errorf("unexpected upstream authorization: %q", r.Header.Get("Authorization")) + } + var request map[string]any + if err := json.NewDecoder(r.Body).Decode(&request); err != nil { + t.Error(err) + } + if request["model"] != "embedding-upstream" || request["input"] != "hello vector" { + t.Errorf("unexpected Embeddings request: %+v", request) + } + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"object":"list","data":[{"object":"embedding","index":0,"embedding":[0.1,0.2]}],"model":"embedding-upstream","usage":{"prompt_tokens":8,"total_tokens":8}}`) + })) + defer upstream.Close() + + meter := &fakeBillingMeter{authorized: make(chan billing.Authorization, 1), settled: make(chan domain.UsageEvent, 1)} + gateway, _ := newTestGatewayWithBilling(t, []config.ProviderConfig{{ + ID: "embeddings", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: upstream.URL + "/v1", APIKey: "upstream-secret", + }}, []config.RouteConfig{{Provider: "embeddings", UpstreamModel: "embedding-upstream", Weight: 1}}, meter) + defer gateway.Close() + + request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/embeddings", strings.NewReader(`{"model":"public/model","input":"hello vector"}`)) + request.Header.Set("Authorization", "Bearer client-secret") + response, err := http.DefaultClient.Do(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + if response.StatusCode != http.StatusOK { + payload, _ := io.ReadAll(response.Body) + t.Fatalf("unexpected status %d: %s", response.StatusCode, payload) + } + authorization := <-meter.authorized + if authorization.Protocol != domain.ProtocolOpenAIEmbeddings { + t.Fatalf("authorization protocol = %q", authorization.Protocol) + } + event := <-meter.settled + if event.Protocol != domain.ProtocolOpenAIEmbeddings || event.UpstreamModel != "embedding-upstream" || + event.Usage.InputTokens != 8 || event.Usage.OutputTokens != 0 || event.Usage.TotalTokens != 8 || !event.UsageReported { + t.Fatalf("unexpected Embeddings usage event: %+v", event) + } +} + +func TestOpenAIEmbeddingsRequiresDeclaredCapability(t *testing.T) { + providerConfig := config.ProviderConfig{ID: "embeddings", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: "https://example.invalid/v1", APIKey: "secret"} + modelCatalog := catalog.New(config.Config{Providers: []config.ProviderConfig{providerConfig}, Models: []config.ModelConfig{{ + ID: "public/model", Routes: []config.RouteConfig{{Provider: "embeddings", UpstreamModel: "embedding-upstream", Weight: 1}}, + }}}) + authenticator, err := auth.NewStatic(`[{"key":"client-secret","key_id":"key-1","tenant_id":"tenant-1","project_id":"project-1","scopes":["inference"]}]`, false) + if err != nil { + t.Fatal(err) + } + api := New(Options{Authenticator: authenticator, Catalog: modelCatalog, Router: routing.New(modelCatalog), Metrics: &telemetry.Metrics{}, Logger: slog.New(slog.NewTextHandler(io.Discard, nil)), MaxBodyBytes: 1 << 20}) + request := httptest.NewRequest(http.MethodPost, "/v1/embeddings", strings.NewReader(`{"model":"public/model","input":"hello"}`)) + request.Header.Set("Authorization", "Bearer client-secret") + response := httptest.NewRecorder() + api.Handler().ServeHTTP(response, request) + if response.Code != http.StatusBadRequest || !strings.Contains(response.Body.String(), "unsupported_capability") { + t.Fatalf("unexpected response %d: %s", response.Code, response.Body.String()) + } +} + +func TestOpenAIEmbeddingsRejectsNonJSONSuccessBeforeResponseStarts(t *testing.T) { + upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/html; charset=utf-8") + _, _ = io.WriteString(w, "provider console") + })) + defer upstream.Close() + meter := &fakeBillingMeter{settled: make(chan domain.UsageEvent, 1)} + gateway, _ := newTestGatewayWithBilling(t, []config.ProviderConfig{{ + ID: "embeddings", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: upstream.URL, APIKey: "secret", + }}, []config.RouteConfig{{Provider: "embeddings", UpstreamModel: "embedding-model", Weight: 1}}, meter) + defer gateway.Close() + request, _ := http.NewRequest(http.MethodPost, gateway.URL+"/v1/embeddings", strings.NewReader(`{"model":"public/model","input":"hello"}`)) + request.Header.Set("Authorization", "Bearer client-secret") + response, err := http.DefaultClient.Do(request) + if err != nil { + t.Fatal(err) + } + defer response.Body.Close() + body, _ := io.ReadAll(response.Body) + if response.StatusCode != http.StatusBadGateway || !strings.Contains(string(body), "invalid_provider_response") || strings.Contains(string(body), "provider console") { + t.Fatalf("unexpected response %d: %s", response.StatusCode, body) + } + event := <-meter.settled + if event.Success || event.StatusCode != http.StatusBadGateway || event.ErrorType != "invalid_provider_response" || event.UsageReported { + t.Fatalf("unexpected invalid provider usage event: %+v", event) + } +} + func TestInsufficientBalanceRejectsBeforeCallingUpstream(t *testing.T) { var calls atomic.Int64 upstream := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { @@ -542,9 +688,16 @@ func newTestGatewayWithBilling(t *testing.T, providers []config.ProviderConfig, if err != nil { t.Fatal(err) } + modelConfig := config.ModelConfig{ID: "public/model", OwnedBy: "test", Routes: routes} + for _, providerConfig := range providers { + if providerConfig.WireAPI == "embeddings" { + modelConfig.Capabilities = []string{"embeddings"} + break + } + } cfg := config.Config{ Providers: providers, - Models: []config.ModelConfig{{ID: "public/model", OwnedBy: "test", Routes: routes}}, + Models: []config.ModelConfig{modelConfig}, UpstreamHTTP: config.UpstreamHTTPConfig{ MaxIdleConnections: 100, MaxIdleConnectionsPerHost: 20, IdleConnectionTimeoutSecs: 10, ResponseHeaderTimeoutSecs: 2, diff --git a/internal/limits/limits.go b/internal/limits/limits.go index 6e0bf74..255e5e0 100644 --- a/internal/limits/limits.go +++ b/internal/limits/limits.go @@ -64,20 +64,37 @@ const ( ) const acquireScript = ` -local req = tonumber(ARGV[1]) -local tok = tonumber(ARGV[2]) -local conc = tonumber(ARGV[3]) -local estimate = tonumber(ARGV[4]) -local ttl = tonumber(ARGV[5]) -local r = 0 -local t = 0 +local project_req = tonumber(ARGV[1]) +local project_tok = tonumber(ARGV[2]) +local project_conc = tonumber(ARGV[3]) +local key_req = tonumber(ARGV[4]) +local key_tok = tonumber(ARGV[5]) +local estimate = tonumber(ARGV[6]) +local ttl = tonumber(ARGV[7]) +local pr = 0 +local pt = 0 local c = 0 -if req > 0 then r = redis.call('INCR', KEYS[1]); if r == 1 then redis.call('PEXPIRE', KEYS[1], ttl) end end -if tok > 0 then t = redis.call('INCRBY', KEYS[2], estimate); if t == estimate then redis.call('PEXPIRE', KEYS[2], ttl) end end -if conc > 0 then c = redis.call('INCR', KEYS[3]); redis.call('PEXPIRE', KEYS[3], 3600000) end -if (req > 0 and r > req) then if req > 0 then redis.call('DECR', KEYS[1]) end; if tok > 0 then redis.call('DECRBY', KEYS[2], estimate) end; if conc > 0 then redis.call('DECR', KEYS[3]) end; return {0,1} end -if (tok > 0 and t > tok) then if req > 0 then redis.call('DECR', KEYS[1]) end; if tok > 0 then redis.call('DECRBY', KEYS[2], estimate) end; if conc > 0 then redis.call('DECR', KEYS[3]) end; return {0,2} end -if (conc > 0 and c > conc) then if req > 0 then redis.call('DECR', KEYS[1]) end; if tok > 0 then redis.call('DECRBY', KEYS[2], estimate) end; if conc > 0 then redis.call('DECR', KEYS[3]) end; return {0,3} end +local kr = 0 +local kt = 0 +if project_req > 0 then pr = redis.call('INCR', KEYS[1]); if pr == 1 then redis.call('PEXPIRE', KEYS[1], ttl) end end +if project_tok > 0 then pt = redis.call('INCRBY', KEYS[2], estimate); if pt == estimate then redis.call('PEXPIRE', KEYS[2], ttl) end end +if project_conc > 0 then c = redis.call('INCR', KEYS[3]); redis.call('PEXPIRE', KEYS[3], 3600000) end +if key_req > 0 then kr = redis.call('INCR', KEYS[4]); if kr == 1 then redis.call('PEXPIRE', KEYS[4], ttl) end end +if key_tok > 0 then kt = redis.call('INCRBY', KEYS[5], estimate); if kt == estimate then redis.call('PEXPIRE', KEYS[5], ttl) end end +local reason = 0 +if project_req > 0 and pr > project_req then reason = 1 +elseif project_tok > 0 and pt > project_tok then reason = 2 +elseif project_conc > 0 and c > project_conc then reason = 3 +elseif key_req > 0 and kr > key_req then reason = 4 +elseif key_tok > 0 and kt > key_tok then reason = 5 end +if reason > 0 then + if project_req > 0 then redis.call('DECR', KEYS[1]) end + if project_tok > 0 then redis.call('DECRBY', KEYS[2], estimate) end + if project_conc > 0 then redis.call('DECR', KEYS[3]) end + if key_req > 0 then redis.call('DECR', KEYS[4]) end + if key_tok > 0 then redis.call('DECRBY', KEYS[5], estimate) end + return {0,reason} +end return {1,0} ` @@ -139,15 +156,15 @@ func (l *Limiter) Policy(projectID string) (domain.LimitPolicy, bool) { } func (l *Limiter) Acquire(ctx context.Context, principal domain.Principal, body []byte) (Lease, error) { - policy, ok := l.Policy(principal.ProjectID) - if !ok || (policy.RequestsPerMinute == 0 && policy.TokensPerMinute == 0 && policy.Concurrent == 0) { + policy, _ := l.Policy(principal.ProjectID) + if policy.RequestsPerMinute == 0 && policy.TokensPerMinute == 0 && policy.Concurrent == 0 && principal.RequestsPerMinute == 0 && principal.TokensPerMinute == 0 { return noopLease{}, nil } estimate := EstimateTokens(body, l.defaultMaxOutput) minute := time.Now().Unix() / 60 if l.redis != nil && time.Now().UnixNano() >= l.redisRetryAt.Load() { redisContext, cancel := context.WithTimeout(ctx, redisCommandTimeout) - lease, err := l.acquireRedis(redisContext, principal.ProjectID, minute, policy, estimate) + lease, err := l.acquireRedis(redisContext, principal, minute, policy, estimate) cancel() if err == nil { return lease, nil @@ -158,7 +175,7 @@ func (l *Limiter) Acquire(ctx context.Context, principal domain.Principal, body } l.markRedis(false, err) } - return l.acquireLocal(principal.ProjectID, minute, policy, estimate) + return l.acquireLocal(principal, minute, policy, estimate) } func EstimateTokens(body []byte, defaultMaxOutput int64) int64 { @@ -192,10 +209,13 @@ func EstimateTokens(body []byte, defaultMaxOutput int64) int64 { return input + maxOutput } -func (l *Limiter) acquireRedis(ctx context.Context, project string, minute int64, policy domain.LimitPolicy, estimate int64) (Lease, error) { - base := l.prefix + ":" + project + ":" + strconv.FormatInt(minute, 10) - keys := []string{base + ":requests", base + ":tokens", l.prefix + ":" + project + ":concurrent"} - values, err := l.redis.Eval(ctx, acquireScript, keys, policy.RequestsPerMinute, policy.TokensPerMinute, policy.Concurrent, estimate, 125000).Result() +func (l *Limiter) acquireRedis(ctx context.Context, principal domain.Principal, minute int64, policy domain.LimitPolicy, estimate int64) (Lease, error) { + projectBase := l.prefix + ":project:" + principal.ProjectID + ":" + strconv.FormatInt(minute, 10) + keyBase := l.prefix + ":key:" + principal.KeyID + ":" + strconv.FormatInt(minute, 10) + concurrencyKey := l.prefix + ":project:" + principal.ProjectID + ":concurrent" + keys := []string{projectBase + ":requests", projectBase + ":tokens", concurrencyKey, keyBase + ":requests", keyBase + ":tokens"} + values, err := l.redis.Eval(ctx, acquireScript, keys, policy.RequestsPerMinute, policy.TokensPerMinute, policy.Concurrent, + principal.RequestsPerMinute, principal.TokensPerMinute, estimate, 125000).Result() if err != nil { return nil, err } @@ -207,19 +227,19 @@ func (l *Limiter) acquireRedis(ctx context.Context, project string, minute int64 reason, _ := toInt64(items[1]) if allowed == 0 { switch reason { - case 1: + case 1, 4: return nil, ErrRequestsExceeded - case 2: + case 2, 5: return nil, ErrTokensExceeded default: return nil, ErrConcurrencyLimit } } l.markRedis(true, nil) - return &redisLease{limiter: l, key: keys[2]}, nil + return &redisLease{limiter: l, key: concurrencyKey}, nil } -func (l *Limiter) acquireLocal(project string, minute int64, policy domain.LimitPolicy, estimate int64) (Lease, error) { +func (l *Limiter) acquireLocal(principal domain.Principal, minute int64, policy domain.LimitPolicy, estimate int64) (Lease, error) { l.localMu.Lock() defer l.localMu.Unlock() for key, value := range l.local { @@ -227,28 +247,44 @@ func (l *Limiter) acquireLocal(project string, minute int64, policy domain.Limit delete(l.local, key) } } - window := l.local[project] + projectKey := "project:" + principal.ProjectID + keyKey := "key:" + principal.KeyID + projectWindow := l.localWindow(projectKey, minute) + keyWindow := l.localWindow(keyKey, minute) + if policy.RequestsPerMinute > 0 && projectWindow.requests >= policy.RequestsPerMinute { + return nil, ErrRequestsExceeded + } + if policy.TokensPerMinute > 0 && (estimate > policy.TokensPerMinute || projectWindow.tokens > policy.TokensPerMinute-estimate) { + return nil, ErrTokensExceeded + } + if policy.Concurrent > 0 && projectWindow.concurrent >= policy.Concurrent { + return nil, ErrConcurrencyLimit + } + if principal.RequestsPerMinute > 0 && keyWindow.requests >= principal.RequestsPerMinute { + return nil, ErrRequestsExceeded + } + if principal.TokensPerMinute > 0 && (estimate > principal.TokensPerMinute || keyWindow.tokens > principal.TokensPerMinute-estimate) { + return nil, ErrTokensExceeded + } + projectWindow.requests++ + projectWindow.tokens += estimate + projectWindow.concurrent++ + keyWindow.requests++ + keyWindow.tokens += estimate + return &localLease{limiter: l, project: projectKey}, nil +} + +func (l *Limiter) localWindow(key string, minute int64) *localWindow { + window := l.local[key] if window == nil { window = &localWindow{minute: minute} - l.local[project] = window + l.local[key] = window } else if window.minute != minute { window.minute = minute window.requests = 0 window.tokens = 0 } - if policy.RequestsPerMinute > 0 && window.requests >= policy.RequestsPerMinute { - return nil, ErrRequestsExceeded - } - if policy.TokensPerMinute > 0 && (estimate > policy.TokensPerMinute || window.tokens > policy.TokensPerMinute-estimate) { - return nil, ErrTokensExceeded - } - if policy.Concurrent > 0 && window.concurrent >= policy.Concurrent { - return nil, ErrConcurrencyLimit - } - window.requests++ - window.tokens += estimate - window.concurrent++ - return &localLease{limiter: l, project: project}, nil + return window } func (l *Limiter) releaseLocal(project string) { diff --git a/internal/limits/limits_test.go b/internal/limits/limits_test.go index 78b346e..8e48d42 100644 --- a/internal/limits/limits_test.go +++ b/internal/limits/limits_test.go @@ -56,6 +56,44 @@ func TestLocalTokenAndConcurrencyLimits(t *testing.T) { retry.Release() } +func TestLocalKeyLimitDoesNotConsumeProjectQuotaWhenRejected(t *testing.T) { + limiter := New("", "test", 0, nil) + limiter.ReplacePolicies([]domain.LimitPolicy{{ProjectID: "project-1", RequestsPerMinute: 2}}) + firstKey := domain.Principal{ProjectID: "project-1", KeyID: "key-1", RequestsPerMinute: 1} + lease, err := limiter.Acquire(context.Background(), firstKey, []byte(`{}`)) + if err != nil { + t.Fatal(err) + } + lease.Release() + if _, err := limiter.Acquire(context.Background(), firstKey, []byte(`{}`)); !errors.Is(err, ErrRequestsExceeded) { + t.Fatalf("second key request error = %v, want request limit", err) + } + secondKey := domain.Principal{ProjectID: "project-1", KeyID: "key-2", RequestsPerMinute: 1} + lease, err = limiter.Acquire(context.Background(), secondKey, []byte(`{}`)) + if err != nil { + t.Fatalf("key rejection consumed project quota: %v", err) + } + lease.Release() + if _, err := limiter.Acquire(context.Background(), domain.Principal{ProjectID: "project-1", KeyID: "key-3"}, []byte(`{}`)); !errors.Is(err, ErrRequestsExceeded) { + t.Fatalf("project request limit error = %v", err) + } +} + +func TestLocalKeyTokenLimit(t *testing.T) { + limiter := New("", "test", 0, nil) + body := []byte(`{"max_tokens":4}`) + estimate := EstimateTokens(body, 0) + principal := domain.Principal{ProjectID: "project-1", KeyID: "key-1", TokensPerMinute: estimate} + lease, err := limiter.Acquire(context.Background(), principal, body) + if err != nil { + t.Fatal(err) + } + lease.Release() + if _, err := limiter.Acquire(context.Background(), principal, body); !errors.Is(err, ErrTokensExceeded) { + t.Fatalf("second key token request error = %v, want token limit", err) + } +} + func TestEstimateTokensUsesLargestExplicitOutputLimit(t *testing.T) { body := []byte(`{"max_tokens":10,"max_completion_tokens":25}`) want := int64((len(body)+3)/4 + 25) diff --git a/internal/provider/forwarder.go b/internal/provider/forwarder.go index d7d850d..f816f8a 100644 --- a/internal/provider/forwarder.go +++ b/internal/provider/forwarder.go @@ -19,9 +19,10 @@ import ( ) type Result struct { - Response *http.Response - Route domain.Route - Attempts int + Response *http.Response + Route domain.Route + Attempts int + AttemptStartedAt time.Time } type Forwarder struct { @@ -83,7 +84,7 @@ func (f *Forwarder) Forward(ctx context.Context, protocol domain.Protocol, reque lastErr = fmt.Errorf("upstream %s returned %d", route.Provider.ID, response.StatusCode) continue } - return Result{Response: response, Route: route, Attempts: attempts}, nil + return Result{Response: response, Route: route, Attempts: attempts, AttemptStartedAt: attemptStarted}, nil } if lastErr == nil { lastErr = errors.New("all upstream routes failed") @@ -99,6 +100,19 @@ func (f *Forwarder) observe(modelID string, route domain.Route, statusCode int, providerhealth.Observation{StatusCode: statusCode, Latency: latency, Failed: failed}) } +// ObserveTTFT feeds the first user-visible output latency into adaptive route +// selection. It is deliberately separate from the header/circuit observation +// because streaming TTFT is only known after response forwarding begins. +func (f *Forwarder) ObserveTTFT(modelID string, route domain.Route, latency time.Duration) { + if f.health == nil || latency <= 0 { + return + } + f.health.ObserveTTFT(providerhealth.RouteKey{ModelID: modelID, ProviderID: route.Provider.ID, WireAPI: route.Provider.EffectiveWireAPI()}, latency) + if f.metrics != nil { + f.metrics.UpstreamTTFT(latency) + } +} + func rewriteModel(body []byte, upstreamModel string) ([]byte, error) { return rewriteRequest(body, upstreamModel, domain.ProtocolOpenAI) } @@ -141,6 +155,8 @@ func rewriteRequestWithWireAPI(body []byte, upstreamModel string, protocol domai func endpointURL(provider domain.Provider, _ domain.Protocol) string { baseURL := strings.TrimRight(provider.BaseURL, "/") switch provider.EffectiveWireAPI() { + case "embeddings": + return baseURL + "/embeddings" case "responses": return baseURL + "/responses" case "messages": diff --git a/internal/provider/forwarder_test.go b/internal/provider/forwarder_test.go index e9751ce..2ae81af 100644 --- a/internal/provider/forwarder_test.go +++ b/internal/provider/forwarder_test.go @@ -58,3 +58,21 @@ func TestResponsesWireAPIUsesResponsesEndpointWithoutChatStreamOptions(t *testin t.Fatalf("endpoint URL = %q", got) } } + +func TestEmbeddingsWireAPIUsesEmbeddingsEndpoint(t *testing.T) { + result, err := rewriteRequestWithWireAPI([]byte(`{"model":"public/model","input":["one","two"]}`), "embedding-upstream", domain.ProtocolOpenAIEmbeddings, "embeddings") + if err != nil { + t.Fatal(err) + } + var body map[string]json.RawMessage + if err := json.Unmarshal(result, &body); err != nil { + t.Fatal(err) + } + if string(body["model"]) != `"embedding-upstream"` || string(body["input"]) != `["one","two"]` { + t.Fatalf("unexpected rewritten body: %s", result) + } + provider := domain.Provider{BaseURL: "https://example.test/v1", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings"} + if got := endpointURL(provider, domain.ProtocolOpenAIEmbeddings); got != "https://example.test/v1/embeddings" { + t.Fatalf("endpoint URL = %q", got) + } +} diff --git a/internal/providerhealth/prober.go b/internal/providerhealth/prober.go new file mode 100644 index 0000000..820fd04 --- /dev/null +++ b/internal/providerhealth/prober.go @@ -0,0 +1,171 @@ +package providerhealth + +import ( + "context" + "encoding/json" + "io" + "log/slog" + "net/http" + "strings" + "time" + + "aigw/internal/catalog" + "aigw/internal/domain" +) + +type ProbeMetrics interface { + ProviderProbe(success bool) +} + +type ProbeOptions struct { + Enabled bool + Interval time.Duration + Timeout time.Duration + Catalog *catalog.Catalog + Tracker *Tracker + Metrics ProbeMetrics + Logger *slog.Logger + Client *http.Client +} + +type Prober struct { + enabled bool + interval time.Duration + timeout time.Duration + catalog *catalog.Catalog + tracker *Tracker + metrics ProbeMetrics + logger *slog.Logger + client *http.Client +} + +type probeTarget struct { + provider domain.Provider + keys []RouteKey +} + +func NewProber(options ProbeOptions) *Prober { + if options.Interval <= 0 { + options.Interval = 30 * time.Second + } + if options.Timeout <= 0 { + options.Timeout = 5 * time.Second + } + if options.Logger == nil { + options.Logger = slog.Default() + } + if options.Client == nil { + transport := http.DefaultTransport.(*http.Transport).Clone() + transport.Proxy = http.ProxyFromEnvironment + options.Client = &http.Client{Transport: transport, Timeout: options.Timeout} + } + return &Prober{enabled: options.Enabled, interval: options.Interval, timeout: options.Timeout, + catalog: options.Catalog, tracker: options.Tracker, metrics: options.Metrics, logger: options.Logger, client: options.Client} +} + +func (p *Prober) Run(ctx context.Context) { + if p == nil || !p.enabled || p.catalog == nil || p.tracker == nil { + return + } + p.ProbeOnce(ctx) + ticker := time.NewTicker(p.interval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + p.ProbeOnce(ctx) + } + } +} + +func (p *Prober) ProbeOnce(ctx context.Context) { + for _, target := range p.targets() { + if ctx.Err() != nil { + return + } + p.probe(ctx, target) + } +} + +func (p *Prober) targets() []probeTarget { + providerIndexes := make(map[string]int) + seen := make(map[RouteKey]struct{}) + result := make([]probeTarget, 0) + for _, model := range p.catalog.AllModels() { + for _, route := range model.Routes { + key := RouteKey{ModelID: model.ID, ProviderID: route.Provider.ID, WireAPI: route.Provider.EffectiveWireAPI()} + if key.ProviderID == "" { + continue + } + if _, exists := seen[key]; exists { + continue + } + seen[key] = struct{}{} + index, exists := providerIndexes[route.Provider.ID] + if !exists { + index = len(result) + providerIndexes[route.Provider.ID] = index + result = append(result, probeTarget{provider: route.Provider}) + } + result[index].keys = append(result[index].keys, key) + } + } + return result +} + +func (p *Prober) probe(parent context.Context, target probeTarget) { + ctx, cancel := context.WithTimeout(parent, p.timeout) + defer cancel() + request, err := http.NewRequestWithContext(ctx, http.MethodGet, providerModelsURL(target.provider), nil) + if err != nil { + p.observe(target, 0, 0, false) + return + } + request.Header.Set("Accept", "application/json") + request.Header.Set("User-Agent", "aigw-health/0.1") + if target.provider.Protocol == domain.ProtocolAnthropic { + request.Header.Set("x-api-key", target.provider.APIKey) + request.Header.Set("anthropic-version", "2023-06-01") + } else { + request.Header.Set("Authorization", "Bearer "+target.provider.APIKey) + } + started := time.Now() + response, err := p.client.Do(request) + latency := time.Since(started) + if err != nil { + p.observe(target, 0, latency, false) + p.logger.Warn("provider_probe_failed", "provider_id", target.provider.ID, "error", err) + return + } + body, readErr := io.ReadAll(io.LimitReader(response.Body, 1<<20)) + _ = response.Body.Close() + success := response.StatusCode >= 200 && response.StatusCode < 300 && readErr == nil && + strings.HasPrefix(strings.ToLower(response.Header.Get("Content-Type")), "application/json") && validModelsEnvelope(body) + p.observe(target, response.StatusCode, latency, success) + if !success { + p.logger.Warn("provider_probe_failed", "provider_id", target.provider.ID, "status_code", response.StatusCode) + } +} + +func (p *Prober) observe(target probeTarget, statusCode int, latency time.Duration, success bool) { + for _, key := range target.keys { + p.tracker.Observe(key, Observation{StatusCode: statusCode, Latency: latency, Failed: !success, Active: true}) + } + if p.metrics != nil { + p.metrics.ProviderProbe(success) + } +} + +func validModelsEnvelope(body []byte) bool { + var envelope struct { + Data []json.RawMessage `json:"data"` + } + return json.Unmarshal(body, &envelope) == nil && envelope.Data != nil +} + +func providerModelsURL(provider domain.Provider) string { + baseURL := strings.TrimRight(provider.BaseURL, "/") + return baseURL + "/models" +} diff --git a/internal/providerhealth/prober_test.go b/internal/providerhealth/prober_test.go new file mode 100644 index 0000000..69b6e80 --- /dev/null +++ b/internal/providerhealth/prober_test.go @@ -0,0 +1,91 @@ +package providerhealth + +import ( + "context" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "sync/atomic" + "testing" + "time" + + "aigw/internal/catalog" + "aigw/internal/domain" +) + +type probeMetricCounter struct { + total atomic.Int64 + failed atomic.Int64 +} + +func (m *probeMetricCounter) ProviderProbe(success bool) { + m.total.Add(1) + if !success { + m.failed.Add(1) + } +} + +func TestProberAuthenticatesDeduplicatesAndOpensRouteCircuits(t *testing.T) { + var calls atomic.Int64 + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + calls.Add(1) + if r.URL.Path != "/v1/models" || r.Header.Get("Authorization") != "Bearer secret" { + t.Errorf("unexpected probe path=%s auth=%q", r.URL.Path, r.Header.Get("Authorization")) + } + http.Error(w, "unavailable", http.StatusServiceUnavailable) + })) + defer server.Close() + + provider := domain.Provider{ID: "provider", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: server.URL + "/v1", APIKey: "secret"} + model := domain.Model{ID: "public/model", Routes: []domain.Route{{Provider: provider, UpstreamModel: "one"}, {Provider: provider, UpstreamModel: "two"}}} + tracker := New(Options{FailureThreshold: 1, OpenDuration: time.Minute}) + metrics := &probeMetricCounter{} + prober := NewProber(ProbeOptions{Enabled: true, Catalog: catalog.NewModels([]domain.Model{model}), Tracker: tracker, + Metrics: metrics, Timeout: time.Second, Logger: slog.New(slog.NewTextHandler(io.Discard, nil))}) + prober.ProbeOnce(context.Background()) + + key := RouteKey{ModelID: model.ID, ProviderID: provider.ID, WireAPI: provider.WireAPI} + if calls.Load() != 1 || metrics.total.Load() != 1 || metrics.failed.Load() != 1 || !tracker.CircuitOpen(key) { + t.Fatalf("calls=%d total=%d failed=%d open=%v", calls.Load(), metrics.total.Load(), metrics.failed.Load(), tracker.CircuitOpen(key)) + } + status := tracker.Snapshot()[0] + if status.ActiveProbes != 1 || status.LastProbeAt == nil || status.LastStatusCode != http.StatusServiceUnavailable { + t.Fatalf("unexpected probe status: %+v", status) + } +} + +func TestProberRejectsHTMLSuccessResponse(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.Header().Set("Content-Type", "text/html") + _, _ = io.WriteString(w, "provider console") + })) + defer server.Close() + provider := domain.Provider{ID: "provider", Protocol: domain.ProtocolOpenAI, BaseURL: server.URL, APIKey: "secret"} + model := domain.Model{ID: "model", Routes: []domain.Route{{Provider: provider, UpstreamModel: "model"}}} + tracker := New(Options{FailureThreshold: 1}) + NewProber(ProbeOptions{Enabled: true, Catalog: catalog.NewModels([]domain.Model{model}), Tracker: tracker, Timeout: time.Second, + Logger: slog.New(slog.NewTextHandler(io.Discard, nil))}).ProbeOnce(context.Background()) + if !tracker.CircuitOpen(RouteKey{ModelID: model.ID, ProviderID: provider.ID, WireAPI: "chat_completions"}) { + t.Fatal("HTML success response must not be treated as a healthy API probe") + } +} + +func TestProberUsesAnthropicAuthentication(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Header.Get("x-api-key") != "anthropic-secret" || r.Header.Get("anthropic-version") == "" { + t.Errorf("unexpected Anthropic headers") + } + w.Header().Set("Content-Type", "application/json") + _, _ = io.WriteString(w, `{"data":[]}`) + })) + defer server.Close() + provider := domain.Provider{ID: "anthropic", Protocol: domain.ProtocolAnthropic, WireAPI: "messages", BaseURL: server.URL + "/v1", APIKey: "anthropic-secret"} + model := domain.Model{ID: "anthropic/model", Routes: []domain.Route{{Provider: provider, UpstreamModel: "model"}}} + tracker := New(Options{}) + NewProber(ProbeOptions{Enabled: true, Catalog: catalog.NewModels([]domain.Model{model}), Tracker: tracker, Timeout: time.Second}).ProbeOnce(context.Background()) + status := tracker.Snapshot()[0] + if status.State != "healthy" || status.ActiveProbes != 1 { + t.Fatalf("unexpected status: %+v", status) + } +} diff --git a/internal/providerhealth/redis_history.go b/internal/providerhealth/redis_history.go new file mode 100644 index 0000000..322acce --- /dev/null +++ b/internal/providerhealth/redis_history.go @@ -0,0 +1,347 @@ +package providerhealth + +import ( + "context" + "encoding/json" + "errors" + "log/slog" + "strings" + "sync" + "sync/atomic" + "time" + + "github.com/redis/go-redis/v9" +) + +const ( + defaultHistoryStream = "aigw:provider-health:events" + defaultHistoryQueueSize = 4096 + defaultHistoryMaxEvents = 20_000 + defaultHistoryTTL = 15 * time.Minute + historyPublishBatchSize = 64 + historyCommandTimeout = 500 * time.Millisecond + historyRetryCooldown = 5 * time.Second +) + +type HistoryOptions struct { + Enabled bool + RedisURL string + Stream string + Instance string + QueueSize int + MaxEvents int64 + TTL time.Duration + Logger *slog.Logger + Client *redis.Client + Metrics SharedHistoryMetrics +} + +type SharedHistoryMetrics interface { + ProviderHealthSharedPublished() + ProviderHealthSharedImported() + ProviderHealthSharedDropped() + ProviderHealthSharedRedisFailure() + ProviderHealthSharedConnected(bool) +} + +type RedisHistory struct { + redis *redis.Client + ownedClient bool + stream string + instance string + queue chan Event + maxEvents int64 + ttl time.Duration + logger *slog.Logger + metrics SharedHistoryMetrics + retryAt atomic.Int64 + dropped atomic.Uint64 + closed atomic.Bool + failureSeen atomic.Bool +} + +func NewRedisHistory(options HistoryOptions) *RedisHistory { + if !options.Enabled || (strings.TrimSpace(options.RedisURL) == "" && options.Client == nil) { + return nil + } + if options.Logger == nil { + options.Logger = slog.Default() + } + if options.Stream == "" { + options.Stream = defaultHistoryStream + } + if options.QueueSize <= 0 { + options.QueueSize = defaultHistoryQueueSize + } + if options.MaxEvents <= 0 { + options.MaxEvents = defaultHistoryMaxEvents + } + if options.TTL <= 0 { + options.TTL = defaultHistoryTTL + } + client := options.Client + ownedClient := false + if client == nil { + redisOptions, err := redis.ParseURL(options.RedisURL) + if err != nil { + options.Logger.Warn("provider_health_redis_config_invalid", "error", err, "fallback", "local") + return nil + } + redisOptions.MaxRetries = -1 + redisOptions.DialerRetries = 1 + redisOptions.DialTimeout = historyCommandTimeout + redisOptions.ReadTimeout = historyCommandTimeout + redisOptions.WriteTimeout = historyCommandTimeout + redisOptions.PoolTimeout = historyCommandTimeout + client = redis.NewClient(redisOptions) + ownedClient = true + } + return &RedisHistory{redis: client, ownedClient: ownedClient, stream: options.Stream, instance: options.Instance, + queue: make(chan Event, options.QueueSize), maxEvents: options.MaxEvents, ttl: options.TTL, + logger: options.Logger, metrics: options.Metrics} +} + +func (h *RedisHistory) Enqueue(event Event) { + if h == nil || h.closed.Load() { + return + } + select { + case h.queue <- event: + default: + if h.recordDropped(1) == 1 { + h.logger.Warn("provider_health_share_queue_full", "fallback", "local") + } + } +} + +func (h *RedisHistory) Run(ctx context.Context, tracker *Tracker) { + if h == nil || tracker == nil { + return + } + var workers sync.WaitGroup + workers.Add(2) + go func() { + defer workers.Done() + h.runPublisher(ctx) + }() + go func() { + defer workers.Done() + h.runReader(ctx, tracker) + }() + workers.Wait() +} + +func (h *RedisHistory) runPublisher(ctx context.Context) { + batch := make([]Event, 0, historyPublishBatchSize) + for { + select { + case <-ctx.Done(): + return + case event := <-h.queue: + batch = append(batch[:0], event) + } + for len(batch) < historyPublishBatchSize { + select { + case event := <-h.queue: + batch = append(batch, event) + default: + h.publishBatch(ctx, batch) + batch = batch[:0] + goto nextBatch + } + } + h.publishBatch(ctx, batch) + batch = batch[:0] + nextBatch: + } +} + +func (h *RedisHistory) runReader(ctx context.Context, tracker *Tracker) { + lastID := h.importRecent(ctx, tracker) + if lastID == "" { + lastID = "0-0" + } + for ctx.Err() == nil { + if time.Now().UnixNano() < h.retryAt.Load() { + if !waitHistoryRetry(ctx, 100*time.Millisecond) { + return + } + continue + } + lastID = h.read(ctx, tracker, lastID) + } +} + +func (h *RedisHistory) Close() error { + if h == nil || !h.closed.CompareAndSwap(false, true) || !h.ownedClient { + return nil + } + return h.redis.Close() +} + +func (h *RedisHistory) Dropped() uint64 { + if h == nil { + return 0 + } + return h.dropped.Load() +} + +func (h *RedisHistory) publishBatch(parent context.Context, events []Event) { + if len(events) == 0 { + return + } + if time.Now().UnixNano() < h.retryAt.Load() { + h.recordDropped(uint64(len(events))) + return + } + ctx, cancel := context.WithTimeout(parent, historyCommandTimeout) + defer cancel() + pipe := h.redis.Pipeline() + published := 0 + for _, event := range events { + payload, err := json.Marshal(event) + if err != nil { + h.recordDropped(1) + continue + } + pipe.XAdd(ctx, &redis.XAddArgs{Stream: h.stream, MaxLen: h.maxEvents, Approx: true, + Values: map[string]any{"instance": h.instance, "event": string(payload)}}) + published++ + } + if published == 0 { + pipe.Discard() + return + } + pipe.Expire(ctx, h.stream, h.ttl) + if _, err := pipe.Exec(ctx); err != nil { + h.recordDropped(uint64(published)) + h.fail(err) + return + } + h.recovered() + for range published { + if h.metrics != nil { + h.metrics.ProviderHealthSharedPublished() + } + } +} + +func (h *RedisHistory) importRecent(parent context.Context, tracker *Tracker) string { + ctx, cancel := context.WithTimeout(parent, historyCommandTimeout) + defer cancel() + messages, err := h.redis.XRevRangeN(ctx, h.stream, "+", "-", h.maxEvents).Result() + if err != nil { + if !errors.Is(err, redis.Nil) { + h.fail(err) + } + return "" + } + if len(messages) == 0 { + h.recovered() + return "" + } + lastID := messages[0].ID + for index := len(messages) - 1; index >= 0; index-- { + h.applyMessage(tracker, messages[index]) + } + h.recovered() + return lastID +} + +func (h *RedisHistory) read(parent context.Context, tracker *Tracker, lastID string) string { + if time.Now().UnixNano() < h.retryAt.Load() { + return lastID + } + ctx, cancel := context.WithTimeout(parent, historyCommandTimeout) + defer cancel() + streams, err := h.redis.XRead(ctx, &redis.XReadArgs{Streams: []string{h.stream, lastID}, Count: 512, Block: 200 * time.Millisecond}).Result() + if errors.Is(err, redis.Nil) { + h.recovered() + return lastID + } + if err != nil { + h.fail(err) + return lastID + } + h.recovered() + for _, stream := range streams { + for _, message := range stream.Messages { + h.applyMessage(tracker, message) + lastID = message.ID + } + } + return lastID +} + +func (h *RedisHistory) applyMessage(tracker *Tracker, message redis.XMessage) { + if asString(message.Values["instance"]) == h.instance && h.instance != "" { + return + } + payload := asString(message.Values["event"]) + var event Event + if payload == "" || json.Unmarshal([]byte(payload), &event) != nil { + return + } + if event.ObservedAt.Before(time.Now().Add(-h.ttl)) { + return + } + tracker.ApplyShared(event) + if h.metrics != nil { + h.metrics.ProviderHealthSharedImported() + } +} + +func (h *RedisHistory) fail(err error) { + h.retryAt.Store(time.Now().Add(historyRetryCooldown).UnixNano()) + if h.metrics != nil { + h.metrics.ProviderHealthSharedConnected(false) + } + if h.failureSeen.CompareAndSwap(false, true) { + if h.metrics != nil { + h.metrics.ProviderHealthSharedRedisFailure() + } + h.logger.Warn("provider_health_redis_unavailable", "error", err, "fallback", "local", "retry_after", historyRetryCooldown) + } +} + +func (h *RedisHistory) recovered() { + h.retryAt.Store(0) + if h.metrics != nil { + h.metrics.ProviderHealthSharedConnected(true) + } + if h.failureSeen.Swap(false) { + h.logger.Info("provider_health_redis_recovered") + } +} + +func (h *RedisHistory) recordDropped(count uint64) uint64 { + total := h.dropped.Add(count) + if h.metrics != nil { + for range count { + h.metrics.ProviderHealthSharedDropped() + } + } + return total +} + +func waitHistoryRetry(ctx context.Context, delay time.Duration) bool { + timer := time.NewTimer(delay) + defer timer.Stop() + select { + case <-ctx.Done(): + return false + case <-timer.C: + return true + } +} + +func asString(value any) string { + switch item := value.(type) { + case string: + return item + case []byte: + return string(item) + default: + return "" + } +} diff --git a/internal/providerhealth/tracker.go b/internal/providerhealth/tracker.go index a5af2b2..6394def 100644 --- a/internal/providerhealth/tracker.go +++ b/internal/providerhealth/tracker.go @@ -8,31 +8,59 @@ import ( const recentWindow = 100 +type EventKind string + +const ( + EventOutcome EventKind = "outcome" + EventTTFT EventKind = "ttft" +) + type RouteKey struct { - ModelID string - ProviderID string - WireAPI string + ModelID string `json:"model_id"` + ProviderID string `json:"provider_id"` + WireAPI string `json:"wire_api"` } type Observation struct { StatusCode int Latency time.Duration Failed bool + Active bool ObservedAt time.Time } +type Event struct { + Kind EventKind `json:"kind"` + Key RouteKey `json:"key"` + StatusCode int `json:"status_code,omitempty"` + LatencyMillis int64 `json:"latency_ms,omitempty"` + Failed bool `json:"failed,omitempty"` + Active bool `json:"active,omitempty"` + ObservedAt time.Time `json:"observed_at"` +} + +type EventSink interface { + Enqueue(Event) +} + type Status struct { ModelID string `json:"model_id"` ProviderID string `json:"provider_id"` WireAPI string `json:"wire_api"` State string `json:"state"` Attempts uint64 `json:"attempts"` + ActiveProbes uint64 `json:"active_probes"` RecentSamples int `json:"recent_samples"` AvailabilityPercent float64 `json:"availability_percent"` HeaderLatencyEWMA int64 `json:"header_latency_ewma_ms"` + TTFTSamples uint64 `json:"ttft_samples"` + TTFTEWMA int64 `json:"ttft_ewma_ms"` + SharedAttempts uint64 `json:"shared_attempts"` + SharedTTFTSamples uint64 `json:"shared_ttft_samples"` ConsecutiveFailures uint64 `json:"consecutive_failures"` LastStatusCode int `json:"last_status_code,omitempty"` LastObservedAt *time.Time `json:"last_observed_at,omitempty"` + LastProbeAt *time.Time `json:"last_probe_at,omitempty"` LastHealthyAt *time.Time `json:"last_healthy_at,omitempty"` CircuitOpenUntil *time.Time `json:"circuit_open_until,omitempty"` } @@ -41,6 +69,7 @@ type Options struct { FailureThreshold uint64 OpenDuration time.Duration Now func() time.Time + Sink EventSink } type Tracker struct { @@ -48,23 +77,63 @@ type Tracker struct { failureThreshold uint64 openDuration time.Duration now func() time.Time + sinkMu sync.RWMutex + sink EventSink } type routeState struct { mu sync.RWMutex attempts uint64 + activeProbes uint64 consecutiveFailures uint64 lastStatusCode int lastObservedAt time.Time + lastProbeAt time.Time lastHealthyAt time.Time openUntil time.Time headerLatencyEWMA float64 + ttftSamples uint64 + ttftEWMA float64 + sharedAttempts uint64 + sharedTTFTSamples uint64 recent [recentWindow]bool recentCount int recentPosition int recentHealthy int } +// ObserveTTFT records user-visible response latency without counting a second +// request outcome. Forwarder observations already update availability and the +// circuit when response headers arrive. +func (t *Tracker) ObserveTTFT(key RouteKey, latency time.Duration) { + if t == nil || key.ProviderID == "" || latency <= 0 { + return + } + observedAt := t.now() + t.observeTTFT(key, latency, false) + t.publish(Event{Kind: EventTTFT, Key: key, LatencyMillis: durationMillis(latency), ObservedAt: observedAt}) +} + +func (t *Tracker) observeTTFT(key RouteKey, latency time.Duration, shared bool) { + value, _ := t.states.LoadOrStore(key, &routeState{}) + state := value.(*routeState) + state.mu.Lock() + defer state.mu.Unlock() + valueMS := float64(durationMillis(latency)) + if valueMS < 1 { + valueMS = 1 + } + state.ttftSamples++ + if shared { + state.sharedTTFTSamples++ + } + if state.ttftEWMA == 0 { + state.ttftEWMA = valueMS + } else { + state.ttftEWMA = state.ttftEWMA*0.8 + valueMS*0.2 + } +} + func New(options Options) *Tracker { if options.FailureThreshold == 0 { options.FailureThreshold = 3 @@ -75,7 +144,7 @@ func New(options Options) *Tracker { if options.Now == nil { options.Now = time.Now } - return &Tracker{failureThreshold: options.FailureThreshold, openDuration: options.OpenDuration, now: options.Now} + return &Tracker{failureThreshold: options.FailureThreshold, openDuration: options.OpenDuration, now: options.Now, sink: options.Sink} } func (t *Tracker) Observe(key RouteKey, observation Observation) { @@ -85,12 +154,26 @@ func (t *Tracker) Observe(key RouteKey, observation Observation) { if observation.ObservedAt.IsZero() { observation.ObservedAt = t.now() } + t.observe(key, observation, false) + t.publish(Event{Kind: EventOutcome, Key: key, StatusCode: observation.StatusCode, + LatencyMillis: durationMillis(observation.Latency), Failed: observation.Failed, + Active: observation.Active, ObservedAt: observation.ObservedAt}) +} + +func (t *Tracker) observe(key RouteKey, observation Observation, shared bool) { value, _ := t.states.LoadOrStore(key, &routeState{}) state := value.(*routeState) state.mu.Lock() defer state.mu.Unlock() state.attempts++ + if shared { + state.sharedAttempts++ + } + if observation.Active { + state.activeProbes++ + state.lastProbeAt = observation.ObservedAt + } state.lastStatusCode = observation.StatusCode state.lastObservedAt = observation.ObservedAt if observation.Latency > 0 { @@ -117,6 +200,53 @@ func (t *Tracker) Observe(key RouteKey, observation Observation) { state.lastHealthyAt = observation.ObservedAt } +func (t *Tracker) SetSink(sink EventSink) { + if t == nil { + return + } + t.sinkMu.Lock() + t.sink = sink + t.sinkMu.Unlock() +} + +func (t *Tracker) ApplyShared(event Event) { + if t == nil || event.Key.ProviderID == "" { + return + } + if event.ObservedAt.IsZero() { + event.ObservedAt = t.now() + } + switch event.Kind { + case EventOutcome: + t.observe(event.Key, Observation{StatusCode: event.StatusCode, Latency: time.Duration(event.LatencyMillis) * time.Millisecond, + Failed: event.Failed, Active: event.Active, ObservedAt: event.ObservedAt}, true) + case EventTTFT: + if event.LatencyMillis > 0 { + t.observeTTFT(event.Key, time.Duration(event.LatencyMillis)*time.Millisecond, true) + } + } +} + +func (t *Tracker) publish(event Event) { + t.sinkMu.RLock() + sink := t.sink + t.sinkMu.RUnlock() + if sink != nil { + sink.Enqueue(event) + } +} + +func durationMillis(value time.Duration) int64 { + if value <= 0 { + return 0 + } + milliseconds := value.Milliseconds() + if milliseconds < 1 { + return 1 + } + return milliseconds +} + func (s *routeState) addRecent(healthy bool) { if s.recentCount == recentWindow { if s.recent[s.recentPosition] { @@ -146,6 +276,21 @@ func (t *Tracker) CircuitOpen(key RouteKey) bool { return state.openUntil.After(t.now()) } +// StatusFor returns one immutable route-health snapshot for routing decisions. +func (t *Tracker) StatusFor(key RouteKey) (Status, bool) { + if t == nil { + return Status{}, false + } + value, ok := t.states.Load(key) + if !ok { + return Status{}, false + } + state := value.(*routeState) + state.mu.RLock() + defer state.mu.RUnlock() + return statusFromState(key, state, t.now()), true +} + func (t *Tracker) Snapshot() []Status { if t == nil { return []Status{} @@ -172,8 +317,9 @@ func (t *Tracker) Snapshot() []Status { func statusFromState(key RouteKey, state *routeState, now time.Time) Status { item := Status{ModelID: key.ModelID, ProviderID: key.ProviderID, WireAPI: key.WireAPI, - Attempts: state.attempts, RecentSamples: state.recentCount, HeaderLatencyEWMA: int64(state.headerLatencyEWMA + 0.5), - ConsecutiveFailures: state.consecutiveFailures, LastStatusCode: state.lastStatusCode} + Attempts: state.attempts, ActiveProbes: state.activeProbes, RecentSamples: state.recentCount, HeaderLatencyEWMA: int64(state.headerLatencyEWMA + 0.5), + TTFTSamples: state.ttftSamples, TTFTEWMA: int64(state.ttftEWMA + 0.5), SharedAttempts: state.sharedAttempts, + SharedTTFTSamples: state.sharedTTFTSamples, ConsecutiveFailures: state.consecutiveFailures, LastStatusCode: state.lastStatusCode} if state.recentCount > 0 { item.AvailabilityPercent = float64(state.recentHealthy) / float64(state.recentCount) * 100 } @@ -181,6 +327,10 @@ func statusFromState(key RouteKey, state *routeState, now time.Time) Status { value := state.lastObservedAt item.LastObservedAt = &value } + if !state.lastProbeAt.IsZero() { + value := state.lastProbeAt + item.LastProbeAt = &value + } if !state.lastHealthyAt.IsZero() { value := state.lastHealthyAt item.LastHealthyAt = &value diff --git a/internal/providerhealth/tracker_test.go b/internal/providerhealth/tracker_test.go index 73bd6d5..19752df 100644 --- a/internal/providerhealth/tracker_test.go +++ b/internal/providerhealth/tracker_test.go @@ -1,11 +1,23 @@ package providerhealth import ( + "context" + "fmt" + "log/slog" + "os" "sync" "testing" "time" + + "github.com/redis/go-redis/v9" ) +type captureSink struct { + events []Event +} + +func (s *captureSink) Enqueue(event Event) { s.events = append(s.events, event) } + func TestTrackerOpensAndRecoversCircuit(t *testing.T) { now := time.Date(2026, time.August, 6, 0, 0, 0, 0, time.UTC) tracker := New(Options{FailureThreshold: 3, OpenDuration: 30 * time.Second, Now: func() time.Time { return now }}) @@ -31,6 +43,146 @@ func TestTrackerOpensAndRecoversCircuit(t *testing.T) { } } +func TestTrackerRecordsActiveProbeMetadata(t *testing.T) { + now := time.Date(2026, time.August, 6, 1, 0, 0, 0, time.UTC) + tracker := New(Options{Now: func() time.Time { return now }}) + key := RouteKey{ModelID: "model", ProviderID: "provider", WireAPI: "embeddings"} + tracker.Observe(key, Observation{StatusCode: 200, Latency: 20 * time.Millisecond, Active: true}) + status := tracker.Snapshot()[0] + if status.ActiveProbes != 1 || status.LastProbeAt == nil || !status.LastProbeAt.Equal(now) || status.State != "healthy" { + t.Fatalf("unexpected active probe status: %+v", status) + } +} + +func TestTrackerRecordsTTFTWithoutDoubleCountingAvailability(t *testing.T) { + tracker := New(Options{}) + key := RouteKey{ModelID: "model", ProviderID: "provider", WireAPI: "responses"} + tracker.Observe(key, Observation{StatusCode: 200, Latency: 20 * time.Millisecond}) + tracker.ObserveTTFT(key, 100*time.Millisecond) + tracker.ObserveTTFT(key, 200*time.Millisecond) + status, found := tracker.StatusFor(key) + if !found || status.Attempts != 1 || status.RecentSamples != 1 || status.TTFTSamples != 2 || status.TTFTEWMA != 120 { + t.Fatalf("unexpected TTFT status: %+v", status) + } +} + +func TestTrackerSharesLocalEventsWithoutRepublishingImportedEvents(t *testing.T) { + sink := &captureSink{} + tracker := New(Options{Sink: sink}) + key := RouteKey{ModelID: "model", ProviderID: "provider", WireAPI: "responses"} + tracker.Observe(key, Observation{StatusCode: 200, Latency: 20 * time.Millisecond}) + tracker.ObserveTTFT(key, 100*time.Millisecond) + if len(sink.events) != 2 || sink.events[0].Kind != EventOutcome || sink.events[1].Kind != EventTTFT { + t.Fatalf("unexpected published events %+v", sink.events) + } + tracker.ApplyShared(Event{Kind: EventOutcome, Key: key, StatusCode: 503, Failed: true, ObservedAt: time.Now()}) + tracker.ApplyShared(Event{Kind: EventTTFT, Key: key, LatencyMillis: 250, ObservedAt: time.Now()}) + if len(sink.events) != 2 { + t.Fatalf("imported observations were republished: %d events", len(sink.events)) + } + status, found := tracker.StatusFor(key) + if !found || status.Attempts != 2 || status.SharedAttempts != 1 || status.TTFTSamples != 2 || status.SharedTTFTSamples != 1 { + t.Fatalf("unexpected shared status %+v", status) + } +} + +func TestSharedFailuresOpenLocalCircuit(t *testing.T) { + now := time.Date(2026, time.August, 6, 2, 0, 0, 0, time.UTC) + tracker := New(Options{FailureThreshold: 3, OpenDuration: 30 * time.Second, Now: func() time.Time { return now }}) + key := RouteKey{ModelID: "model", ProviderID: "provider", WireAPI: "responses"} + for range 3 { + tracker.ApplyShared(Event{Kind: EventOutcome, Key: key, StatusCode: 503, Failed: true, ObservedAt: now}) + } + status, found := tracker.StatusFor(key) + if !found || !tracker.CircuitOpen(key) || status.State != "open" || status.Attempts != 3 || status.SharedAttempts != 3 { + t.Fatalf("shared failures did not open the circuit: %+v", status) + } +} + +func TestRedisHistoryDropsWithoutBlockingWhenQueueIsFull(t *testing.T) { + client := redis.NewClient(&redis.Options{Addr: "127.0.0.1:1"}) + defer client.Close() + history := NewRedisHistory(HistoryOptions{Enabled: true, Client: client, QueueSize: 1, + Logger: slog.New(slog.NewTextHandler(ioDiscard{}, nil))}) + event := Event{Kind: EventOutcome, Key: RouteKey{ModelID: "model", ProviderID: "provider"}, ObservedAt: time.Now()} + history.Enqueue(event) + history.Enqueue(event) + if history.Dropped() != 1 { + t.Fatalf("dropped = %d, want 1", history.Dropped()) + } +} + +func TestRedisHistorySharesAndReplaysObservations(t *testing.T) { + redisURL := os.Getenv("AIGW_TEST_REDIS_URL") + if redisURL == "" { + t.Skip("AIGW_TEST_REDIS_URL is not set") + } + stream := fmt.Sprintf("aigw:test:provider-health:%d", time.Now().UnixNano()) + options, err := redis.ParseURL(redisURL) + if err != nil { + t.Fatal(err) + } + cleanupClient := redis.NewClient(options) + t.Cleanup(func() { + _, _ = cleanupClient.Del(context.Background(), stream).Result() + _ = cleanupClient.Close() + }) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + key := RouteKey{ModelID: "model", ProviderID: "provider", WireAPI: "responses"} + first := New(Options{}) + second := New(Options{}) + firstHistory := NewRedisHistory(HistoryOptions{Enabled: true, RedisURL: redisURL, Stream: stream, Instance: "first", + TTL: time.Minute, MaxEvents: 100, Logger: slog.New(slog.NewTextHandler(ioDiscard{}, nil))}) + secondHistory := NewRedisHistory(HistoryOptions{Enabled: true, RedisURL: redisURL, Stream: stream, Instance: "second", + TTL: time.Minute, MaxEvents: 100, Logger: slog.New(slog.NewTextHandler(ioDiscard{}, nil))}) + if firstHistory == nil || secondHistory == nil { + t.Fatal("Redis histories were not configured") + } + t.Cleanup(func() { _ = firstHistory.Close() }) + t.Cleanup(func() { _ = secondHistory.Close() }) + first.SetSink(firstHistory) + go firstHistory.Run(ctx, first) + go secondHistory.Run(ctx, second) + + first.Observe(key, Observation{StatusCode: 200, Latency: 10 * time.Millisecond}) + first.ObserveTTFT(key, 40*time.Millisecond) + waitForSharedStatus(t, second, key, 1, 1) + status, _ := first.StatusFor(key) + if status.SharedAttempts != 0 || status.SharedTTFTSamples != 0 { + t.Fatalf("publisher imported its own events: %+v", status) + } + + restarted := New(Options{}) + restartedHistory := NewRedisHistory(HistoryOptions{Enabled: true, RedisURL: redisURL, Stream: stream, Instance: "restarted", + TTL: time.Minute, MaxEvents: 100, Logger: slog.New(slog.NewTextHandler(ioDiscard{}, nil))}) + if restartedHistory == nil { + t.Fatal("restart history was not configured") + } + t.Cleanup(func() { _ = restartedHistory.Close() }) + go restartedHistory.Run(ctx, restarted) + waitForSharedStatus(t, restarted, key, 1, 1) +} + +func waitForSharedStatus(t *testing.T, tracker *Tracker, key RouteKey, attempts, ttft uint64) { + t.Helper() + deadline := time.Now().Add(5 * time.Second) + for time.Now().Before(deadline) { + status, found := tracker.StatusFor(key) + if found && status.SharedAttempts >= attempts && status.SharedTTFTSamples >= ttft { + return + } + time.Sleep(25 * time.Millisecond) + } + status, _ := tracker.StatusFor(key) + t.Fatalf("shared status did not converge: %+v", status) +} + +type ioDiscard struct{} + +func (ioDiscard) Write(data []byte) (int, error) { return len(data), nil } + func TestTrackerConcurrentObservations(t *testing.T) { tracker := New(Options{FailureThreshold: 1000}) key := RouteKey{ModelID: "model", ProviderID: "provider", WireAPI: "chat_completions"} diff --git a/internal/routing/router.go b/internal/routing/router.go index 37de633..4465c09 100644 --- a/internal/routing/router.go +++ b/internal/routing/router.go @@ -24,6 +24,13 @@ type Router struct { counters sync.Map } +const ( + adaptiveMinimumAvailabilitySamples = 5 + adaptiveMinimumTTFTSamples = 3 + adaptiveExplorationInterval = 20 + adaptivePreferenceThreshold = 0.90 +) + func New(catalog *catalog.Catalog, trackers ...*providerhealth.Tracker) *Router { router := &Router{catalog: catalog} if len(trackers) > 0 { @@ -90,6 +97,8 @@ func protocolCompatible(provider domain.Provider, requestProtocol domain.Protoco return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "chat_completions" case domain.ProtocolOpenAIResponses: return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "responses" + case domain.ProtocolOpenAIEmbeddings: + return provider.Protocol == domain.ProtocolOpenAI && provider.EffectiveWireAPI() == "embeddings" case domain.ProtocolAnthropic: return provider.Protocol == domain.ProtocolAnthropic && provider.EffectiveWireAPI() == "messages" default: @@ -104,25 +113,113 @@ func (r *Router) rotate(modelID string, protocol domain.Protocol, routes []domai key := modelID + "\x00" + string(protocol) + "\x00" + strconv.Itoa(routes[0].Priority) counterValue, _ := r.counters.LoadOrStore(key, &atomic.Uint64{}) counter := counterValue.(*atomic.Uint64).Add(1) - 1 - - totalWeight := 0 - for _, route := range routes { - totalWeight += route.Weight + indices := make([]int, len(routes)) + for index := range routes { + indices[index] = index } - position := int(counter % uint64(totalWeight)) - selected := 0 - for i, route := range routes { - if position < route.Weight { - selected = i - break + primaryPool := indices + if r.health != nil && counter%adaptiveExplorationInterval != adaptiveExplorationInterval-1 { + if preferred := r.preferredRoutes(modelID, routes); len(preferred) > 0 && len(preferred) < len(routes) { + indices = append(preferred, difference(indices, preferred)...) + primaryPool = preferred } - position -= route.Weight } + selectedPosition := weightedPosition(routes, primaryPool, counter) + selected := primaryPool[selectedPosition] result := make([]domain.Route, 0, len(routes)) result = append(result, routes[selected]) - for offset := 1; offset < len(routes); offset++ { - result = append(result, routes[(selected+offset)%len(routes)]) + for _, index := range indices { + if index != selected { + result = append(result, routes[index]) + } } return result } + +func (r *Router) preferredRoutes(modelID string, routes []domain.Route) []int { + type measuredRoute struct { + qualified bool + status providerhealth.Status + } + measured := make([]measuredRoute, len(routes)) + fastestTTFT := int64(0) + for index, route := range routes { + status, exists := r.health.StatusFor(providerhealth.RouteKey{ModelID: modelID, ProviderID: route.Provider.ID, WireAPI: route.Provider.EffectiveWireAPI()}) + if !exists { + continue + } + qualified := status.RecentSamples >= adaptiveMinimumAvailabilitySamples || status.TTFTSamples >= adaptiveMinimumTTFTSamples + measured[index] = measuredRoute{qualified: qualified, status: status} + if status.TTFTSamples >= adaptiveMinimumTTFTSamples && status.TTFTEWMA > 0 && (fastestTTFT == 0 || status.TTFTEWMA < fastestTTFT) { + fastestTTFT = status.TTFTEWMA + } + } + + scores := make([]float64, len(routes)) + best := 0.0 + hasQualified := false + for index, item := range measured { + score := 1.0 + if item.qualified { + hasQualified = true + if item.status.RecentSamples >= adaptiveMinimumAvailabilitySamples { + availability := item.status.AvailabilityPercent / 100 + if availability < 0.05 { + availability = 0.05 + } + score *= availability + } + if fastestTTFT > 0 && item.status.TTFTSamples >= adaptiveMinimumTTFTSamples && item.status.TTFTEWMA > 0 { + latencyFactor := float64(fastestTTFT) / float64(item.status.TTFTEWMA) + if latencyFactor < 0.10 { + latencyFactor = 0.10 + } + score *= latencyFactor + } + } + scores[index] = score + if score > best { + best = score + } + } + if !hasQualified { + return nil + } + result := make([]int, 0, len(routes)) + for index, score := range scores { + if score >= best*adaptivePreferenceThreshold { + result = append(result, index) + } + } + return result +} + +func difference(all, selected []int) []int { + included := make(map[int]struct{}, len(selected)) + for _, index := range selected { + included[index] = struct{}{} + } + result := make([]int, 0, len(all)-len(selected)) + for _, index := range all { + if _, exists := included[index]; !exists { + result = append(result, index) + } + } + return result +} + +func weightedPosition(routes []domain.Route, indices []int, counter uint64) int { + totalWeight := 0 + for _, index := range indices { + totalWeight += routes[index].Weight + } + position := int(counter % uint64(totalWeight)) + for positionIndex, routeIndex := range indices { + if position < routes[routeIndex].Weight { + return positionIndex + } + position -= routes[routeIndex].Weight + } + return 0 +} diff --git a/internal/routing/router_test.go b/internal/routing/router_test.go index 2ad4685..f421184 100644 --- a/internal/routing/router_test.go +++ b/internal/routing/router_test.go @@ -96,15 +96,56 @@ func TestPlanUsesWeightsForPrimarySelection(t *testing.T) { } } +func TestPlanPrefersLowerTTFTAndStillExplores(t *testing.T) { + health := providerhealth.New(providerhealth.Options{}) + cfg := config.Config{ + Providers: []config.ProviderConfig{ + {ID: "fast", Protocol: domain.ProtocolOpenAI, BaseURL: "https://fast.test", APIKey: "one"}, + {ID: "slow", Protocol: domain.ProtocolOpenAI, BaseURL: "https://slow.test", APIKey: "two"}, + }, + Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{ + {Provider: "fast", UpstreamModel: "model", Weight: 1}, + {Provider: "slow", UpstreamModel: "model", Weight: 1}, + }}}, + } + for _, providerID := range []string{"fast", "slow"} { + key := providerhealth.RouteKey{ModelID: "public/model", ProviderID: providerID, WireAPI: "chat_completions"} + for range 5 { + health.Observe(key, providerhealth.Observation{StatusCode: 200, Latency: 20 * time.Millisecond}) + } + latency := 100 * time.Millisecond + if providerID == "slow" { + latency = 500 * time.Millisecond + } + for range 3 { + health.ObserveTTFT(key, latency) + } + } + router := New(catalog.New(cfg), health) + counts := map[string]int{} + for range 200 { + plan, err := router.Plan("public/model", domain.ProtocolOpenAI) + if err != nil { + t.Fatal(err) + } + counts[plan[0].Provider.ID]++ + } + if counts["fast"] != 190 || counts["slow"] != 10 { + t.Fatalf("adaptive selection must prefer fast route while preserving 5%% exploration: %+v", counts) + } +} + func TestPlanSeparatesOpenAIWireAPIs(t *testing.T) { cfg := config.Config{ Providers: []config.ProviderConfig{ {ID: "chat", Protocol: domain.ProtocolOpenAI, WireAPI: "chat_completions", BaseURL: "https://chat.test", APIKey: "one"}, {ID: "responses", Protocol: domain.ProtocolOpenAI, WireAPI: "responses", BaseURL: "https://responses.test", APIKey: "two"}, + {ID: "embeddings", Protocol: domain.ProtocolOpenAI, WireAPI: "embeddings", BaseURL: "https://embeddings.test", APIKey: "three"}, }, Models: []config.ModelConfig{{ID: "public/model", Routes: []config.RouteConfig{ {Provider: "chat", UpstreamModel: "chat-model", Weight: 1}, {Provider: "responses", UpstreamModel: "responses-model", Weight: 1}, + {Provider: "embeddings", UpstreamModel: "embedding-model", Weight: 1}, }}}, } router := New(catalog.New(cfg)) @@ -116,6 +157,10 @@ func TestPlanSeparatesOpenAIWireAPIs(t *testing.T) { if err != nil || len(responses) != 1 || responses[0].Provider.ID != "responses" { t.Fatalf("unexpected Responses plan: %+v err=%v", responses, err) } + embeddings, err := router.Plan("public/model", domain.ProtocolOpenAIEmbeddings) + if err != nil || len(embeddings) != 1 || embeddings[0].Provider.ID != "embeddings" { + t.Fatalf("unexpected Embeddings plan: %+v err=%v", embeddings, err) + } } func TestPlanProviderPinsWithoutFallbackToOtherProviders(t *testing.T) { diff --git a/internal/telemetry/metrics.go b/internal/telemetry/metrics.go index 04cc494..0631de4 100644 --- a/internal/telemetry/metrics.go +++ b/internal/telemetry/metrics.go @@ -4,24 +4,34 @@ import ( "fmt" "net/http" "sync/atomic" + "time" "aigw/internal/billing" ) type Metrics struct { - requests atomic.Uint64 - failed atomic.Uint64 - inFlight atomic.Int64 - attempts atomic.Uint64 - droppedUsage atomic.Uint64 - settlementBacklog atomic.Int64 - settlementSpool atomic.Int64 - stripeRefundBacklog atomic.Int64 - stripeUncollected atomic.Int64 - stripeMismatches atomic.Int64 - stripeWebhooks atomic.Int64 - unmeteredSuccesses atomic.Int64 - ready atomic.Int64 + requests atomic.Uint64 + failed atomic.Uint64 + inFlight atomic.Int64 + attempts atomic.Uint64 + upstreamTTFTCount atomic.Uint64 + upstreamTTFTMSSum atomic.Uint64 + providerProbes atomic.Uint64 + providerProbeFailed atomic.Uint64 + providerSharedPublished atomic.Uint64 + providerSharedImported atomic.Uint64 + providerSharedDropped atomic.Uint64 + providerSharedFailures atomic.Uint64 + providerSharedConnected atomic.Int64 + droppedUsage atomic.Uint64 + settlementBacklog atomic.Int64 + settlementSpool atomic.Int64 + stripeRefundBacklog atomic.Int64 + stripeUncollected atomic.Int64 + stripeMismatches atomic.Int64 + stripeWebhooks atomic.Int64 + unmeteredSuccesses atomic.Int64 + ready atomic.Int64 } func (m *Metrics) SetSettlementQueue(backlog int64, spool int) { @@ -60,6 +70,39 @@ func (m *Metrics) UpstreamAttempt() { m.attempts.Add(1) } +func (m *Metrics) UpstreamTTFT(latency time.Duration) { + if latency <= 0 { + return + } + milliseconds := latency.Milliseconds() + if milliseconds < 1 { + milliseconds = 1 + } + m.upstreamTTFTCount.Add(1) + m.upstreamTTFTMSSum.Add(uint64(milliseconds)) +} + +func (m *Metrics) ProviderProbe(success bool) { + m.providerProbes.Add(1) + if !success { + m.providerProbeFailed.Add(1) + } +} + +func (m *Metrics) ProviderHealthSharedPublished() { m.providerSharedPublished.Add(1) } +func (m *Metrics) ProviderHealthSharedImported() { m.providerSharedImported.Add(1) } +func (m *Metrics) ProviderHealthSharedDropped() { m.providerSharedDropped.Add(1) } +func (m *Metrics) ProviderHealthSharedRedisFailure() { + m.providerSharedFailures.Add(1) +} +func (m *Metrics) ProviderHealthSharedConnected(connected bool) { + if connected { + m.providerSharedConnected.Store(1) + return + } + m.providerSharedConnected.Store(0) +} + func (m *Metrics) UsageDropped() { m.droppedUsage.Add(1) } @@ -70,6 +113,15 @@ func (m *Metrics) ServeHTTP(w http.ResponseWriter, _ *http.Request) { fmt.Fprintf(w, "# TYPE aigw_requests_failed_total counter\naigw_requests_failed_total %d\n", m.failed.Load()) fmt.Fprintf(w, "# TYPE aigw_requests_in_flight gauge\naigw_requests_in_flight %d\n", m.inFlight.Load()) fmt.Fprintf(w, "# TYPE aigw_upstream_attempts_total counter\naigw_upstream_attempts_total %d\n", m.attempts.Load()) + fmt.Fprintf(w, "# TYPE aigw_upstream_ttft_ms_count counter\naigw_upstream_ttft_ms_count %d\n", m.upstreamTTFTCount.Load()) + fmt.Fprintf(w, "# TYPE aigw_upstream_ttft_ms_sum counter\naigw_upstream_ttft_ms_sum %d\n", m.upstreamTTFTMSSum.Load()) + fmt.Fprintf(w, "# TYPE aigw_provider_probes_total counter\naigw_provider_probes_total %d\n", m.providerProbes.Load()) + fmt.Fprintf(w, "# TYPE aigw_provider_probe_failures_total counter\naigw_provider_probe_failures_total %d\n", m.providerProbeFailed.Load()) + fmt.Fprintf(w, "# TYPE aigw_provider_health_shared_published_total counter\naigw_provider_health_shared_published_total %d\n", m.providerSharedPublished.Load()) + fmt.Fprintf(w, "# TYPE aigw_provider_health_shared_imported_total counter\naigw_provider_health_shared_imported_total %d\n", m.providerSharedImported.Load()) + fmt.Fprintf(w, "# TYPE aigw_provider_health_shared_dropped_total counter\naigw_provider_health_shared_dropped_total %d\n", m.providerSharedDropped.Load()) + fmt.Fprintf(w, "# TYPE aigw_provider_health_shared_redis_failures_total counter\naigw_provider_health_shared_redis_failures_total %d\n", m.providerSharedFailures.Load()) + fmt.Fprintf(w, "# TYPE aigw_provider_health_shared_connected gauge\naigw_provider_health_shared_connected %d\n", m.providerSharedConnected.Load()) fmt.Fprintf(w, "# TYPE aigw_usage_events_dropped_total counter\naigw_usage_events_dropped_total %d\n", m.droppedUsage.Load()) fmt.Fprintf(w, "# TYPE aigw_billing_settlement_backlog gauge\naigw_billing_settlement_backlog %d\n", m.settlementBacklog.Load()) fmt.Fprintf(w, "# TYPE aigw_billing_settlement_spool_records gauge\naigw_billing_settlement_spool_records %d\n", m.settlementSpool.Load()) diff --git a/internal/usage/observer.go b/internal/usage/observer.go index 71c1379..4c7b2b4 100644 --- a/internal/usage/observer.go +++ b/internal/usage/observer.go @@ -4,6 +4,7 @@ import ( "bytes" "encoding/json" "strings" + "time" "aigw/internal/domain" ) @@ -18,21 +19,32 @@ type Observer struct { usage domain.Usage found bool explicitTotal bool + firstOutputAt time.Time + now func() time.Time } func NewObserver(protocol domain.Protocol, stream bool) *Observer { - return &Observer{protocol: protocol, stream: stream} + return &Observer{protocol: protocol, stream: stream, now: time.Now} } func (o *Observer) Write(p []byte) (int, error) { if o.stream { o.observeSSE(p) } else { + if len(p) > 0 { + o.markFirstOutput() + } o.captureTail(p) } return len(p), nil } +// FirstOutputAt is the arrival time of the first user-visible output. For +// streaming responses, metadata-only and heartbeat events are ignored. +func (o *Observer) FirstOutputAt() time.Time { + return o.firstOutputAt +} + func (o *Observer) Usage() domain.Usage { if o.stream { if len(o.line) > 0 { @@ -85,16 +97,88 @@ func (o *Observer) observeSSE(p []byte) { } func (o *Observer) parseSSELine(line []byte) { - if !bytes.HasPrefix(line, []byte("data:")) || !bytes.Contains(line, []byte("\"usage\"")) { + if !bytes.HasPrefix(line, []byte("data:")) { return } payload := bytes.TrimSpace(bytes.TrimPrefix(line, []byte("data:"))) if bytes.Equal(payload, []byte("[DONE]")) { return } + o.observeStreamOutput(payload) + if !bytes.Contains(payload, []byte("\"usage\"")) { + return + } o.parseJSON(payload) } +type streamOutputEvent struct { + Type string `json:"type"` + Delta json.RawMessage `json:"delta"` + Choices []struct { + Delta struct { + Content json.RawMessage `json:"content"` + } `json:"delta"` + } `json:"choices"` +} + +func (o *Observer) observeStreamOutput(payload []byte) { + if !o.firstOutputAt.IsZero() { + return + } + var event streamOutputEvent + if json.Unmarshal(payload, &event) != nil { + return + } + for _, choice := range event.Choices { + if rawContainsVisibleText(choice.Delta.Content) { + o.markFirstOutput() + return + } + } + switch event.Type { + case "response.output_text.delta": + if rawContainsVisibleText(event.Delta) { + o.markFirstOutput() + } + case "content_block_delta": + var delta struct { + Text string `json:"text"` + } + if json.Unmarshal(event.Delta, &delta) == nil && delta.Text != "" { + o.markFirstOutput() + } + } +} + +func rawContainsVisibleText(raw json.RawMessage) bool { + if len(raw) == 0 || bytes.Equal(raw, []byte("null")) { + return false + } + var text string + if json.Unmarshal(raw, &text) == nil { + return text != "" + } + var parts []struct { + Text string `json:"text"` + } + if json.Unmarshal(raw, &parts) != nil { + return false + } + for _, part := range parts { + if part.Text != "" { + return true + } + } + return false +} + +func (o *Observer) markFirstOutput() { + if !o.firstOutputAt.IsZero() { + return + } + o.firstOutputAt = o.now() +} + type tokenDetails struct { CachedTokens *int64 `json:"cached_tokens"` CacheWriteTokens *int64 `json:"cache_write_tokens"` diff --git a/internal/usage/observer_test.go b/internal/usage/observer_test.go index 4a4eba6..e96a38a 100644 --- a/internal/usage/observer_test.go +++ b/internal/usage/observer_test.go @@ -2,6 +2,7 @@ package usage import ( "testing" + "time" "aigw/internal/domain" ) @@ -57,3 +58,73 @@ func TestObserverReadsResponsesUsage(t *testing.T) { t.Fatalf("unexpected streaming Responses usage: %+v reported=%v", got, stream.Reported()) } } + +func TestObserverReadsEmbeddingsUsage(t *testing.T) { + observer := NewObserver(domain.ProtocolOpenAIEmbeddings, false) + _, _ = observer.Write([]byte(`{"object":"list","data":[],"usage":{"prompt_tokens":17,"total_tokens":17}}`)) + got := observer.Usage() + if !observer.Reported() || got.InputTokens != 17 || got.OutputTokens != 0 || got.TotalTokens != 17 { + t.Fatalf("unexpected Embeddings usage: %+v reported=%v", got, observer.Reported()) + } +} + +func TestObserverMarksFirstVisibleStreamingOutput(t *testing.T) { + base := time.Date(2026, 8, 6, 1, 2, 3, 0, time.UTC) + tests := []struct { + name string + protocol domain.Protocol + metadata string + output string + }{ + { + name: "openai chat", protocol: domain.ProtocolOpenAI, + metadata: "data: {\"choices\":[{\"delta\":{\"role\":\"assistant\",\"content\":null}}]}\n\n", + output: "data: {\"choices\":[{\"delta\":{\"content\":\"Hello\"}}]}\n\n", + }, + { + name: "responses", protocol: domain.ProtocolOpenAIResponses, + metadata: "data: {\"type\":\"response.created\"}\n\n", + output: "data: {\"type\":\"response.output_text.delta\",\"delta\":\"Hello\"}\n\n", + }, + { + name: "anthropic", protocol: domain.ProtocolAnthropic, + metadata: "data: {\"type\":\"content_block_start\",\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n", + output: "data: {\"type\":\"content_block_delta\",\"delta\":{\"type\":\"text_delta\",\"text\":\"Hello\"}}\n\n", + }, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + observer := NewObserver(test.protocol, true) + current := base + observer.now = func() time.Time { return current } + _, _ = observer.Write([]byte(test.metadata)) + if got := observer.FirstOutputAt(); !got.IsZero() { + t.Fatalf("metadata marked first output at %v", got) + } + current = base.Add(275 * time.Millisecond) + _, _ = observer.Write([]byte(test.output)) + if got := observer.FirstOutputAt(); !got.Equal(current) { + t.Fatalf("first output = %v, want %v", got, current) + } + current = base.Add(time.Second) + _, _ = observer.Write([]byte(test.output)) + if got := observer.FirstOutputAt(); !got.Equal(base.Add(275 * time.Millisecond)) { + t.Fatalf("first output changed to %v", got) + } + }) + } +} + +func TestObserverMarksFirstNonStreamingBodyWrite(t *testing.T) { + base := time.Date(2026, 8, 6, 1, 2, 3, 0, time.UTC) + observer := NewObserver(domain.ProtocolOpenAI, false) + observer.now = func() time.Time { return base } + _, _ = observer.Write(nil) + if !observer.FirstOutputAt().IsZero() { + t.Fatal("empty write must not mark first output") + } + _, _ = observer.Write([]byte(`{"choices":[]}`)) + if got := observer.FirstOutputAt(); !got.Equal(base) { + t.Fatalf("first output = %v, want %v", got, base) + } +} -- cgit v1.2.3