summaryrefslogtreecommitdiff
path: root/internal/controlplane
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-06 09:29:41 +1200
committerChia <Chia@93.nz>2026-08-06 09:32:46 +1200
commit41e322c53d7b4b796eb377d0df9c29ecd10ba431 (patch)
treec730526150e55e39b822d5197e4a20318ecaa449 /internal/controlplane
parenteadb2ffe85c43cf6fc741c9823cd28eedb4a844c (diff)
feat: complete commercial control plane, billing, auth, and model catalog
- add PostgreSQL control-plane persistence with Redis-degraded hot reload - implement prepaid balance, usage ledger, Stripe top-up and reconciliation - add registration, email verification, password reset, invitations and RBAC - support TOTP, Passkey MFA, device sessions, quotas and rate limits - add tenant billing profiles, audit logs and operational readiness checks - build authenticated admin console, Quickstart, Playground and usage analytics - add public model catalog with pricing, filtering and cost estimation - support OpenAI Responses providers and provider health failover - validate real upstream usage reporting and balance settlement
Diffstat (limited to 'internal/controlplane')
-rw-r--r--internal/controlplane/access.go11
-rw-r--r--internal/controlplane/access_test.go4
-rw-r--r--internal/controlplane/mail_operations.go13
-rw-r--r--internal/controlplane/mutations.go82
-rw-r--r--internal/controlplane/preferences.go156
-rw-r--r--internal/controlplane/preferences_test.go31
-rw-r--r--internal/controlplane/queries.go191
-rw-r--r--internal/controlplane/queries_test.go80
-rw-r--r--internal/controlplane/schema.sql110
-rw-r--r--internal/controlplane/snapshot.go35
-rw-r--r--internal/controlplane/store.go4
-rw-r--r--internal/controlplane/store_integration_test.go167
-rw-r--r--internal/controlplane/types.go223
-rw-r--r--internal/controlplane/usage.go128
-rw-r--r--internal/controlplane/usage_analytics.go215
-rw-r--r--internal/controlplane/usage_analytics_test.go34
-rw-r--r--internal/controlplane/usage_integration_test.go152
17 files changed, 1564 insertions, 72 deletions
diff --git a/internal/controlplane/access.go b/internal/controlplane/access.go
index 47e9a8f..1b04db6 100644
--- a/internal/controlplane/access.go
+++ b/internal/controlplane/access.go
@@ -43,23 +43,24 @@ func (a ConsoleActor) Can(permission string) bool {
case RoleTenantAdmin:
switch permission {
case "overview.read", "tenants.read", "projects.read", "projects.write", "keys.read", "keys.write",
- "billing.read", "billing.topup", "usage.read", "audit.read", "limits.read", "limits.write", "users.read", "users.write":
+ "billing.read", "billing.topup", "usage.read", "audit.read", "limits.read", "limits.write", "users.read", "users.write",
+ "preferences.read", "developer.preferences.write", "billing.preferences.write":
return true
}
return false
case RoleTenantBilling:
- return permission == "overview.read" || permission == "billing.read" || permission == "billing.topup" || permission == "usage.read" || permission == "audit.read"
+ return permission == "overview.read" || permission == "billing.read" || permission == "billing.topup" || permission == "usage.read" || permission == "audit.read" || permission == "preferences.read" || permission == "billing.preferences.write"
case RoleTenantDeveloper:
- return permission == "overview.read" || permission == "tenants.read" || permission == "projects.read" || permission == "keys.read" || permission == "keys.write" || permission == "usage.read" || permission == "limits.read"
+ return permission == "overview.read" || permission == "tenants.read" || permission == "projects.read" || permission == "keys.read" || permission == "keys.write" || permission == "usage.read" || permission == "limits.read" || permission == "preferences.read" || permission == "developer.preferences.write"
case RoleTenantViewer:
- return permission == "overview.read" || permission == "tenants.read" || permission == "projects.read" || permission == "keys.read" || permission == "billing.read" || permission == "usage.read" || permission == "limits.read" || permission == "audit.read"
+ return permission == "overview.read" || permission == "tenants.read" || permission == "projects.read" || permission == "keys.read" || permission == "billing.read" || permission == "usage.read" || permission == "limits.read" || permission == "audit.read" || permission == "preferences.read"
default:
return false
}
}
func (a ConsoleActor) Permissions() []string {
- all := []string{"overview.read", "tenants.read", "tenants.write", "projects.read", "projects.write", "keys.read", "keys.write", "platform.read", "platform.write", "billing.read", "billing.topup", "billing.adjust", "usage.read", "limits.read", "limits.write", "users.read", "users.write", "audit.read"}
+ all := []string{"overview.read", "preferences.read", "developer.preferences.write", "billing.preferences.write", "tenants.read", "tenants.write", "projects.read", "projects.write", "keys.read", "keys.write", "platform.read", "platform.write", "billing.read", "billing.topup", "billing.adjust", "usage.read", "limits.read", "limits.write", "users.read", "users.write", "audit.read"}
result := make([]string, 0, len(all))
for _, permission := range all {
if a.Can(permission) {
diff --git a/internal/controlplane/access_test.go b/internal/controlplane/access_test.go
index 7767e0d..d871024 100644
--- a/internal/controlplane/access_test.go
+++ b/internal/controlplane/access_test.go
@@ -13,8 +13,12 @@ func TestConsoleRolePermissions(t *testing.T) {
{RoleTenantAdmin, "keys.write", true},
{RoleTenantAdmin, "limits.write", true},
{RoleTenantBilling, "billing.topup", true},
+ {RoleTenantBilling, "billing.preferences.write", true},
+ {RoleTenantBilling, "developer.preferences.write", false},
{RoleTenantBilling, "keys.read", false},
{RoleTenantDeveloper, "keys.write", true},
+ {RoleTenantDeveloper, "developer.preferences.write", true},
+ {RoleTenantDeveloper, "billing.preferences.write", false},
{RoleTenantDeveloper, "billing.read", false},
{RoleTenantViewer, "usage.read", true},
{RoleTenantViewer, "users.read", false},
diff --git a/internal/controlplane/mail_operations.go b/internal/controlplane/mail_operations.go
index aaaf2c5..d42a328 100644
--- a/internal/controlplane/mail_operations.go
+++ b/internal/controlplane/mail_operations.go
@@ -148,23 +148,26 @@ func (s *Store) queueBillingNotifications(ctx context.Context, config MailNotifi
(COALESCE(sum(-amount_micros) FILTER (WHERE kind='usage' AND created_at>=date_trunc('day',now())-interval '7 days' AND created_at<date_trunc('day',now())),0)/7)::bigint baseline
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(spend.today,0),COALESCE(spend.baseline,0),
+ COALESCE(pref.low_balance_enabled,TRUE),COALESCE(pref.low_balance_threshold_micros,$1)
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))`)
+ AND NOT EXISTS (SELECT 1 FROM mail_suppressions s WHERE lower(s.recipient)=lower(u.email))`, config.LowBalanceMicros)
if err != nil {
return err
}
defer rows.Close()
for rows.Next() {
var tenantID, currency, email, name string
- var available, today, baseline int64
- if err := rows.Scan(&tenantID, &currency, &available, &email, &name, &today, &baseline); err != nil {
+ var available, today, baseline, lowBalanceThreshold int64
+ var lowBalanceEnabled bool
+ if err := rows.Scan(&tenantID, &currency, &available, &email, &name, &today, &baseline, &lowBalanceEnabled, &lowBalanceThreshold); err != nil {
return err
}
day := time.Now().UTC().Format("2006-01-02")
- if available <= config.LowBalanceMicros {
+ if lowBalanceEnabled && available <= lowBalanceThreshold {
body := fmt.Sprintf("Hi %s,\n\nYour AIGW prepaid balance is low: %.6f %s remains available. Add funds to avoid interrupted API access.\n", displayName(name), float64(available)/1_000_000, strings.ToUpper(currency))
if err := s.queueNotification(ctx, tenantID, email, "low_balance", day, "AIGW balance is low", body); err != nil {
return err
diff --git a/internal/controlplane/mutations.go b/internal/controlplane/mutations.go
index c2cb7d8..9f81d6e 100644
--- a/internal/controlplane/mutations.go
+++ b/internal/controlplane/mutations.go
@@ -17,8 +17,9 @@ import (
)
var (
- ErrNotFound = errors.New("control-plane resource not found")
- slugPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$`)
+ ErrNotFound = errors.New("control-plane resource not found")
+ slugPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$`)
+ nonSlugCharacters = regexp.MustCompile(`[^a-z0-9]+`)
)
func (s *Store) CreateTenant(ctx context.Context, input CreateTenantInput) (Tenant, int64, error) {
@@ -84,11 +85,28 @@ 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 input.ExpiresAt != nil && !input.ExpiresAt.After(time.Now()) {
+ return CreatedAPIKey{}, 0, errors.New("API key expiry must be in the future")
+ }
if len(input.Scopes) == 0 {
input.Scopes = []string{"inference"}
}
scopes := uniqueStrings(input.Scopes)
+ tags := uniqueStrings(input.Tags)
+ allowedModels := uniqueStrings(input.AllowedModels)
+ if len(scopes) > 20 || len(tags) > 20 || len(allowedModels) > 200 {
+ return CreatedAPIKey{}, 0, errors.New("API key has too many scopes, tags, or model restrictions")
+ }
+ for _, value := range append(append(append([]string{}, scopes...), tags...), allowedModels...) {
+ if len(value) > 160 {
+ return CreatedAPIKey{}, 0, errors.New("API key scope, tag, or model ID is too long")
+ }
+ }
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)
@@ -104,15 +122,32 @@ 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)
- VALUES ($1, $2, $3, $4, $5, $6)
- RETURNING id::text, tenant_id::text, project_id::text, name, key_prefix, scopes, status, created_at`,
- input.TenantID, input.ProjectID, input.Name, prefix, hash[:], scopesJSON,
- ).Scan(&result.ID, &result.TenantID, &result.ProjectID, &result.Name, &result.KeyPrefix, &scopesJSON, &result.Status, &result.CreatedAt)
+ 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,
+ &result.LastUsedAt, &result.CreatedAt)
if err != nil {
return CreatedAPIKey{}, 0, fmt.Errorf("create API key: %w", err)
}
+ if len(allowedModels) > 0 {
+ command, err := tx.Exec(ctx, `
+ INSERT INTO api_key_model_restrictions (api_key_id, model_id)
+ SELECT $1, id FROM models WHERE public_id = ANY($2::text[])`, result.ID, allowedModels)
+ if err != nil {
+ return CreatedAPIKey{}, 0, fmt.Errorf("restrict API key models: %w", err)
+ }
+ if command.RowsAffected() != int64(len(allowedModels)) {
+ return CreatedAPIKey{}, 0, errors.New("one or more allowed model IDs do not exist")
+ }
+ }
result.Scopes = scopes
+ result.Tags = tags
+ result.AllowedModels = allowedModels
result.Key = rawKey
generation, err := bumpGeneration(ctx, tx)
if err != nil {
@@ -129,10 +164,29 @@ func (s *Store) RevokeAPIKey(ctx context.Context, id string) (int64, error) {
}
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)
+ if input.Slug == "" {
+ input.Slug = strings.Trim(nonSlugCharacters.ReplaceAllString(strings.ToLower(input.Name), "-"), "-")
+ if len(input.Slug) > 64 {
+ input.Slug = strings.TrimRight(input.Slug[:64], "-")
+ }
+ }
input.BaseURL = strings.TrimRight(strings.TrimSpace(input.BaseURL), "/")
- if input.Name == "" || input.APIKey == "" || (input.Protocol != "openai" && input.Protocol != "anthropic") {
- return Provider{}, 0, errors.New("provider requires name, protocol openai|anthropic, base_url, and api_key")
+ input.WireAPI = strings.TrimSpace(input.WireAPI)
+ if input.WireAPI == "" {
+ if input.Protocol == "anthropic" {
+ input.WireAPI = "messages"
+ } else {
+ input.WireAPI = "chat_completions"
+ }
+ }
+ 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") ||
+ (input.Protocol == "anthropic" && input.WireAPI != "messages") {
+ return Provider{}, 0, errors.New("provider wire_api is incompatible with protocol")
}
parsed, err := url.Parse(input.BaseURL)
if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") {
@@ -149,11 +203,11 @@ func (s *Store) CreateProvider(ctx context.Context, input CreateProviderInput) (
defer tx.Rollback(ctx)
var result Provider
err = tx.QueryRow(ctx, `
- INSERT INTO providers (name, protocol, base_url, api_key_ciphertext)
- VALUES ($1, $2, $3, $4)
- RETURNING id::text, name, protocol, base_url, enabled, created_at`,
- input.Name, input.Protocol, input.BaseURL, ciphertext,
- ).Scan(&result.ID, &result.Name, &result.Protocol, &result.BaseURL, &result.Enabled, &result.CreatedAt)
+ INSERT INTO providers (slug, name, protocol, wire_api, base_url, api_key_ciphertext)
+ VALUES ($1, $2, $3, $4, $5, $6)
+ RETURNING id::text, slug, name, protocol, wire_api, base_url, enabled, created_at`,
+ input.Slug, input.Name, input.Protocol, input.WireAPI, input.BaseURL, ciphertext,
+ ).Scan(&result.ID, &result.Slug, &result.Name, &result.Protocol, &result.WireAPI, &result.BaseURL, &result.Enabled, &result.CreatedAt)
if err != nil {
return Provider{}, 0, fmt.Errorf("create provider: %w", err)
}
diff --git a/internal/controlplane/preferences.go b/internal/controlplane/preferences.go
new file mode 100644
index 0000000..73bcce6
--- /dev/null
+++ b/internal/controlplane/preferences.go
@@ -0,0 +1,156 @@
+package controlplane
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "strings"
+ "time"
+
+ "github.com/jackc/pgx/v5"
+)
+
+const maxLowBalanceThresholdMicros 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
+ }
+ result := TenantPreferences{TenantID: strings.TrimSpace(tenantID), LowBalanceEnabled: true, LowBalanceThresholdMicros: defaultThresholdMicros}
+ 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)
+ if errors.Is(err, pgx.ErrNoRows) {
+ return result, nil
+ }
+ if err != nil {
+ return TenantPreferences{}, fmt.Errorf("query tenant preferences: %w", err)
+ }
+ if defaultModel != nil {
+ result.DefaultModel = *defaultModel
+ }
+ if fallbackModel != nil {
+ result.FallbackModel = *fallbackModel
+ }
+ result.UpdatedAt = &updatedAt
+ return result, nil
+}
+
+func (s *Store) SetDeveloperPreferences(ctx context.Context, input SetDeveloperPreferencesInput) (TenantPreferences, error) {
+ input.TenantID = strings.TrimSpace(input.TenantID)
+ input.DefaultModel = strings.TrimSpace(input.DefaultModel)
+ input.FallbackModel = strings.TrimSpace(input.FallbackModel)
+ if input.TenantID == "" {
+ return TenantPreferences{}, errors.New("tenant_id is required")
+ }
+ if input.DefaultModel != "" && input.DefaultModel == input.FallbackModel {
+ return TenantPreferences{}, errors.New("default_model and fallback_model must be different")
+ }
+ available, err := s.ListDeveloperModels(ctx, input.TenantID)
+ if err != nil {
+ return TenantPreferences{}, err
+ }
+ allowed := make(map[string]struct{}, len(available))
+ for _, model := range available {
+ allowed[model.PublicID] = struct{}{}
+ }
+ for field, model := range map[string]string{"default_model": input.DefaultModel, "fallback_model": input.FallbackModel} {
+ if model != "" {
+ if _, ok := allowed[model]; !ok {
+ return TenantPreferences{}, fmt.Errorf("%s is not available to this tenant", field)
+ }
+ }
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return TenantPreferences{}, err
+ }
+ defer tx.Rollback(ctx)
+ if err := tx.QueryRow(ctx, `SELECT id FROM tenants WHERE id=$1`, input.TenantID).Scan(new(string)); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return TenantPreferences{}, ErrNotFound
+ }
+ return TenantPreferences{}, err
+ }
+ var result TenantPreferences
+ var defaultModel, fallbackModel *string
+ var updatedAt time.Time
+ if err := tx.QueryRow(ctx, `
+ INSERT INTO tenant_preferences (tenant_id, default_model, fallback_model)
+ VALUES ($1, NULLIF($2,''), NULLIF($3,''))
+ 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(
+ &result.TenantID, &defaultModel, &fallbackModel, &result.LowBalanceEnabled,
+ &result.LowBalanceThresholdMicros, &updatedAt); err != nil {
+ return TenantPreferences{}, fmt.Errorf("save developer preferences: %w", err)
+ }
+ if defaultModel != nil {
+ result.DefaultModel = *defaultModel
+ }
+ if fallbackModel != nil {
+ result.FallbackModel = *fallbackModel
+ }
+ result.UpdatedAt = &updatedAt
+ if err := tx.Commit(ctx); err != nil {
+ return TenantPreferences{}, err
+ }
+ return result, nil
+}
+
+func (s *Store) SetBillingPreferences(ctx context.Context, input SetBillingPreferencesInput, defaultThresholdMicros int64) (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) {
+ return TenantPreferences{}, errors.New("low balance threshold is outside the supported range")
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return TenantPreferences{}, err
+ }
+ defer tx.Rollback(ctx)
+ if err := tx.QueryRow(ctx, `SELECT id FROM tenants WHERE id=$1`, input.TenantID).Scan(new(string)); err != nil {
+ if errors.Is(err, pgx.ErrNoRows) {
+ return TenantPreferences{}, ErrNotFound
+ }
+ return TenantPreferences{}, err
+ }
+ var result TenantPreferences
+ 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))
+ 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()
+ 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(
+ &result.TenantID, &defaultModel, &fallbackModel, &result.LowBalanceEnabled,
+ &result.LowBalanceThresholdMicros, &updatedAt); err != nil {
+ return TenantPreferences{}, fmt.Errorf("save billing preferences: %w", err)
+ }
+ if defaultModel != nil {
+ result.DefaultModel = *defaultModel
+ }
+ if fallbackModel != nil {
+ result.FallbackModel = *fallbackModel
+ }
+ result.UpdatedAt = &updatedAt
+ if err := tx.Commit(ctx); err != nil {
+ return TenantPreferences{}, err
+ }
+ return result, nil
+}
diff --git a/internal/controlplane/preferences_test.go b/internal/controlplane/preferences_test.go
new file mode 100644
index 0000000..88a3d50
--- /dev/null
+++ b/internal/controlplane/preferences_test.go
@@ -0,0 +1,31 @@
+package controlplane
+
+import (
+ "context"
+ "testing"
+)
+
+func TestGetTenantPreferencesWithoutTenantUsesConfiguredDefault(t *testing.T) {
+ result, err := (&Store{}).GetTenantPreferences(context.Background(), "", 12_500_000)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if !result.LowBalanceEnabled || result.LowBalanceThresholdMicros != 12_500_000 {
+ t.Fatalf("unexpected defaults: %+v", result)
+ }
+}
+
+func TestPreferenceValidationRejectsUnsafeValuesBeforeDatabaseAccess(t *testing.T) {
+ store := &Store{}
+ if _, err := store.SetDeveloperPreferences(context.Background(), SetDeveloperPreferencesInput{
+ TenantID: "tenant", DefaultModel: "same", FallbackModel: "same",
+ }); err == nil {
+ t.Fatal("expected identical default and fallback models to fail")
+ }
+ threshold := maxLowBalanceThresholdMicros + 1
+ if _, err := store.SetBillingPreferences(context.Background(), SetBillingPreferencesInput{
+ TenantID: "tenant", LowBalanceThresholdMicros: &threshold,
+ }, 5_000_000); err == nil {
+ t.Fatal("expected excessive low balance threshold to fail")
+ }
+}
diff --git a/internal/controlplane/queries.go b/internal/controlplane/queries.go
index 73d869f..48610c7 100644
--- a/internal/controlplane/queries.go
+++ b/internal/controlplane/queries.go
@@ -4,6 +4,7 @@ import (
"context"
"encoding/json"
"fmt"
+ "time"
)
func (s *Store) Overview(ctx context.Context) (Overview, error) {
@@ -84,15 +85,35 @@ 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)
+ periodEnd := periodStart.AddDate(0, 1, 0)
query := `
- SELECT id::text, tenant_id::text, project_id::text, name, key_prefix, scopes, status, created_at
- FROM api_keys`
- args := []any{}
+ 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,
+ usage.month_spend, usage.month_requests, pending.month_reserved,
+ COALESCE((
+ SELECT jsonb_agg(m.public_id ORDER BY m.public_id)
+ FROM api_key_model_restrictions r
+ JOIN models m ON m.id = r.model_id
+ WHERE r.api_key_id = k.id
+ ), '[]'::jsonb)
+ FROM api_keys k
+ CROSS JOIN LATERAL (
+ SELECT COALESCE(SUM(u.cost_micros), 0)::bigint AS month_spend, COUNT(*)::bigint AS month_requests
+ FROM usage_events u WHERE u.key_id = k.id AND u.started_at >= $1 AND u.started_at < $2
+ ) usage
+ CROSS JOIN LATERAL (
+ SELECT COALESCE(SUM(b.reserved_micros), 0)::bigint AS month_reserved
+ 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}
if tenantID != "" {
- query += ` WHERE tenant_id=$1`
+ query += ` WHERE k.tenant_id=$3`
args = append(args, tenantID)
}
- query += ` ORDER BY created_at DESC`
+ query += ` ORDER BY k.created_at DESC`
rows, err := s.db.Query(ctx, query, args...)
if err != nil {
return nil, fmt.Errorf("query API keys: %w", err)
@@ -101,13 +122,22 @@ func (s *Store) ListAPIKeysFor(ctx context.Context, tenantID string) ([]APIKey,
result := make([]APIKey, 0)
for rows.Next() {
var item APIKey
- var scopesJSON []byte
- if err := rows.Scan(&item.ID, &item.TenantID, &item.ProjectID, &item.Name, &item.KeyPrefix, &scopesJSON, &item.Status, &item.CreatedAt); err != nil {
+ 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,
+ &item.LastUsedAt, &item.CreatedAt, &item.CurrentMonthSpendMicros, &item.CurrentMonthRequests,
+ &item.CurrentMonthReservedMicros, &allowedModelsJSON); err != nil {
return nil, fmt.Errorf("scan API key: %w", err)
}
if err := json.Unmarshal(scopesJSON, &item.Scopes); err != nil {
return nil, fmt.Errorf("decode API key scopes: %w", err)
}
+ if err := json.Unmarshal(tagsJSON, &item.Tags); err != nil {
+ return nil, fmt.Errorf("decode API key tags: %w", err)
+ }
+ if err := json.Unmarshal(allowedModelsJSON, &item.AllowedModels); err != nil {
+ return nil, fmt.Errorf("decode API key model restrictions: %w", err)
+ }
result = append(result, item)
}
return result, rows.Err()
@@ -153,7 +183,7 @@ func (s *Store) ResourceTenantID(ctx context.Context, resource, id string) (stri
func (s *Store) ListProviders(ctx context.Context) ([]Provider, error) {
rows, err := s.db.Query(ctx, `
- SELECT p.id::text, p.name, p.protocol, p.base_url, p.enabled, count(r.id), p.created_at
+ SELECT p.id::text, p.slug, p.name, p.protocol, p.wire_api, p.base_url, p.enabled, count(r.id), p.created_at
FROM providers p LEFT JOIN model_routes r ON r.provider_id = p.id
GROUP BY p.id ORDER BY p.created_at DESC`)
if err != nil {
@@ -163,7 +193,7 @@ func (s *Store) ListProviders(ctx context.Context) ([]Provider, error) {
result := make([]Provider, 0)
for rows.Next() {
var item Provider
- if err := rows.Scan(&item.ID, &item.Name, &item.Protocol, &item.BaseURL, &item.Enabled, &item.RouteCount, &item.CreatedAt); err != nil {
+ if err := rows.Scan(&item.ID, &item.Slug, &item.Name, &item.Protocol, &item.WireAPI, &item.BaseURL, &item.Enabled, &item.RouteCount, &item.CreatedAt); err != nil {
return nil, fmt.Errorf("scan provider: %w", err)
}
result = append(result, item)
@@ -182,6 +212,7 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) {
m.enabled, m.created_at
FROM models m JOIN LATERAL (
SELECT * FROM model_price_versions v WHERE v.model_id=m.id
+ AND v.effective_from <= now() AND (v.effective_to IS NULL OR v.effective_to > now())
ORDER BY v.effective_from DESC LIMIT 1
) pv ON TRUE ORDER BY m.public_id`)
if err != nil {
@@ -264,7 +295,7 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) {
keyRows.Close()
routeRows, err := s.db.Query(ctx, `
- SELECT r.id::text, r.model_id::text, r.provider_id::text, p.name, p.protocol,
+ SELECT r.id::text, r.model_id::text, r.provider_id::text, p.name, p.protocol, p.wire_api, p.enabled,
r.upstream_model, r.priority, r.weight, r.enabled
FROM model_routes r JOIN providers p ON p.id = r.provider_id
ORDER BY r.priority, r.created_at`)
@@ -275,7 +306,7 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) {
for routeRows.Next() {
var route Route
var modelID string
- if err := routeRows.Scan(&route.ID, &modelID, &route.ProviderID, &route.ProviderName, &route.Protocol, &route.UpstreamModel, &route.Priority, &route.Weight, &route.Enabled); err != nil {
+ if err := routeRows.Scan(&route.ID, &modelID, &route.ProviderID, &route.ProviderName, &route.Protocol, &route.WireAPI, &route.ProviderEnabled, &route.UpstreamModel, &route.Priority, &route.Weight, &route.Enabled); err != nil {
return nil, fmt.Errorf("scan model route: %w", err)
}
if position, ok := positions[modelID]; ok {
@@ -284,3 +315,141 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) {
}
return models, routeRows.Err()
}
+
+func (s *Store) ListDeveloperModels(ctx context.Context, tenantID string) ([]DeveloperModel, error) {
+ models, err := s.ListModels(ctx)
+ if err != nil {
+ return nil, err
+ }
+ keyIDs := map[string]struct{}{}
+ if tenantID != "" {
+ rows, err := s.db.Query(ctx, `SELECT id::text FROM api_keys WHERE tenant_id=$1 AND status='active'`, tenantID)
+ if err != nil {
+ return nil, fmt.Errorf("query developer API keys: %w", err)
+ }
+ for rows.Next() {
+ var id string
+ if err := rows.Scan(&id); err != nil {
+ rows.Close()
+ return nil, err
+ }
+ keyIDs[id] = struct{}{}
+ }
+ if err := rows.Err(); err != nil {
+ rows.Close()
+ return nil, err
+ }
+ rows.Close()
+ }
+ return developerModelsFor(models, tenantID, keyIDs), nil
+}
+
+func (s *Store) ListPublicModels(ctx context.Context) ([]PublicModel, error) {
+ models, err := s.ListModels(ctx)
+ if err != nil {
+ return nil, err
+ }
+ return publicModelsFor(models), nil
+}
+
+func publicModelsFor(models []Model) []PublicModel {
+ result := make([]PublicModel, 0, len(models))
+ for _, model := range models {
+ if !model.Enabled || model.Lifecycle == "retired" || len(model.AllowedTenantIDs) != 0 || len(model.AllowedKeyIDs) != 0 {
+ continue
+ }
+ wireSet := make(map[string]struct{})
+ providerSet := make(map[string]struct{})
+ wireAPIs := make([]string, 0, len(model.Routes))
+ for _, route := range model.Routes {
+ if !route.Enabled || !route.ProviderEnabled {
+ continue
+ }
+ wireAPI := route.WireAPI
+ if wireAPI == "" && route.Protocol == "anthropic" {
+ wireAPI = "messages"
+ } else if wireAPI == "" {
+ wireAPI = "chat_completions"
+ }
+ if _, exists := wireSet[wireAPI]; !exists {
+ wireSet[wireAPI] = struct{}{}
+ wireAPIs = append(wireAPIs, wireAPI)
+ }
+ providerSet[route.ProviderID] = struct{}{}
+ }
+ if len(wireAPIs) == 0 {
+ continue
+ }
+ result = append(result, PublicModel{PublicID: model.PublicID, DisplayName: model.DisplayName,
+ Description: model.Description, OwnedBy: model.OwnedBy, InputModalities: model.InputModalities,
+ OutputModalities: model.OutputModalities, ContextWindow: model.ContextWindow, MaxOutputTokens: model.MaxOutputTokens,
+ Capabilities: model.Capabilities, Regions: model.Regions, Lifecycle: model.Lifecycle, ReleasedAt: model.ReleasedAt,
+ ReplacementModel: model.ReplacementModel, Aliases: model.Aliases, PriceCurrency: model.PriceCurrency,
+ InputPriceMicrosPerMillion: model.InputPriceMicrosPerMillion, OutputPriceMicrosPerMillion: model.OutputPriceMicrosPerMillion,
+ CacheReadPriceMicrosPerMillion: model.CacheReadPriceMicrosPerMillion, CacheWritePriceMicrosPerMillion: model.CacheWritePriceMicrosPerMillion,
+ SupportedWireAPIs: wireAPIs, ProviderCount: len(providerSet), AvailableProviderCount: len(providerSet), HealthStatus: "available"})
+ }
+ return result
+}
+
+func developerModelsFor(models []Model, tenantID string, keyIDs map[string]struct{}) []DeveloperModel {
+ result := make([]DeveloperModel, 0, len(models))
+ for _, model := range models {
+ if !model.Enabled || model.Lifecycle == "retired" || !stringAllowed(model.AllowedTenantIDs, tenantID) || !keyAllowed(model.AllowedKeyIDs, keyIDs) {
+ continue
+ }
+ wireSet := map[string]struct{}{}
+ wireAPIs := make([]string, 0, len(model.Routes))
+ for _, route := range model.Routes {
+ if !route.Enabled || !route.ProviderEnabled {
+ continue
+ }
+ wireAPI := route.WireAPI
+ if wireAPI == "" && route.Protocol == "anthropic" {
+ wireAPI = "messages"
+ } else if wireAPI == "" {
+ wireAPI = "chat_completions"
+ }
+ if _, exists := wireSet[wireAPI]; !exists {
+ wireSet[wireAPI] = struct{}{}
+ wireAPIs = append(wireAPIs, wireAPI)
+ }
+ }
+ if len(wireAPIs) == 0 {
+ continue
+ }
+ result = append(result, DeveloperModel{ID: model.ID, PublicID: model.PublicID, DisplayName: model.DisplayName,
+ Description: model.Description, OwnedBy: model.OwnedBy, InputModalities: model.InputModalities,
+ OutputModalities: model.OutputModalities, ContextWindow: model.ContextWindow, MaxOutputTokens: model.MaxOutputTokens,
+ Capabilities: model.Capabilities, Regions: model.Regions, Lifecycle: model.Lifecycle, ReleasedAt: model.ReleasedAt,
+ ReplacementModel: model.ReplacementModel, Aliases: model.Aliases, PriceCurrency: model.PriceCurrency,
+ InputPriceMicrosPerMillion: model.InputPriceMicrosPerMillion, OutputPriceMicrosPerMillion: model.OutputPriceMicrosPerMillion,
+ CacheReadPriceMicrosPerMillion: model.CacheReadPriceMicrosPerMillion, CacheWritePriceMicrosPerMillion: model.CacheWritePriceMicrosPerMillion,
+ SupportedWireAPIs: wireAPIs})
+ }
+ return result
+}
+
+func stringAllowed(allowed []string, value string) bool {
+ if len(allowed) == 0 || value == "" {
+ return true
+ }
+ for _, item := range allowed {
+ if item == value {
+ return true
+ }
+ }
+ return false
+}
+
+func keyAllowed(allowed []string, keys map[string]struct{}) bool {
+ if len(allowed) == 0 {
+ return true
+ }
+ for _, id := range allowed {
+ if _, ok := keys[id]; ok {
+ return true
+ }
+ }
+ return false
+}
diff --git a/internal/controlplane/queries_test.go b/internal/controlplane/queries_test.go
new file mode 100644
index 0000000..4cd44b9
--- /dev/null
+++ b/internal/controlplane/queries_test.go
@@ -0,0 +1,80 @@
+package controlplane
+
+import (
+ "encoding/json"
+ "strings"
+ "testing"
+)
+
+func TestDeveloperModelsForFiltersScopeAndRedactsRouting(t *testing.T) {
+ models := []Model{
+ {ID: "public", PublicID: "acme/public", DisplayName: "Public", Enabled: true, Lifecycle: "active",
+ AllowedTenantIDs: nil, Routes: []Route{{Protocol: "openai", WireAPI: "responses", Enabled: true, ProviderEnabled: true, ProviderName: "secret-provider", UpstreamModel: "secret-model", Priority: 1, Weight: 100}}},
+ {ID: "tenant", PublicID: "acme/private", Enabled: true, Lifecycle: "active", AllowedTenantIDs: []string{"tenant-a"},
+ Routes: []Route{{Protocol: "anthropic", WireAPI: "messages", Enabled: true, ProviderEnabled: true}}},
+ {ID: "key", PublicID: "acme/key", Enabled: true, Lifecycle: "active", AllowedKeyIDs: []string{"key-a"},
+ Routes: []Route{{Protocol: "openai", Enabled: true, ProviderEnabled: true}}},
+ {ID: "retired", PublicID: "acme/retired", Enabled: true, Lifecycle: "retired", Routes: []Route{{Protocol: "openai", Enabled: true, ProviderEnabled: true}}},
+ {ID: "disabled", PublicID: "acme/disabled", Enabled: false, Lifecycle: "active", Routes: []Route{{Protocol: "openai", Enabled: true, ProviderEnabled: true}}},
+ {ID: "noroute", PublicID: "acme/noroute", Enabled: true, Lifecycle: "active", Routes: []Route{{Protocol: "openai", Enabled: false, ProviderEnabled: true}}},
+ {ID: "provider-off", PublicID: "acme/provider-off", Enabled: true, Lifecycle: "active", Routes: []Route{{Protocol: "openai", Enabled: true, ProviderEnabled: false}}},
+ }
+ got := developerModelsFor(models, "tenant-a", map[string]struct{}{"key-a": {}})
+ if len(got) != 3 {
+ t.Fatalf("developer model count = %d, want 3", len(got))
+ }
+ if got[0].PublicID != "acme/key" && got[1].PublicID != "acme/key" && got[2].PublicID != "acme/key" {
+ t.Fatal("key-allowlisted model was not included for an active tenant key")
+ }
+ for _, item := range got {
+ if len(item.SupportedWireAPIs) == 0 {
+ t.Fatalf("unexpected developer model: %+v", item)
+ }
+ if item.PublicID == "acme/public" && item.SupportedWireAPIs[0] != "responses" {
+ t.Fatalf("responses wire API was not preserved: %+v", item)
+ }
+ }
+ for _, item := range got {
+ if item.PublicID == "acme/private" && item.SupportedWireAPIs[0] != "messages" {
+ t.Fatalf("anthropic wire API was not preserved: %+v", item)
+ }
+ }
+ encoded, err := json.Marshal(got)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if strings.Contains(string(encoded), "secret-provider") || strings.Contains(string(encoded), "secret-model") || strings.Contains(string(encoded), "provider_id") {
+ t.Fatalf("developer model response leaked routing metadata: %s", encoded)
+ }
+}
+
+func TestDeveloperModelsForDoesNotExposeKeyRestrictedModelsWithoutTenantKey(t *testing.T) {
+ models := []Model{{ID: "key", PublicID: "key-only", Enabled: true, Lifecycle: "active", AllowedKeyIDs: []string{"key-a"}, Routes: []Route{{Protocol: "openai", Enabled: true, ProviderEnabled: true}}}}
+ if got := developerModelsFor(models, "tenant-a", map[string]struct{}{}); len(got) != 0 {
+ t.Fatalf("key-restricted models visible without a matching key: %+v", got)
+ }
+}
+
+func TestPublicModelsForOnlyExposesUnrestrictedCatalogData(t *testing.T) {
+ models := []Model{
+ {ID: "public", PublicID: "acme/public", DisplayName: "Public", Enabled: true, Lifecycle: "active",
+ Routes: []Route{{ProviderID: "provider-a", ProviderName: "internal provider", Protocol: "openai", WireAPI: "responses", UpstreamModel: "secret-model", Enabled: true, ProviderEnabled: true}}},
+ {ID: "tenant", PublicID: "acme/tenant", Enabled: true, Lifecycle: "active", AllowedTenantIDs: []string{"tenant-a"},
+ Routes: []Route{{ProviderID: "provider-a", Protocol: "openai", Enabled: true, ProviderEnabled: true}}},
+ {ID: "key", PublicID: "acme/key", Enabled: true, Lifecycle: "active", AllowedKeyIDs: []string{"key-a"},
+ Routes: []Route{{ProviderID: "provider-a", Protocol: "openai", Enabled: true, ProviderEnabled: true}}},
+ }
+ got := publicModelsFor(models)
+ if len(got) != 1 || got[0].PublicID != "acme/public" || got[0].ProviderCount != 1 || got[0].SupportedWireAPIs[0] != "responses" {
+ t.Fatalf("unexpected public catalog: %+v", got)
+ }
+ encoded, err := json.Marshal(got)
+ if err != nil {
+ t.Fatal(err)
+ }
+ for _, forbidden := range []string{"secret-model", "internal provider", "provider_id", "upstream_model", "allowed_tenant"} {
+ if strings.Contains(string(encoded), forbidden) {
+ t.Fatalf("public catalog leaked %q: %s", forbidden, encoded)
+ }
+ }
+}
diff --git a/internal/controlplane/schema.sql b/internal/controlplane/schema.sql
index 88db49b..ef1ccdd 100644
--- a/internal/controlplane/schema.sql
+++ b/internal/controlplane/schema.sql
@@ -49,17 +49,53 @@ CREATE TABLE IF NOT EXISTS api_keys (
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 expires_at TIMESTAMPTZ;
+ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS tags JSONB NOT NULL DEFAULT '[]'::jsonb;
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')),
base_url TEXT NOT NULL,
api_key_ciphertext BYTEA NOT NULL,
enabled BOOLEAN NOT NULL DEFAULT TRUE,
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
);
+ALTER TABLE providers ADD COLUMN IF NOT EXISTS slug TEXT;
+WITH normalized AS (
+ SELECT id,
+ left(trim(both '-' from regexp_replace(lower(name), '[^a-z0-9]+', '-', 'g')), 64) AS base
+ FROM providers
+), ranked AS (
+ SELECT id, base, count(*) OVER (PARTITION BY base) AS base_count
+ FROM normalized
+)
+UPDATE providers p
+SET slug = CASE
+ WHEN length(r.base) BETWEEN 3 AND 64
+ AND r.base ~ '^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$'
+ AND r.base_count = 1 THEN r.base
+ ELSE 'provider-' || left(replace(p.id::text, '-', ''), 12)
+END
+FROM ranked r
+WHERE p.id = r.id AND (p.slug IS NULL OR p.slug = '');
+ALTER TABLE providers ALTER COLUMN slug SET NOT NULL;
+CREATE UNIQUE INDEX IF NOT EXISTS providers_slug_unique_idx ON providers (slug);
+ALTER TABLE providers DROP CONSTRAINT IF EXISTS providers_slug_check;
+ALTER TABLE providers ADD CONSTRAINT providers_slug_check CHECK (slug ~ '^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$');
+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 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 = 'anthropic' AND wire_api = 'messages')
+);
CREATE TABLE IF NOT EXISTS models (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
@@ -145,6 +181,15 @@ CREATE TABLE IF NOT EXISTS api_key_model_allowlist (
PRIMARY KEY (api_key_id, model_id)
);
+-- Per-key restrictions are separate from the platform model allowlist above:
+-- an empty set means that the key may use every model visible to its tenant.
+CREATE TABLE IF NOT EXISTS api_key_model_restrictions (
+ api_key_id UUID NOT NULL REFERENCES api_keys(id) ON DELETE CASCADE,
+ model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ PRIMARY KEY (api_key_id, model_id)
+);
+
CREATE TABLE IF NOT EXISTS model_routes (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE,
@@ -171,6 +216,18 @@ CREATE TABLE IF NOT EXISTS tenant_wallets (
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
);
+-- Tenant-scoped developer and billing preferences. These values are control
+-- plane data, but are intentionally not loaded into the inference snapshot.
+CREATE TABLE IF NOT EXISTS tenant_preferences (
+ tenant_id UUID PRIMARY KEY REFERENCES tenants(id) ON DELETE CASCADE,
+ default_model TEXT REFERENCES models(public_id) ON DELETE SET NULL,
+ 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),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ CHECK (default_model IS NULL OR fallback_model IS NULL OR default_model <> fallback_model)
+);
+
CREATE TABLE IF NOT EXISTS billing_reservations (
request_id TEXT PRIMARY KEY,
tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE CASCADE,
@@ -286,6 +343,10 @@ CREATE TABLE IF NOT EXISTS topup_orders (
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
paid_at TIMESTAMPTZ
);
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS trigger_type TEXT NOT NULL DEFAULT 'manual';
+ALTER TABLE topup_orders DROP CONSTRAINT IF EXISTS topup_orders_trigger_type_check;
+ALTER TABLE topup_orders ADD CONSTRAINT topup_orders_trigger_type_check
+ CHECK (trigger_type IN ('manual','auto'));
ALTER TABLE topup_orders DROP CONSTRAINT IF EXISTS topup_orders_status_check;
ALTER TABLE topup_orders ADD CONSTRAINT topup_orders_status_check
CHECK (status IN ('pending','paid','failed','expired','partially_refunded','refunded','disputed','reversed'));
@@ -306,6 +367,8 @@ ALTER TABLE topup_orders ADD CONSTRAINT topup_orders_reconciliation_status_check
CHECK (reconciliation_status IN ('unknown','ok','repaired','missing','mismatch','resolved'));
CREATE INDEX IF NOT EXISTS topup_orders_payment_intent_idx ON topup_orders (stripe_payment_intent_id) WHERE stripe_payment_intent_id IS NOT NULL;
CREATE INDEX IF NOT EXISTS topup_orders_customer_idx ON topup_orders (stripe_customer_id) WHERE stripe_customer_id IS NOT NULL;
+CREATE UNIQUE INDEX IF NOT EXISTS topup_orders_auto_pending_idx ON topup_orders (tenant_id)
+ WHERE trigger_type='auto' AND status='pending';
CREATE TABLE IF NOT EXISTS billing_reconciliation_resolutions (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
@@ -325,6 +388,51 @@ CREATE TABLE IF NOT EXISTS stripe_customers (
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
);
+CREATE TABLE IF NOT EXISTS tenant_billing_profiles (
+ tenant_id UUID PRIMARY KEY REFERENCES tenants(id) ON DELETE CASCADE,
+ legal_name TEXT NOT NULL,
+ billing_email TEXT NOT NULL,
+ address_line1 TEXT NOT NULL,
+ address_line2 TEXT NOT NULL DEFAULT '',
+ city TEXT NOT NULL,
+ region TEXT NOT NULL DEFAULT '',
+ postal_code TEXT NOT NULL,
+ country TEXT NOT NULL CHECK (country ~ '^[A-Z]{2}$'),
+ stripe_sync_status TEXT NOT NULL DEFAULT 'pending'
+ CHECK (stripe_sync_status IN ('pending','synced','failed','disabled')),
+ stripe_synced_at TIMESTAMPTZ,
+ stripe_sync_error TEXT NOT NULL DEFAULT '',
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE TABLE IF NOT EXISTS tenant_auto_topup_settings (
+ tenant_id UUID PRIMARY KEY REFERENCES tenants(id) ON DELETE CASCADE,
+ enabled BOOLEAN NOT NULL DEFAULT FALSE,
+ threshold_micros BIGINT NOT NULL CHECK (threshold_micros >= 0),
+ topup_amount_minor BIGINT NOT NULL CHECK (topup_amount_minor > 0),
+ stripe_payment_method_id TEXT,
+ payment_method_type TEXT NOT NULL DEFAULT '',
+ payment_method_brand TEXT NOT NULL DEFAULT '',
+ payment_method_last4 TEXT NOT NULL DEFAULT '',
+ payment_method_exp_month INTEGER NOT NULL DEFAULT 0 CHECK (payment_method_exp_month BETWEEN 0 AND 12),
+ payment_method_exp_year INTEGER NOT NULL DEFAULT 0 CHECK (payment_method_exp_year >= 0),
+ stripe_setup_session_id TEXT UNIQUE,
+ status TEXT NOT NULL DEFAULT 'not_configured'
+ CHECK (status IN ('not_configured','ready','charging','action_required','failed')),
+ last_error TEXT NOT NULL DEFAULT '',
+ failure_count INTEGER NOT NULL DEFAULT 0 CHECK (failure_count >= 0),
+ last_attempt_at TIMESTAMPTZ,
+ last_succeeded_at TIMESTAMPTZ,
+ next_attempt_at TIMESTAMPTZ,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ CHECK (stripe_payment_method_id IS NOT NULL OR enabled = FALSE)
+);
+CREATE INDEX IF NOT EXISTS tenant_auto_topup_ready_idx
+ ON tenant_auto_topup_settings (next_attempt_at, tenant_id)
+ WHERE enabled AND stripe_payment_method_id IS NOT NULL;
+
CREATE TABLE IF NOT EXISTS stripe_refunds (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE RESTRICT,
@@ -658,8 +766,10 @@ CREATE INDEX IF NOT EXISTS billing_ledger_tenant_idx ON billing_ledger (tenant_i
CREATE INDEX IF NOT EXISTS usage_events_tenant_idx ON usage_events (tenant_id, created_at DESC);
CREATE INDEX IF NOT EXISTS usage_events_project_idx ON usage_events (project_id, created_at DESC);
CREATE INDEX IF NOT EXISTS usage_events_model_idx ON usage_events (public_model, created_at DESC);
+CREATE INDEX IF NOT EXISTS usage_events_key_idx ON usage_events (key_id, started_at DESC);
CREATE INDEX IF NOT EXISTS billing_reservations_pending_idx ON billing_reservations (status, created_at) WHERE status = 'pending';
CREATE INDEX IF NOT EXISTS billing_reservations_project_pending_idx ON billing_reservations (project_id, created_at) WHERE status = 'pending';
+CREATE INDEX IF NOT EXISTS billing_reservations_key_period_idx ON billing_reservations (key_id, created_at DESC) WHERE status IN ('pending', 'metering_failed');
CREATE INDEX IF NOT EXISTS console_users_tenant_idx ON console_users (tenant_id, created_at DESC);
CREATE INDEX IF NOT EXISTS console_sessions_user_idx ON console_sessions (user_id, created_at DESC);
CREATE INDEX IF NOT EXISTS console_sessions_expiry_idx ON console_sessions (expires_at) WHERE revoked_at IS NULL;
diff --git a/internal/controlplane/snapshot.go b/internal/controlplane/snapshot.go
index b9bddaf..75a3a80 100644
--- a/internal/controlplane/snapshot.go
+++ b/internal/controlplane/snapshot.go
@@ -74,7 +74,7 @@ func loadLimitPolicies(ctx context.Context, tx pgx.Tx) ([]domain.LimitPolicy, er
func (s *Store) loadProviders(ctx context.Context, tx pgx.Tx) (map[string]domain.Provider, error) {
rows, err := tx.Query(ctx, `
- SELECT id::text, name, protocol, base_url, api_key_ciphertext
+ SELECT id::text, slug, name, protocol, wire_api, base_url, api_key_ciphertext
FROM providers
WHERE enabled = TRUE
ORDER BY name`)
@@ -84,16 +84,16 @@ func (s *Store) loadProviders(ctx context.Context, tx pgx.Tx) (map[string]domain
defer rows.Close()
providers := make(map[string]domain.Provider)
for rows.Next() {
- var id, name, protocol, baseURL string
+ var id, slug, name, protocol, wireAPI, baseURL string
var ciphertext []byte
- if err := rows.Scan(&id, &name, &protocol, &baseURL, &ciphertext); err != nil {
+ if err := rows.Scan(&id, &slug, &name, &protocol, &wireAPI, &baseURL, &ciphertext); err != nil {
return nil, fmt.Errorf("scan provider: %w", err)
}
apiKey, err := s.cipher.Decrypt(ciphertext)
if err != nil {
return nil, fmt.Errorf("decrypt provider %q credential: %w", name, err)
}
- providers[id] = domain.Provider{ID: id, Protocol: domain.Protocol(protocol), BaseURL: strings.TrimRight(baseURL, "/"), APIKey: apiKey}
+ providers[id] = domain.Provider{ID: id, Slug: slug, Protocol: domain.Protocol(protocol), WireAPI: wireAPI, BaseURL: strings.TrimRight(baseURL, "/"), APIKey: apiKey}
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("read providers: %w", err)
@@ -245,11 +245,18 @@ 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
+ SELECT k.id::text, k.key_hash, k.tenant_id::text, k.project_id::text, k.scopes,
+ k.monthly_spend_micros, k.expires_at,
+ COALESCE((
+ SELECT jsonb_agg(m.public_id ORDER BY m.public_id)
+ FROM api_key_model_restrictions r
+ JOIN models m ON m.id = r.model_id
+ WHERE r.api_key_id = k.id
+ ), '[]'::jsonb)
FROM api_keys k
JOIN tenants t ON t.id = k.tenant_id AND t.status = 'active'
JOIN projects p ON p.id = k.project_id AND p.status = 'active'
- WHERE k.status = 'active'`)
+ WHERE k.status = 'active' AND (k.expires_at IS NULL OR k.expires_at > now())`)
if err != nil {
return nil, fmt.Errorf("query API keys: %w", err)
}
@@ -257,8 +264,11 @@ func loadAPIKeys(ctx context.Context, tx pgx.Tx) ([]auth.HashedKeyRecord, error)
records := make([]auth.HashedKeyRecord, 0)
for rows.Next() {
var keyID, tenantID, projectID string
- var hashBytes, scopesJSON []byte
- if err := rows.Scan(&keyID, &hashBytes, &tenantID, &projectID, &scopesJSON); err != nil {
+ var hashBytes, scopesJSON, allowedModelsJSON []byte
+ var monthlySpendMicros int64
+ var expiresAt *time.Time
+ if err := rows.Scan(&keyID, &hashBytes, &tenantID, &projectID, &scopesJSON,
+ &monthlySpendMicros, &expiresAt, &allowedModelsJSON); err != nil {
return nil, fmt.Errorf("scan API key: %w", err)
}
if len(hashBytes) != sha256.Size {
@@ -270,8 +280,17 @@ func loadAPIKeys(ctx context.Context, tx pgx.Tx) ([]auth.HashedKeyRecord, error)
if err := json.Unmarshal(scopesJSON, &scopes); err != nil {
return nil, fmt.Errorf("decode API key %s scopes: %w", keyID, err)
}
+ var modelIDs []string
+ if err := json.Unmarshal(allowedModelsJSON, &modelIDs); err != nil {
+ return nil, fmt.Errorf("decode API key %s model restrictions: %w", keyID, err)
+ }
+ allowedModels := make(map[string]struct{}, len(modelIDs))
+ for _, modelID := range modelIDs {
+ allowedModels[modelID] = struct{}{}
+ }
records = append(records, auth.HashedKeyRecord{Hash: hash, Principal: domain.Principal{
KeyID: keyID, TenantID: tenantID, ProjectID: projectID, Scopes: scopes,
+ AllowedModels: allowedModels, MonthlySpendMicros: monthlySpendMicros, ExpiresAt: expiresAt,
}})
}
if err := rows.Err(); err != nil {
diff --git a/internal/controlplane/store.go b/internal/controlplane/store.go
index fb48d78..c4d7016 100644
--- a/internal/controlplane/store.go
+++ b/internal/controlplane/store.go
@@ -23,7 +23,7 @@ var schemaSQL string
var ErrRedisDisabled = errors.New("Redis propagation is disabled")
-const migrationVersion int64 = 2026080504
+const migrationVersion int64 = 2026080605
type Options struct {
DatabaseURL string
@@ -133,7 +133,7 @@ func applySchema(ctx context.Context, db *pgxpool.Pool) error {
if !errors.Is(err, pgx.ErrNoRows) && err != nil {
return err
}
- if _, err := tx.Exec(ctx, `INSERT INTO schema_migrations(version,name,checksum) VALUES ($1,$2,$3) ON CONFLICT DO NOTHING`, migrationVersion, "commercial-control-plane", 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-billing-profiles", 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 4f55e10..9bf093c 100644
--- a/internal/controlplane/store_integration_test.go
+++ b/internal/controlplane/store_integration_test.go
@@ -2,6 +2,7 @@ package controlplane
import (
"context"
+ "encoding/base64"
"fmt"
"net/url"
"os"
@@ -11,6 +12,128 @@ import (
"github.com/jackc/pgx/v5/pgxpool"
)
+func TestAPIKeyRestrictionsRoundTripIntoRuntimeSnapshotPostgres(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("key_controls_%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: "key-control-test", Name: "Key Control Test"})
+ if err != nil {
+ t.Fatal(err)
+ }
+ project, _, err := store.CreateProject(ctx, CreateProjectInput{TenantID: tenant.ID, Slug: "production", Name: "Production"})
+ 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"})
+ 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 != 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})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if created.Key == "" || created.MonthlySpendMicros != 25_000_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
+ (request_id,tenant_id,project_id,key_id,public_model,protocol,status_code,success,started_at,cost_micros)
+ VALUES ('req_key_current_month',$1,$2,$3,$4,'responses',200,TRUE,now(),42000),
+ ('req_key_previous_month',$1,$2,$3,$4,'responses',200,TRUE,now()-interval '2 months',99000)`,
+ tenant.ID, project.ID, created.ID, model.PublicID); err != nil {
+ t.Fatal(err)
+ }
+ if _, err := store.db.Exec(ctx, `INSERT INTO billing_reservations
+ (request_id,tenant_id,project_id,key_id,public_model,currency,reserved_micros,status,
+ input_price_micros_per_million,output_price_micros_per_million,cache_read_price_micros_per_million,cache_write_price_micros_per_million)
+ VALUES ('req_key_pending',$1,$2,$3,$4,'usd',9000,'pending',100000,200000,0,0)`,
+ tenant.ID, project.ID, created.ID, model.PublicID); err != nil {
+ t.Fatal(err)
+ }
+ keys, err := store.ListAPIKeysFor(ctx, tenant.ID)
+ 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) {
+ 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])
+ }
+ snapshot, err := store.LoadSnapshot(ctx)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(snapshot.APIKeys) != 1 {
+ t.Fatalf("snapshot API keys = %d, want 1", len(snapshot.APIKeys))
+ }
+ principal := snapshot.APIKeys[0].Principal
+ if principal.MonthlySpendMicros != 25_000_000 || principal.ExpiresAt == nil {
+ t.Fatalf("snapshot lost API key controls: %+v", principal)
+ }
+ if _, ok := principal.AllowedModels[model.PublicID]; !ok {
+ t.Fatalf("snapshot lost allowed model: %+v", principal.AllowedModels)
+ }
+ otherTenant, _, err := store.CreateTenant(ctx, CreateTenantInput{Slug: "other-key-control-test", Name: "Other Key Control Test"})
+ if err != nil {
+ t.Fatal(err)
+ }
+ otherProject, _, err := store.CreateProject(ctx, CreateProjectInput{TenantID: otherTenant.ID, Slug: "default", Name: "Other Default"})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if _, _, err := store.CreateAPIKey(ctx, CreateAPIKeyInput{TenantID: tenant.ID, ProjectID: otherProject.ID, Name: "cross tenant"}); err == nil {
+ t.Fatal("cross-tenant project binding should be rejected by the composite foreign key")
+ }
+ if _, _, err := store.CreateAPIKey(ctx, CreateAPIKeyInput{TenantID: tenant.ID, ProjectID: project.ID,
+ Name: "invalid model", AllowedModels: []string{"model/does-not-exist"}}); err == nil {
+ t.Fatal("unknown allowed model should reject the whole API key transaction")
+ }
+ keys, err = store.ListAPIKeysFor(ctx, tenant.ID)
+ if err != nil || len(keys) != 1 {
+ t.Fatalf("failed key transaction leaked a row: keys=%d err=%v", len(keys), err)
+ }
+}
+
func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) {
databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL")
if databaseURL == "" {
@@ -63,4 +186,48 @@ func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) {
if count != 2 {
t.Fatalf("migration history contains %d rows, want previous and current", count)
}
+
+ scopedDB, err := pgxpool.New(ctx, isolatedURL)
+ if err != nil {
+ t.Fatal(err)
+ }
+ defer scopedDB.Close()
+ var tenantID, 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 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)
+ }
+ if err := scopedDB.QueryRow(ctx, `INSERT INTO models (public_id,display_name) VALUES ('model/preference-test','Preference Model') RETURNING id::text`).Scan(&modelID); err != nil {
+ t.Fatal(err)
+ }
+ if _, err := scopedDB.Exec(ctx, `INSERT INTO model_price_versions (model_id,version,currency) VALUES ($1,1,'usd')`, modelID); err != nil {
+ t.Fatal(err)
+ }
+ if _, err := scopedDB.Exec(ctx, `INSERT INTO model_routes (model_id,provider_id,upstream_model) VALUES ($1,$2,'upstream-test')`, modelID, providerID); err != nil {
+ t.Fatal(err)
+ }
+ store := &Store{db: scopedDB}
+ prefs, err := store.SetDeveloperPreferences(ctx, SetDeveloperPreferencesInput{TenantID: tenantID, DefaultModel: "model/preference-test"})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if prefs.DefaultModel != "model/preference-test" || prefs.FallbackModel != "" {
+ t.Fatalf("unexpected developer preferences: %+v", prefs)
+ }
+ enabled := false
+ threshold := int64(9_750_000)
+ if _, err := store.SetBillingPreferences(ctx, SetBillingPreferencesInput{TenantID: tenantID,
+ LowBalanceEnabled: &enabled, LowBalanceThresholdMicros: &threshold}, 5_000_000); err != nil {
+ t.Fatal(err)
+ }
+ prefs, err = store.GetTenantPreferences(ctx, tenantID, 5_000_000)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if prefs.LowBalanceEnabled || prefs.LowBalanceThresholdMicros != threshold || 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 c24dde9..f6402d8 100644
--- a/internal/controlplane/types.go
+++ b/internal/controlplane/types.go
@@ -31,14 +31,22 @@ type Project struct {
}
type APIKey struct {
- ID string `json:"id"`
- TenantID string `json:"tenant_id"`
- ProjectID string `json:"project_id"`
- Name string `json:"name"`
- KeyPrefix string `json:"key_prefix"`
- Scopes []string `json:"scopes"`
- Status string `json:"status"`
- CreatedAt time.Time `json:"created_at"`
+ ID string `json:"id"`
+ TenantID string `json:"tenant_id"`
+ ProjectID string `json:"project_id"`
+ Name string `json:"name"`
+ KeyPrefix string `json:"key_prefix"`
+ Scopes []string `json:"scopes"`
+ Tags []string `json:"tags"`
+ AllowedModels []string `json:"allowed_models"`
+ MonthlySpendMicros int64 `json:"monthly_spend_micros"`
+ CurrentMonthSpendMicros int64 `json:"current_month_spend_micros"`
+ CurrentMonthReservedMicros int64 `json:"current_month_reserved_micros"`
+ CurrentMonthRequests int64 `json:"current_month_requests"`
+ Status string `json:"status"`
+ ExpiresAt *time.Time `json:"expires_at,omitempty"`
+ LastUsedAt *time.Time `json:"last_used_at,omitempty"`
+ CreatedAt time.Time `json:"created_at"`
}
type CreatedAPIKey struct {
@@ -48,8 +56,10 @@ type CreatedAPIKey struct {
type Provider struct {
ID string `json:"id"`
+ Slug string `json:"slug"`
Name string `json:"name"`
Protocol string `json:"protocol"`
+ WireAPI string `json:"wire_api"`
BaseURL string `json:"base_url"`
Enabled bool `json:"enabled"`
RouteCount int `json:"route_count"`
@@ -57,14 +67,112 @@ type Provider struct {
}
type Route struct {
- ID string `json:"id"`
- ProviderID string `json:"provider_id"`
- ProviderName string `json:"provider_name"`
- Protocol string `json:"protocol"`
- UpstreamModel string `json:"upstream_model"`
- Priority int `json:"priority"`
- Weight int `json:"weight"`
- Enabled bool `json:"enabled"`
+ ID string `json:"id"`
+ ProviderID string `json:"provider_id"`
+ ProviderName string `json:"provider_name"`
+ Protocol string `json:"protocol"`
+ WireAPI string `json:"wire_api"`
+ UpstreamModel string `json:"upstream_model"`
+ Priority int `json:"priority"`
+ Weight int `json:"weight"`
+ Enabled bool `json:"enabled"`
+ ProviderEnabled bool `json:"provider_enabled"`
+}
+
+// DeveloperModel is the customer-safe model catalog view. It intentionally
+// omits internal provider URLs, upstream model names, routing weights, and
+// allowlist membership.
+type DeveloperModel struct {
+ ID string `json:"-"`
+ PublicID string `json:"public_id"`
+ DisplayName string `json:"display_name"`
+ Description string `json:"description"`
+ OwnedBy string `json:"owned_by"`
+ InputModalities []string `json:"input_modalities"`
+ OutputModalities []string `json:"output_modalities"`
+ ContextWindow int64 `json:"context_window"`
+ MaxOutputTokens int64 `json:"max_output_tokens"`
+ Capabilities []string `json:"capabilities"`
+ Regions []string `json:"regions"`
+ Lifecycle string `json:"lifecycle"`
+ ReleasedAt *time.Time `json:"released_at,omitempty"`
+ ReplacementModel string `json:"replacement_model,omitempty"`
+ Aliases []string `json:"aliases"`
+ PriceCurrency string `json:"price_currency"`
+ 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"`
+ CacheWritePriceMicrosPerMillion int64 `json:"cache_write_price_micros_per_million"`
+ SupportedWireAPIs []string `json:"supported_wire_apis"`
+ ProviderCount int `json:"provider_count"`
+ AvailableProviderCount int `json:"available_provider_count"`
+ HealthStatus string `json:"health_status"`
+ Providers []DeveloperProviderHealth `json:"providers"`
+}
+
+type DeveloperProviderHealth struct {
+ Slug string `json:"slug"`
+ Name string `json:"name"`
+ Protocol string `json:"protocol"`
+ WireAPI string `json:"wire_api"`
+ State string `json:"state"`
+ Attempts uint64 `json:"attempts"`
+ RecentSamples int `json:"recent_samples"`
+ AvailabilityPercent float64 `json:"availability_percent"`
+ HeaderLatencyEWMA int64 `json:"header_latency_ewma_ms"`
+ ConsecutiveFailures uint64 `json:"consecutive_failures"`
+ LastObservedAt *time.Time `json:"last_observed_at,omitempty"`
+ CircuitOpenUntil *time.Time `json:"circuit_open_until,omitempty"`
+}
+
+// PublicModel is the unauthenticated catalog view. It is deliberately smaller
+// than DeveloperModel so authenticated-only fields cannot become public by
+// accident when the tenant catalog evolves.
+type PublicModel struct {
+ PublicID string `json:"public_id"`
+ DisplayName string `json:"display_name"`
+ Description string `json:"description"`
+ OwnedBy string `json:"owned_by"`
+ InputModalities []string `json:"input_modalities"`
+ OutputModalities []string `json:"output_modalities"`
+ ContextWindow int64 `json:"context_window"`
+ MaxOutputTokens int64 `json:"max_output_tokens"`
+ Capabilities []string `json:"capabilities"`
+ Regions []string `json:"regions"`
+ Lifecycle string `json:"lifecycle"`
+ ReleasedAt *time.Time `json:"released_at,omitempty"`
+ ReplacementModel string `json:"replacement_model,omitempty"`
+ Aliases []string `json:"aliases"`
+ PriceCurrency string `json:"price_currency"`
+ 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"`
+ CacheWritePriceMicrosPerMillion int64 `json:"cache_write_price_micros_per_million"`
+ SupportedWireAPIs []string `json:"supported_wire_apis"`
+ ProviderCount int `json:"provider_count"`
+ AvailableProviderCount int `json:"available_provider_count"`
+ HealthStatus string `json:"health_status"`
+}
+
+type TenantPreferences struct {
+ TenantID string `json:"tenant_id,omitempty"`
+ DefaultModel string `json:"default_model,omitempty"`
+ FallbackModel string `json:"fallback_model,omitempty"`
+ LowBalanceEnabled bool `json:"low_balance_enabled"`
+ LowBalanceThresholdMicros int64 `json:"low_balance_threshold_micros"`
+ UpdatedAt *time.Time `json:"updated_at,omitempty"`
+}
+
+type SetDeveloperPreferencesInput struct {
+ TenantID string `json:"tenant_id"`
+ DefaultModel string `json:"default_model"`
+ FallbackModel string `json:"fallback_model"`
+}
+
+type SetBillingPreferencesInput struct {
+ TenantID string `json:"tenant_id"`
+ LowBalanceEnabled *bool `json:"low_balance_enabled"`
+ LowBalanceThresholdMicros *int64 `json:"low_balance_threshold_micros"`
}
type Model struct {
@@ -145,15 +253,21 @@ type CreateProjectInput struct {
}
type CreateAPIKeyInput struct {
- TenantID string `json:"tenant_id"`
- ProjectID string `json:"project_id"`
- Name string `json:"name"`
- Scopes []string `json:"scopes"`
+ TenantID string `json:"tenant_id"`
+ ProjectID string `json:"project_id"`
+ Name string `json:"name"`
+ Scopes []string `json:"scopes"`
+ Tags []string `json:"tags"`
+ AllowedModels []string `json:"allowed_models"`
+ MonthlySpendMicros int64 `json:"monthly_spend_micros"`
+ ExpiresAt *time.Time `json:"expires_at"`
}
type CreateProviderInput struct {
+ Slug string `json:"slug"`
Name string `json:"name"`
Protocol string `json:"protocol"`
+ WireAPI string `json:"wire_api"`
BaseURL string `json:"base_url"`
APIKey string `json:"api_key"`
}
@@ -340,9 +454,12 @@ type UsageRecord struct {
RequestID string `json:"request_id"`
TenantID string `json:"tenant_id"`
ProjectID string `json:"project_id"`
+ ProjectName string `json:"project_name"`
KeyID string `json:"key_id"`
+ KeyName string `json:"key_name"`
PublicModel string `json:"public_model"`
ProviderID string `json:"provider_id,omitempty"`
+ ProviderName string `json:"provider_name,omitempty"`
UpstreamModel string `json:"upstream_model,omitempty"`
Protocol string `json:"protocol"`
Stream bool `json:"stream"`
@@ -360,6 +477,72 @@ type UsageRecord struct {
CostMicros int64 `json:"cost_micros"`
ChargedMicros int64 `json:"charged_micros"`
UncollectedMicros int64 `json:"uncollected_micros"`
+ UsageReported bool `json:"usage_reported"`
+ MeteringStatus string `json:"metering_status"`
+}
+
+type UsageDailyPoint struct {
+ Day time.Time `json:"day"`
+ RequestCount int64 `json:"request_count"`
+ SuccessfulRequests int64 `json:"successful_requests"`
+ 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"`
+ AverageDurationMS int64 `json:"average_duration_ms"`
+ P95DurationMS int64 `json:"p95_duration_ms"`
+}
+
+// UsageAnalytics is a persisted-ledger aggregation used by the developer
+// console. It deliberately contains no prompt or response content.
+type UsageAnalytics struct {
+ RangeStart time.Time `json:"range_start"`
+ RangeEnd time.Time `json:"range_end"`
+ Models []UsageModelAnalytics `json:"models"`
+ Providers []UsageProviderAnalytics `json:"providers"`
+}
+
+type UsageModelAnalytics struct {
+ PublicModel string `json:"public_model"`
+ RequestCount int64 `json:"request_count"`
+ SuccessfulRequests int64 `json:"successful_requests"`
+ ErrorCount int64 `json:"error_count"`
+ ProviderCount int64 `json:"provider_count"`
+ InputTokens int64 `json:"input_tokens"`
+ OutputTokens int64 `json:"output_tokens"`
+ TotalTokens int64 `json:"total_tokens"`
+ CacheReadInputTokens int64 `json:"cache_read_input_tokens"`
+ CacheCreationInputTokens int64 `json:"cache_creation_input_tokens"`
+ ChargedMicros int64 `json:"charged_micros"`
+ UncollectedMicros int64 `json:"uncollected_micros"`
+ MissingUsageRequests int64 `json:"missing_usage_requests"`
+ AverageDurationMS int64 `json:"average_duration_ms"`
+ P95DurationMS int64 `json:"p95_duration_ms"`
+ PreviousChargedMicros int64 `json:"previous_charged_micros"`
+ ChargeChangePercent *float64 `json:"charge_change_percent,omitempty"`
+}
+
+type UsageProviderAnalytics struct {
+ ProviderID string `json:"provider_id"`
+ ProviderName string `json:"provider_name"`
+ WireAPI string `json:"wire_api"`
+ 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"`
+ CacheReadInputTokens int64 `json:"cache_read_input_tokens"`
+ CacheCreationInputTokens int64 `json:"cache_creation_input_tokens"`
+ ChargedMicros int64 `json:"charged_micros"`
+ UncollectedMicros int64 `json:"uncollected_micros"`
+ MissingUsageRequests int64 `json:"missing_usage_requests"`
+ AverageDurationMS int64 `json:"average_duration_ms"`
+ P95DurationMS int64 `json:"p95_duration_ms"`
+ PreviousChargedMicros int64 `json:"previous_charged_micros"`
+ ChargeChangePercent *float64 `json:"charge_change_percent,omitempty"`
}
type UsageSummary struct {
diff --git a/internal/controlplane/usage.go b/internal/controlplane/usage.go
index b436d8c..ebabaf6 100644
--- a/internal/controlplane/usage.go
+++ b/internal/controlplane/usage.go
@@ -14,7 +14,16 @@ import (
type UsageQuery struct {
TenantID string
ProjectID string
+ KeyID string
Model string
+ Provider string
+ Protocol string
+ ErrorType string
+ Stream *bool
+ RequestID string
+ Status string
+ From time.Time
+ To time.Time
Limit int
}
@@ -39,6 +48,11 @@ func (s *Store) RecordUsage(ctx context.Context, event domain.UsageEvent) error
if err != nil {
return fmt.Errorf("persist usage event: %w", err)
}
+ if _, err := tx.Exec(ctx, `UPDATE api_keys
+ SET last_used_at = GREATEST(COALESCE(last_used_at, $2), $2)
+ WHERE id = $1`, event.KeyID, event.StartedAt); err != nil {
+ return fmt.Errorf("update API key last used time: %w", err)
+ }
if command.RowsAffected() == 0 {
return tx.Commit(ctx)
}
@@ -96,20 +110,55 @@ func (s *Store) ListUsage(ctx context.Context, query UsageQuery) ([]UsageRecord,
limit = 200
}
where := []string{"1=1"}
- args := make([]any, 0, 5)
+ args := make([]any, 0, 13)
index := 1
- for _, item := range []struct{ value, clause string }{{query.TenantID, "tenant_id=$"}, {query.ProjectID, "project_id=$"}, {query.Model, "public_model=$"}} {
+ 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) != "" {
where = append(where, item.clause+fmt.Sprint(index))
args = append(args, item.value)
index++
}
}
+ for _, item := range []struct{ value, clause string }{{query.Protocol, "protocol=$"}, {query.ErrorType, "error_type=$"}} {
+ if strings.TrimSpace(item.value) != "" {
+ where = append(where, item.clause+fmt.Sprint(index))
+ args = append(args, item.value)
+ index++
+ }
+ }
+ if strings.TrimSpace(query.Provider) != "" {
+ where = append(where, "EXISTS (SELECT 1 FROM providers filter_provider WHERE filter_provider.id::text=usage_events.provider_id AND filter_provider.slug=$"+fmt.Sprint(index)+")")
+ args = append(args, query.Provider)
+ index++
+ }
+ if query.Stream != nil {
+ where = append(where, "stream=$"+fmt.Sprint(index))
+ args = append(args, *query.Stream)
+ index++
+ }
+ if query.Status == "success" {
+ where = append(where, "success=TRUE")
+ } else if query.Status == "error" {
+ where = append(where, "success=FALSE")
+ }
+ if !query.From.IsZero() {
+ where = append(where, "started_at >= $"+fmt.Sprint(index))
+ args = append(args, query.From)
+ index++
+ }
+ if !query.To.IsZero() {
+ where = append(where, "started_at < $"+fmt.Sprint(index))
+ args = append(args, query.To)
+ index++
+ }
args = append(args, limit)
- rows, err := s.db.Query(ctx, `SELECT request_id, tenant_id::text, project_id::text, key_id::text, public_model,
- COALESCE(provider_id,''), COALESCE(upstream_model,''), protocol, stream, status_code, success, error_type,
+ 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,
- cache_read_input_tokens, cost_micros, charged_micros, uncollected_micros FROM usage_events WHERE `+
+ 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...)
if err != nil {
return nil, fmt.Errorf("query usage events: %w", err)
@@ -118,10 +167,10 @@ func (s *Store) ListUsage(ctx context.Context, query UsageQuery) ([]UsageRecord,
result := make([]UsageRecord, 0)
for rows.Next() {
var item UsageRecord
- if err := rows.Scan(&item.RequestID, &item.TenantID, &item.ProjectID, &item.KeyID, &item.PublicModel, &item.ProviderID, &item.UpstreamModel,
+ 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.CostMicros, &item.ChargedMicros, &item.UncollectedMicros); err != nil {
+ &item.CostMicros, &item.ChargedMicros, &item.UncollectedMicros, &item.UsageReported, &item.MeteringStatus); err != nil {
return nil, fmt.Errorf("scan usage event: %w", err)
}
result = append(result, item)
@@ -129,6 +178,71 @@ func (s *Store) ListUsage(ctx context.Context, query UsageQuery) ([]UsageRecord,
return result, rows.Err()
}
+func (s *Store) UsageDaily(ctx context.Context, query UsageQuery) ([]UsageDailyPoint, error) {
+ where := []string{"1=1"}
+ args := make([]any, 0, 12)
+ 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) != "" {
+ where = append(where, item.clause+fmt.Sprint(index))
+ args = append(args, item.value)
+ index++
+ }
+ }
+ for _, item := range []struct{ value, clause string }{{query.Protocol, "protocol=$"}, {query.ErrorType, "error_type=$"}} {
+ if strings.TrimSpace(item.value) != "" {
+ where = append(where, item.clause+fmt.Sprint(index))
+ args = append(args, item.value)
+ index++
+ }
+ }
+ if strings.TrimSpace(query.Provider) != "" {
+ where = append(where, "EXISTS (SELECT 1 FROM providers filter_provider WHERE filter_provider.id::text=usage_events.provider_id AND filter_provider.slug=$"+fmt.Sprint(index)+")")
+ args = append(args, query.Provider)
+ index++
+ }
+ if query.Stream != nil {
+ where = append(where, "stream=$"+fmt.Sprint(index))
+ args = append(args, *query.Stream)
+ index++
+ }
+ if query.Status == "success" {
+ where = append(where, "success=TRUE")
+ } else if query.Status == "error" {
+ where = append(where, "success=FALSE")
+ }
+ if !query.From.IsZero() {
+ where = append(where, "started_at >= $"+fmt.Sprint(index))
+ args = append(args, query.From)
+ index++
+ }
+ if !query.To.IsZero() {
+ where = append(where, "started_at < $"+fmt.Sprint(index))
+ args = append(args, query.To)
+ index++
+ }
+ rows, err := s.db.Query(ctx, `SELECT date_trunc('day', started_at AT TIME ZONE 'UTC'), count(*),
+ 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
+ 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)
+ }
+ defer rows.Close()
+ result := make([]UsageDailyPoint, 0)
+ 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 {
+ return nil, fmt.Errorf("scan daily usage: %w", err)
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
func (s *Store) UsageSummary(ctx context.Context, tenantID, projectID string) ([]UsageSummary, error) {
where := []string{"1=1"}
args := make([]any, 0, 2)
diff --git a/internal/controlplane/usage_analytics.go b/internal/controlplane/usage_analytics.go
new file mode 100644
index 0000000..7cc042f
--- /dev/null
+++ b/internal/controlplane/usage_analytics.go
@@ -0,0 +1,215 @@
+package controlplane
+
+import (
+ "context"
+ "fmt"
+ "strings"
+ "time"
+)
+
+// UsageAnalytics aggregates the immutable usage ledger for the developer
+// console. Queries run outside the inference path and are scoped by the
+// caller's tenant before reaching this store.
+func (s *Store) UsageAnalytics(ctx context.Context, query UsageQuery) (UsageAnalytics, error) {
+ to := query.To
+ if to.IsZero() {
+ to = time.Now().UTC()
+ }
+ from := query.From
+ if from.IsZero() {
+ from = to.Add(-30 * 24 * time.Hour)
+ }
+ 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)}
+
+ modelPrevious, err := s.usageModelCharges(ctx, query, from.Add(-to.Sub(from)), from)
+ if err != nil {
+ return UsageAnalytics{}, err
+ }
+ models, err := s.usageModelAnalytics(ctx, query, from, to)
+ if err != nil {
+ return UsageAnalytics{}, err
+ }
+ for index := range models {
+ models[index].PreviousChargedMicros = modelPrevious[models[index].PublicModel]
+ models[index].ChargeChangePercent = chargeChange(models[index].ChargedMicros, models[index].PreviousChargedMicros)
+ }
+
+ providerPrevious, err := s.usageProviderCharges(ctx, query, from.Add(-to.Sub(from)), from)
+ if err != nil {
+ return UsageAnalytics{}, err
+ }
+ providers, err := s.usageProviderAnalytics(ctx, query, from, to)
+ if err != nil {
+ return UsageAnalytics{}, err
+ }
+ for index := range providers {
+ providers[index].PreviousChargedMicros = providerPrevious[providers[index].ProviderID]
+ providers[index].ChargeChangePercent = chargeChange(providers[index].ChargedMicros, providers[index].PreviousChargedMicros)
+ }
+ result.Models = models
+ result.Providers = providers
+ return result, nil
+}
+
+func chargeChange(current, previous int64) *float64 {
+ if previous == 0 {
+ return nil
+ }
+ value := (float64(current) - float64(previous)) / float64(previous) * 100
+ return &value
+}
+
+func analyticsUsageWhere(query UsageQuery, from, to time.Time) (string, []any) {
+ where := []string{"1=1"}
+ args := make([]any, 0, 12)
+ index := 1
+ for _, item := range []struct {
+ value string
+ clause string
+ }{
+ {query.TenantID, "e.tenant_id=$"},
+ {query.ProjectID, "e.project_id=$"},
+ {query.KeyID, "e.key_id=$"},
+ {query.Model, "e.public_model=$"},
+ {query.RequestID, "e.request_id=$"},
+ } {
+ if strings.TrimSpace(item.value) != "" {
+ where = append(where, item.clause+fmt.Sprint(index))
+ args = append(args, item.value)
+ index++
+ }
+ }
+ for _, item := range []struct {
+ value string
+ clause string
+ }{
+ {query.Protocol, "e.protocol=$"},
+ {query.ErrorType, "e.error_type=$"},
+ } {
+ if strings.TrimSpace(item.value) != "" {
+ where = append(where, item.clause+fmt.Sprint(index))
+ args = append(args, item.value)
+ index++
+ }
+ }
+ if strings.TrimSpace(query.Provider) != "" {
+ where = append(where, "EXISTS (SELECT 1 FROM providers filter_provider WHERE filter_provider.id::text=e.provider_id AND filter_provider.slug=$"+fmt.Sprint(index)+")")
+ args = append(args, query.Provider)
+ index++
+ }
+ if query.Stream != nil {
+ where = append(where, "e.stream=$"+fmt.Sprint(index))
+ args = append(args, *query.Stream)
+ index++
+ }
+ if query.Status == "success" {
+ where = append(where, "e.success=TRUE")
+ } else if query.Status == "error" {
+ where = append(where, "e.success=FALSE")
+ }
+ if !from.IsZero() {
+ where = append(where, "e.started_at >= $"+fmt.Sprint(index))
+ args = append(args, from)
+ index++
+ }
+ if !to.IsZero() {
+ where = append(where, "e.started_at < $"+fmt.Sprint(index))
+ args = append(args, to)
+ }
+ return strings.Join(where, " AND "), args
+}
+
+func (s *Store) usageModelAnalytics(ctx context.Context, query UsageQuery, from, to time.Time) ([]UsageModelAnalytics, error) {
+ where, args := analyticsUsageWhere(query, from, to)
+ rows, err := s.db.Query(ctx, `SELECT e.public_model, count(*), count(*) FILTER (WHERE e.success), count(*) FILTER (WHERE NOT e.success),
+ 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
+ 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)
+ }
+ defer rows.Close()
+ result := make([]UsageModelAnalytics, 0)
+ for rows.Next() {
+ 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 {
+ return nil, fmt.Errorf("scan usage model analytics: %w", err)
+ }
+ result = append(result, item)
+ }
+ return result, rows.Err()
+}
+
+func (s *Store) usageModelCharges(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 e.public_model, COALESCE(sum(e.charged_micros),0)
+ FROM usage_events e WHERE `+where+` GROUP BY e.public_model`, args...)
+ if err != nil {
+ return nil, fmt.Errorf("query previous model charges: %w", err)
+ }
+ defer rows.Close()
+ result := make(map[string]int64)
+ for rows.Next() {
+ var model string
+ var charged int64
+ if err := rows.Scan(&model, &charged); err != nil {
+ return nil, fmt.Errorf("scan previous model charges: %w", err)
+ }
+ result[model] = charged
+ }
+ return result, rows.Err()
+}
+
+func (s *Store) usageProviderAnalytics(ctx context.Context, query UsageQuery, from, to time.Time) ([]UsageProviderAnalytics, error) {
+ where, args := analyticsUsageWhere(query, from, to)
+ rows, err := s.db.Query(ctx, `SELECT COALESCE(e.provider_id,''), COALESCE(NULLIF(p.name,''),'Unassigned'), COALESCE(p.wire_api,''),
+ 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.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
+ 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 {
+ return nil, fmt.Errorf("query usage provider analytics: %w", err)
+ }
+ defer rows.Close()
+ result := make([]UsageProviderAnalytics, 0)
+ for rows.Next() {
+ 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 {
+ return nil, fmt.Errorf("scan usage provider 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)
+ FROM usage_events e WHERE `+where+` GROUP BY e.provider_id`, args...)
+ if err != nil {
+ return nil, fmt.Errorf("query previous provider charges: %w", err)
+ }
+ defer rows.Close()
+ result := make(map[string]int64)
+ for rows.Next() {
+ var providerID string
+ var charged int64
+ if err := rows.Scan(&providerID, &charged); err != nil {
+ return nil, fmt.Errorf("scan previous provider charges: %w", err)
+ }
+ result[providerID] = charged
+ }
+ return result, rows.Err()
+}
diff --git a/internal/controlplane/usage_analytics_test.go b/internal/controlplane/usage_analytics_test.go
new file mode 100644
index 0000000..2372be6
--- /dev/null
+++ b/internal/controlplane/usage_analytics_test.go
@@ -0,0 +1,34 @@
+package controlplane
+
+import "testing"
+
+func TestChargeChange(t *testing.T) {
+ tests := []struct {
+ name string
+ current, previous int64
+ want *float64
+ }{
+ {name: "no activity", current: 0, previous: 0, want: nil},
+ {name: "new spend has no finite percentage", current: 25, previous: 0, want: nil},
+ {name: "increase", current: 125, previous: 25, want: floatPointer(400)},
+ {name: "decrease", current: 25, previous: 100, want: floatPointer(-75)},
+ }
+ for _, test := range tests {
+ t.Run(test.name, func(t *testing.T) {
+ got := chargeChange(test.current, test.previous)
+ if test.want == nil {
+ if got != nil {
+ t.Fatalf("chargeChange() = %v, want nil", *got)
+ }
+ return
+ }
+ if got == nil || *got != *test.want {
+ t.Fatalf("chargeChange() = %v, want %v", got, *test.want)
+ }
+ })
+ }
+}
+
+func floatPointer(value float64) *float64 {
+ return &value
+}
diff --git a/internal/controlplane/usage_integration_test.go b/internal/controlplane/usage_integration_test.go
new file mode 100644
index 0000000..2545c1d
--- /dev/null
+++ b/internal/controlplane/usage_integration_test.go
@@ -0,0 +1,152 @@
+package controlplane
+
+import (
+ "context"
+ "crypto/sha256"
+ "fmt"
+ "os"
+ "testing"
+ "time"
+
+ "aigw/internal/domain"
+
+ "github.com/jackc/pgx/v5/pgxpool"
+)
+
+func TestUsageFiltersAndDailyAggregationPostgres(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 := MigrateDatabase(ctx, databaseURL); err != nil {
+ t.Fatal(err)
+ }
+ db, err := pgxpool.New(ctx, databaseURL)
+ if err != nil {
+ t.Fatal(err)
+ }
+ t.Cleanup(db.Close)
+ store := &Store{db: db}
+ suffix := time.Now().UnixNano()
+ var tenantID, projectID, keyID, providerID, otherTenantID, otherProjectID, otherKeyID string
+ if err := db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Usage integration') RETURNING id::text`, fmt.Sprintf("usage-%d", suffix)).Scan(&tenantID); err != nil {
+ t.Fatal(err)
+ }
+ if err := db.QueryRow(ctx, `INSERT INTO projects (tenant_id,slug,name) VALUES ($1,'production','Production') RETURNING id::text`, tenantID).Scan(&projectID); err != nil {
+ t.Fatal(err)
+ }
+ keyHash := sha256.Sum256([]byte(fmt.Sprint(suffix)))
+ if err := db.QueryRow(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_hash) VALUES ($1,$2,'Production key','sk-usage',$3) RETURNING id::text`, tenantID, projectID, keyHash[:]).Scan(&keyID); err != nil {
+ t.Fatal(err)
+ }
+ if err := db.QueryRow(ctx, `INSERT INTO providers (slug,name,protocol,wire_api,base_url,api_key_ciphertext)
+ VALUES ($1,$2,'openai','responses','https://usage.test',$3) RETURNING id::text`, fmt.Sprintf("usage-provider-%d", suffix), fmt.Sprintf("Usage provider %d", suffix), []byte("encrypted-test-value")).Scan(&providerID); err != nil {
+ t.Fatal(err)
+ }
+ if err := db.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ($1,'Other usage tenant') RETURNING id::text`, fmt.Sprintf("usage-other-%d", suffix)).Scan(&otherTenantID); err != nil {
+ t.Fatal(err)
+ }
+ if err := db.QueryRow(ctx, `INSERT INTO projects (tenant_id,slug,name) VALUES ($1,'production','Other production') RETURNING id::text`, otherTenantID).Scan(&otherProjectID); err != nil {
+ t.Fatal(err)
+ }
+ otherKeyHash := sha256.Sum256([]byte(fmt.Sprintf("other-%d", suffix)))
+ if err := db.QueryRow(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_hash) VALUES ($1,$2,'Other key','sk-other',$3) RETURNING id::text`, otherTenantID, otherProjectID, otherKeyHash[:]).Scan(&otherKeyID); err != nil {
+ t.Fatal(err)
+ }
+ t.Cleanup(func() {
+ cleanupCtx := context.Background()
+ for _, statement := range []struct{ query, arg string }{
+ {`DELETE FROM usage_monthly_rollups WHERE tenant_id=$1`, tenantID},
+ {`DELETE FROM usage_monthly_rollups WHERE tenant_id=$1`, otherTenantID},
+ {`DELETE FROM usage_events WHERE tenant_id=$1`, tenantID},
+ {`DELETE FROM usage_events WHERE tenant_id=$1`, otherTenantID},
+ {`DELETE FROM providers WHERE id=$1`, providerID},
+ {`DELETE FROM api_keys WHERE id=$1`, keyID},
+ {`DELETE FROM api_keys WHERE id=$1`, otherKeyID},
+ {`DELETE FROM projects WHERE id=$1`, projectID},
+ {`DELETE FROM projects WHERE id=$1`, otherProjectID},
+ {`DELETE FROM tenants WHERE id=$1`, tenantID},
+ {`DELETE FROM tenants WHERE id=$1`, otherTenantID},
+ } {
+ if _, cleanupErr := db.Exec(cleanupCtx, statement.query, statement.arg); cleanupErr != nil {
+ t.Errorf("cleanup usage integration data: %v", cleanupErr)
+ }
+ }
+ })
+
+ 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_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 {
+ if err := store.RecordUsage(ctx, event); err != nil {
+ t.Fatal(err)
+ }
+ }
+ if _, err := db.Exec(ctx, `UPDATE usage_events SET charged_micros=125,cost_micros=125 WHERE request_id=$1`, events[0].RequestID); err != nil {
+ t.Fatal(err)
+ }
+
+ records, err := store.ListUsage(ctx, UsageQuery{TenantID: tenantID, KeyID: keyID, Status: "success", From: started.Add(-time.Minute), To: started.Add(time.Hour), Limit: 20})
+ 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" {
+ t.Fatalf("unexpected filtered usage: %+v", records)
+ }
+ streaming := true
+ providerSlug := fmt.Sprintf("usage-provider-%d", suffix)
+ records, err = store.ListUsage(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), Limit: 20})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(records) != 1 || records[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)})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(filteredPoints) != 1 || filteredPoints[0].RequestCount != 1 || filteredPoints[0].SuccessfulRequests != 0 {
+ t.Fatalf("unexpected filtered daily usage: %+v", filteredPoints)
+ }
+ points, err := store.UsageDaily(ctx, UsageQuery{TenantID: tenantID, From: started.Add(-time.Hour), To: started.Add(2 * time.Hour)})
+ 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 {
+ 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}}
+ 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} {
+ if err := store.RecordUsage(ctx, event); err != nil {
+ t.Fatal(err)
+ }
+ }
+ 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)
+ }
+ analytics, err := store.UsageAnalytics(ctx, UsageQuery{TenantID: tenantID, From: started.Add(-time.Minute), To: started.Add(10 * time.Minute)})
+ 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 {
+ t.Fatalf("unexpected model analytics: %+v", analytics.Models)
+ }
+ 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)
+ }
+ filteredAnalytics, err := store.UsageAnalytics(ctx, UsageQuery{TenantID: tenantID, Provider: providerSlug, Protocol: string(domain.ProtocolOpenAIResponses), ErrorType: "provider_error", Stream: &streaming, From: started.Add(-time.Minute), To: started.Add(10 * time.Minute)})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(filteredAnalytics.Models) != 1 || filteredAnalytics.Models[0].RequestCount != 1 || filteredAnalytics.Models[0].ErrorCount != 1 || len(filteredAnalytics.Providers) != 1 || filteredAnalytics.Providers[0].RequestCount != 1 {
+ t.Fatalf("unexpected filtered analytics: %+v", filteredAnalytics)
+ }
+}