summaryrefslogtreecommitdiff
path: root/internal/adminapi/api.go
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/adminapi/api.go216
1 files changed, 177 insertions, 39 deletions
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)
}