diff options
Diffstat (limited to 'internal/adminapi')
| -rw-r--r-- | internal/adminapi/api.go | 216 | ||||
| -rw-r--r-- | internal/adminapi/bootstrap_test.go | 57 | ||||
| -rw-r--r-- | internal/adminapi/model_page.go | 276 | ||||
| -rw-r--r-- | internal/adminapi/model_page_test.go | 44 | ||||
| -rw-r--r-- | internal/adminapi/usage_test.go | 26 |
5 files changed, 580 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) } 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(`<!doctype html> +<html lang="en"> +<head> + <meta charset="utf-8"> + <meta name="viewport" content="width=device-width, initial-scale=1"> + <meta name="description" content="{{.Description}}"> + <meta property="og:type" content="website"> + <meta property="og:title" content="{{.Title}}"> + <meta property="og:description" content="{{.Description}}"> + <meta property="og:url" content="{{.CanonicalURL}}"> + <link rel="canonical" href="{{.CanonicalURL}}"> + <title>{{.Title}}</title> + <link rel="icon" href="data:image/svg+xml,%3Csvg xmlns='http://www.w3.org/2000/svg' viewBox='0 0 32 32'%3E%3Crect width='32' height='32' fill='%23102a3a'/%3E%3Ctext x='16' y='22' text-anchor='middle' font-family='Arial' font-weight='700' font-size='18' fill='%23b8eef7'%3EA%3C/text%3E%3C/svg%3E"> + <link rel="stylesheet" href="{{.CSSURL}}"> +</head> +<body> + <header class="catalog-header"> + <a class="catalog-brand" href="{{.CatalogURL}}" aria-label="AIGW model catalog"><span>A</span><strong>AIGW</strong><small>MODEL CATALOG</small></a> + <nav aria-label="Account access"><a class="button secondary" href="{{.SignInURL}}">Sign in</a>{{if .RegistrationEnabled}}<a class="button primary" href="{{.RegistrationURL}}">Create account</a>{{end}}</nav> + </header> + <main class="model-page"> + <a class="back-link" href="{{.CatalogURL}}">Back to models</a> + <header class="model-page-header"> + <div><span class="eyebrow">{{if .Model.OwnedBy}}{{.Model.OwnedBy}}{{else}}INDEPENDENT{{end}}</span><h1>{{.Model.DisplayName}}</h1><code>{{.Model.PublicID}}</code></div> + <span class="health {{.HealthClass}}">{{.HealthLabel}}</span> + </header> + <p class="model-page-description">{{.Description}}</p> + <div class="detail-tags">{{range .Tags}}<span>{{.}}</span>{{end}}</div> + + <div class="model-page-layout"> + <section class="model-specs" aria-labelledby="model-specs-heading"> + <h2 id="model-specs-heading">Model details</h2> + <dl class="detail-grid"> + <div><dt>Input price</dt><dd>{{.InputPrice}} / 1M tokens</dd></div> + <div><dt>Output price</dt><dd>{{.OutputPrice}} / 1M tokens</dd></div> + <div><dt>Cached input</dt><dd>{{.CacheReadPrice}} / 1M tokens</dd></div> + <div><dt>Context window</dt><dd>{{.ContextWindow}} tokens</dd></div> + <div><dt>Max output</dt><dd>{{.MaxOutputTokens}} tokens</dd></div> + <div><dt>Regions</dt><dd>{{if .Model.Regions}}{{range $index,$region := .Model.Regions}}{{if $index}}, {{end}}{{$region}}{{end}}{{else}}Global{{end}}</dd></div> + <div><dt>Released</dt><dd>{{.ReleasedAt}}</dd></div> + <div><dt>Lifecycle</dt><dd>{{.Model.Lifecycle}}</dd></div> + </dl> + </section> + <aside class="model-start"> + <span class="eyebrow">API ACCESS</span> + <h2>Start building</h2> + <code>{{.Model.PublicID}}</code> + <div class="model-start-actions"><a class="button secondary" href="{{.SignInURL}}">Sign in</a>{{if .RegistrationEnabled}}<a class="button primary" href="{{.RegistrationURL}}">Create account</a>{{end}}</div> + </aside> + </div> + + <section class="code-examples" aria-labelledby="examples-heading"> + <div><span class="eyebrow">SUPPORTED APIS</span><h2 id="examples-heading">Code examples</h2></div> + <div class="code-example-grid">{{range .Examples}}<article class="code-example"><header><strong>{{.Name}}</strong><code>{{.Endpoint}}</code></header><pre><code>{{.Code}}</code></pre></article>{{end}}</div> + </section> + </main> +</body> +</html>`)) + +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) + } +} |
