From 41e322c53d7b4b796eb377d0df9c29ecd10ba431 Mon Sep 17 00:00:00 2001 From: Chia Date: Thu, 6 Aug 2026 09:29:41 +1200 Subject: 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 --- internal/controlplane/access.go | 11 +- internal/controlplane/access_test.go | 4 + internal/controlplane/mail_operations.go | 13 +- internal/controlplane/mutations.go | 82 +++++++-- internal/controlplane/preferences.go | 156 +++++++++++++++++ internal/controlplane/preferences_test.go | 31 ++++ internal/controlplane/queries.go | 191 ++++++++++++++++++-- internal/controlplane/queries_test.go | 80 +++++++++ internal/controlplane/schema.sql | 110 ++++++++++++ internal/controlplane/snapshot.go | 35 +++- internal/controlplane/store.go | 4 +- internal/controlplane/store_integration_test.go | 167 ++++++++++++++++++ internal/controlplane/types.go | 223 +++++++++++++++++++++--- internal/controlplane/usage.go | 128 +++++++++++++- internal/controlplane/usage_analytics.go | 215 +++++++++++++++++++++++ internal/controlplane/usage_analytics_test.go | 34 ++++ internal/controlplane/usage_integration_test.go | 152 ++++++++++++++++ 17 files changed, 1564 insertions(+), 72 deletions(-) create mode 100644 internal/controlplane/preferences.go create mode 100644 internal/controlplane/preferences_test.go create mode 100644 internal/controlplane/queries_test.go create mode 100644 internal/controlplane/usage_analytics.go create mode 100644 internal/controlplane/usage_analytics_test.go create mode 100644 internal/controlplane/usage_integration_test.go (limited to 'internal/controlplane') 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 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) + } +} -- cgit v1.2.3