diff options
| author | Chia <Chia@93.nz> | 2026-08-06 15:58:57 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-06 15:58:57 +1200 |
| commit | 3f702084d20b3c3a3ea916f3110e99b22bda60b3 (patch) | |
| tree | 517f76c51025ce1ee085ea4898c60f799e5c37ea /internal/controlplane | |
| parent | 41e322c53d7b4b796eb377d0df9c29ecd10ba431 (diff) | |
feat: complete commercial developer workflowspublish-commercial-control-plane
Add tenant-safe usage observability, prepaid billing controls, API key lifecycle management, Embeddings metering, configurable billing alerts, and resilient provider health propagation. Harden Stripe failure handling, migrations, readiness, and the authenticated control-plane UI with end-to-end verification evidence.
Diffstat (limited to 'internal/controlplane')
| -rw-r--r-- | internal/controlplane/api_key_test.go | 26 | ||||
| -rw-r--r-- | internal/controlplane/mail_operations.go | 16 | ||||
| -rw-r--r-- | internal/controlplane/mail_operations_integration_test.go | 144 | ||||
| -rw-r--r-- | internal/controlplane/manager.go | 38 | ||||
| -rw-r--r-- | internal/controlplane/manager_test.go | 49 | ||||
| -rw-r--r-- | internal/controlplane/mutations.go | 130 | ||||
| -rw-r--r-- | internal/controlplane/preferences.go | 75 | ||||
| -rw-r--r-- | internal/controlplane/preferences_test.go | 18 | ||||
| -rw-r--r-- | internal/controlplane/queries.go | 35 | ||||
| -rw-r--r-- | internal/controlplane/schema.sql | 40 | ||||
| -rw-r--r-- | internal/controlplane/snapshot.go | 9 | ||||
| -rw-r--r-- | internal/controlplane/store.go | 48 | ||||
| -rw-r--r-- | internal/controlplane/store_integration_test.go | 84 | ||||
| -rw-r--r-- | internal/controlplane/types.go | 64 | ||||
| -rw-r--r-- | internal/controlplane/usage.go | 82 | ||||
| -rw-r--r-- | internal/controlplane/usage_analytics.go | 61 | ||||
| -rw-r--r-- | internal/controlplane/usage_integration_test.go | 43 | ||||
| -rw-r--r-- | internal/controlplane/usage_test.go | 26 |
18 files changed, 869 insertions, 119 deletions
diff --git a/internal/controlplane/api_key_test.go b/internal/controlplane/api_key_test.go new file mode 100644 index 0000000..51a3a7d --- /dev/null +++ b/internal/controlplane/api_key_test.go @@ -0,0 +1,26 @@ +package controlplane + +import ( + "crypto/sha256" + "strings" + "testing" +) + +func TestGenerateAPIKeySecretReturnsOnlyDisplayFragments(t *testing.T) { + raw, prefix, suffix, hash, err := generateAPIKeySecret() + if err != nil { + t.Fatal(err) + } + if !strings.HasPrefix(raw, "sk-aigw-") || !strings.HasPrefix(raw, strings.TrimSuffix(prefix, "...")) { + t.Fatalf("prefix %q does not identify the generated key", prefix) + } + if len(suffix) != 6 || !strings.HasSuffix(raw, suffix) { + t.Fatalf("suffix %q does not identify the generated key", suffix) + } + if len(prefix)+len(suffix) >= len(raw) { + t.Fatal("display fragments reveal the complete key") + } + if hash != sha256.Sum256([]byte(raw)) { + t.Fatal("generated digest does not authenticate the raw key") + } +} diff --git a/internal/controlplane/mail_operations.go b/internal/controlplane/mail_operations.go index d42a328..fc7b81b 100644 --- a/internal/controlplane/mail_operations.go +++ b/internal/controlplane/mail_operations.go @@ -149,21 +149,25 @@ func (s *Store) queueBillingNotifications(ctx context.Context, config MailNotifi FROM billing_ledger GROUP BY tenant_id) SELECT w.tenant_id::text,w.currency,w.balance_micros-w.reserved_micros,u.email,u.display_name, COALESCE(spend.today,0),COALESCE(spend.baseline,0), - COALESCE(pref.low_balance_enabled,TRUE),COALESCE(pref.low_balance_threshold_micros,$1) + COALESCE(pref.low_balance_enabled,TRUE),COALESCE(pref.low_balance_threshold_micros,$1), + COALESCE(pref.spend_anomaly_enabled,TRUE),COALESCE(pref.spend_anomaly_multiplier,$2), + COALESCE(pref.spend_anomaly_min_micros,$3) FROM tenant_wallets w JOIN console_users u ON u.tenant_id=w.tenant_id LEFT JOIN spend ON spend.tenant_id=w.tenant_id LEFT JOIN tenant_preferences pref ON pref.tenant_id=w.tenant_id WHERE u.status='active' AND u.email_verified_at IS NOT NULL AND u.role IN ('tenant_admin','tenant_billing') - AND NOT EXISTS (SELECT 1 FROM mail_suppressions s WHERE lower(s.recipient)=lower(u.email))`, config.LowBalanceMicros) + AND NOT EXISTS (SELECT 1 FROM mail_suppressions s WHERE lower(s.recipient)=lower(u.email))`, + config.LowBalanceMicros, config.SpendAnomalyMultiplier, config.SpendAnomalyMinMicros) if err != nil { return err } defer rows.Close() for rows.Next() { var tenantID, currency, email, name string - var available, today, baseline, lowBalanceThreshold int64 - var lowBalanceEnabled bool - if err := rows.Scan(&tenantID, ¤cy, &available, &email, &name, &today, &baseline, &lowBalanceEnabled, &lowBalanceThreshold); err != nil { + var available, today, baseline, lowBalanceThreshold, anomalyMultiplier, anomalyMinimum int64 + var lowBalanceEnabled, anomalyEnabled bool + if err := rows.Scan(&tenantID, ¤cy, &available, &email, &name, &today, &baseline, + &lowBalanceEnabled, &lowBalanceThreshold, &anomalyEnabled, &anomalyMultiplier, &anomalyMinimum); err != nil { return err } day := time.Now().UTC().Format("2006-01-02") @@ -173,7 +177,7 @@ func (s *Store) queueBillingNotifications(ctx context.Context, config MailNotifi return err } } - if baseline > 0 && today >= config.SpendAnomalyMinMicros && today >= baseline*config.SpendAnomalyMultiplier { + if anomalyEnabled && baseline > 0 && today >= anomalyMinimum && today >= baseline*anomalyMultiplier { body := fmt.Sprintf("Hi %s,\n\nAIGW detected unusual API spend today: %.6f %s versus a seven-day daily baseline of %.6f %s. Review API keys and usage in the console.\n", displayName(name), float64(today)/1_000_000, strings.ToUpper(currency), float64(baseline)/1_000_000, strings.ToUpper(currency)) if err := s.queueNotification(ctx, tenantID, email, "spend_anomaly", day, "Unusual AIGW API spend detected", body); err != nil { return err diff --git a/internal/controlplane/mail_operations_integration_test.go b/internal/controlplane/mail_operations_integration_test.go new file mode 100644 index 0000000..912ab7f --- /dev/null +++ b/internal/controlplane/mail_operations_integration_test.go @@ -0,0 +1,144 @@ +package controlplane + +import ( + "context" + "encoding/base64" + "fmt" + "net/url" + "os" + "strings" + "testing" + "time" + + "github.com/jackc/pgx/v5/pgxpool" +) + +func TestBillingNotificationsUseTenantPreferencesLedgerAndEncryptedOutboxPostgres(t *testing.T) { + databaseURL := os.Getenv("AIGW_TEST_DATABASE_URL") + if databaseURL == "" { + t.Skip("AIGW_TEST_DATABASE_URL is not set") + } + ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cancel() + rootDB, err := pgxpool.New(ctx, databaseURL) + if err != nil { + t.Fatal(err) + } + defer rootDB.Close() + schema := fmt.Sprintf("mail_notifications_%d", time.Now().UnixNano()) + if _, err := rootDB.Exec(ctx, "CREATE SCHEMA "+schema); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _, _ = rootDB.Exec(context.Background(), "DROP SCHEMA "+schema+" CASCADE") }) + parsed, err := url.Parse(databaseURL) + if err != nil { + t.Fatal(err) + } + query := parsed.Query() + query.Set("search_path", schema) + parsed.RawQuery = query.Encode() + isolatedURL := parsed.String() + if err := MigrateDatabase(ctx, isolatedURL); err != nil { + t.Fatal(err) + } + credentialKey := base64.StdEncoding.EncodeToString([]byte("01234567890123456789012345678901")) + store, err := NewStore(ctx, Options{DatabaseURL: isolatedURL, CredentialKey: credentialKey}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = store.Close() }) + + tenant, _, err := store.CreateTenant(ctx, CreateTenantInput{Slug: "mail-alert-test", Name: "Mail Alert Test"}) + if err != nil { + t.Fatal(err) + } + if _, err := store.db.Exec(ctx, `INSERT INTO tenant_wallets (tenant_id,currency,balance_micros) + VALUES ($1,'usd',1000000)`, tenant.ID); err != nil { + t.Fatal(err) + } + enabled := true + threshold := int64(2_000_000) + anomalyMultiplier := int64(4) + anomalyMinimum := int64(400_000) + if _, err := store.SetBillingPreferences(ctx, SetBillingPreferencesInput{TenantID: tenant.ID, + LowBalanceEnabled: &enabled, LowBalanceThresholdMicros: &threshold, SpendAnomalyEnabled: &enabled, + SpendAnomalyMultiplier: &anomalyMultiplier, SpendAnomalyMinMicros: &anomalyMinimum}, + BillingPreferenceDefaults{LowBalanceThresholdMicros: 5_000_000, SpendAnomalyMultiplier: 10, SpendAnomalyMinMicros: 900_000}); err != nil { + t.Fatal(err) + } + if _, err := store.db.Exec(ctx, `INSERT INTO console_users + (tenant_id,email,display_name,role,status,email_verified_at) VALUES + ($1,'billing-alert@example.test','Billing Owner','tenant_billing','active',now()), + ($1,'developer-no-alert@example.test','Developer','tenant_developer','active',now()), + ($1,'unverified-no-alert@example.test','Unverified Billing','tenant_billing','active',NULL)`, tenant.ID); err != nil { + t.Fatal(err) + } + if _, err := store.db.Exec(ctx, `INSERT INTO billing_ledger + (tenant_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description,created_at) + SELECT $1,'usd',-100000,1000000,'usage','request','historical-'||day::text,'Historical usage', + date_trunc('day',now())-make_interval(days=>day) + FROM generate_series(1,7) AS day`, tenant.ID); err != nil { + t.Fatal(err) + } + if _, err := store.db.Exec(ctx, `INSERT INTO billing_ledger + (tenant_id,currency,amount_micros,balance_after_micros,kind,source_type,source_id,description,created_at) + VALUES ($1,'usd',-500000,1000000,'usage','request','today-spend','Today usage',now())`, tenant.ID); err != nil { + t.Fatal(err) + } + + config := MailNotificationConfig{LowBalanceMicros: 5_000_000, SpendAnomalyMultiplier: 10, SpendAnomalyMinMicros: 900_000} + if err := store.queueBillingNotifications(ctx, config); err != nil { + t.Fatal(err) + } + if err := store.queueBillingNotifications(ctx, config); err != nil { + t.Fatal(err) + } + var billingMessages, otherMessages, notificationEvents int + if err := store.db.QueryRow(ctx, `SELECT + count(*) FILTER (WHERE recipient='billing-alert@example.test'), + count(*) FILTER (WHERE recipient<>'billing-alert@example.test') + FROM console_mail_outbox`).Scan(&billingMessages, &otherMessages); err != nil { + t.Fatal(err) + } + if err := store.db.QueryRow(ctx, `SELECT count(*) FROM mail_notification_events`).Scan(¬ificationEvents); err != nil { + t.Fatal(err) + } + if billingMessages != 2 || otherMessages != 0 || notificationEvents != 2 { + t.Fatalf("notification dedupe or recipient filtering failed: billing=%d other=%d events=%d", billingMessages, otherMessages, notificationEvents) + } + var plaintextLeaks int + if err := store.db.QueryRow(ctx, `SELECT count(*) FROM console_mail_outbox + WHERE convert_from(body_ciphertext,'UTF8') LIKE '%1.000000 USD%'`).Scan(&plaintextLeaks); err == nil { + if plaintextLeaks != 0 { + t.Fatal("notification body was stored as plaintext") + } + } else { + // Authenticated encryption output is arbitrary bytes and usually is not valid UTF-8. + var containsPlaintext bool + if scanErr := store.db.QueryRow(ctx, `SELECT bool_or(position(convert_to('1.000000 USD','UTF8') in body_ciphertext)>0) + FROM console_mail_outbox`).Scan(&containsPlaintext); scanErr != nil { + t.Fatal(scanErr) + } + if containsPlaintext { + t.Fatal("notification body was stored as plaintext") + } + } + + bodies := make([]string, 0, 2) + for range 2 { + message, ok, err := store.ClaimMail(ctx) + if err != nil { + t.Fatal(err) + } + if !ok || message.Recipient != "billing-alert@example.test" { + t.Fatalf("unexpected claimed notification: ok=%v message=%+v", ok, message) + } + bodies = append(bodies, message.Body) + } + joined := strings.Join(bodies, "\n") + for _, expected := range []string{"1.000000 USD", "0.500000 USD", "0.100000 USD"} { + if !strings.Contains(joined, expected) { + t.Fatalf("decrypted notifications do not contain %q: %s", expected, joined) + } + } +} diff --git a/internal/controlplane/manager.go b/internal/controlplane/manager.go index 212963b..58c31fe 100644 --- a/internal/controlplane/manager.go +++ b/internal/controlplane/manager.go @@ -15,11 +15,14 @@ import ( const broadcastQueueSize = 128 +const redisHealthCheckTimeout = 500 * time.Millisecond + type managerStore interface { LoadSnapshot(context.Context) (Snapshot, error) DatabaseGeneration(context.Context) (int64, error) PublishChange(context.Context, ChangeEvent) error Subscribe(context.Context) (<-chan ChangeMessage, func() error, error) + PingRedis(context.Context) error RedisEnabled() bool } @@ -129,7 +132,7 @@ func (m *Manager) RedisConnected() bool { func (m *Manager) Run(ctx context.Context) { var workers sync.WaitGroup if m.store.RedisEnabled() { - workers.Add(2) + workers.Add(3) go func() { defer workers.Done() m.runSubscriptions(ctx) @@ -138,6 +141,10 @@ func (m *Manager) Run(ctx context.Context) { defer workers.Done() m.runBroadcasts(ctx) }() + go func() { + defer workers.Done() + m.runRedisHealth(ctx) + }() } else { m.logger.Info("control_plane_redis_disabled", "fallback", "postgres_polling") } @@ -145,6 +152,35 @@ func (m *Manager) Run(ctx context.Context) { workers.Wait() } +func (m *Manager) runRedisHealth(ctx context.Context) { + interval := m.pollInterval + if interval > time.Second { + interval = time.Second + } + if interval <= 0 { + interval = time.Second + } + ticker := time.NewTicker(interval) + defer ticker.Stop() + for { + pingContext, cancel := context.WithTimeout(ctx, redisHealthCheckTimeout) + err := m.store.PingRedis(pingContext) + cancel() + if err != nil { + if m.redisConnected.Swap(false) { + m.logger.Warn("control_plane_redis_unavailable", "error", err, "fallback", "postgres_polling") + } + } else if !m.redisConnected.Swap(true) { + m.logger.Info("control_plane_redis_recovered") + } + select { + case <-ctx.Done(): + return + case <-ticker.C: + } + } +} + func (m *Manager) runPolling(ctx context.Context) { ticker := time.NewTicker(m.pollInterval) defer ticker.Stop() diff --git a/internal/controlplane/manager_test.go b/internal/controlplane/manager_test.go index b292060..78a0fdb 100644 --- a/internal/controlplane/manager_test.go +++ b/internal/controlplane/manager_test.go @@ -25,6 +25,7 @@ type fakeManagerStore struct { publishErr error published chan ChangeEvent subscribe func(context.Context, int64) (<-chan ChangeMessage, func() error, error) + pingRedis func(context.Context) error } type capturePolicies struct{ values []domain.LimitPolicy } @@ -84,6 +85,13 @@ func (s *fakeManagerStore) RedisEnabled() bool { return s.redisEnabled } +func (s *fakeManagerStore) PingRedis(ctx context.Context) error { + if s.pingRedis != nil { + return s.pingRedis(ctx) + } + return nil +} + func newTestManager(store managerStore, logger *slog.Logger, interval time.Duration) *Manager { return NewManager(store, catalog.NewModels(nil), auth.NewDynamic(nil, false), logger, interval) } @@ -196,6 +204,14 @@ func TestSubscriptionMessageRestoresConnectedStateAfterPublishFailure(t *testing store := newFakeManagerStore(1) store.redisEnabled = true store.publishErr = errors.New("redis unavailable") + var redisAvailable atomic.Bool + redisAvailable.Store(true) + store.pingRedis = func(context.Context) error { + if !redisAvailable.Load() { + return errors.New("redis unavailable") + } + return nil + } messages := make(chan ChangeMessage, 1) store.subscribe = func(_ context.Context, _ int64) (<-chan ChangeMessage, func() error, error) { return messages, func() error { return nil }, nil @@ -212,12 +228,14 @@ func TestSubscriptionMessageRestoresConnectedStateAfterPublishFailure(t *testing close(done) }() waitUntil(t, time.Second, manager.RedisConnected) + redisAvailable.Store(false) if err := manager.AfterMutation(context.Background(), 1, "model", "model-1"); err != nil { t.Fatal(err) } waitUntil(t, time.Second, func() bool { return !manager.RedisConnected() }) store.snapshot.Store(Snapshot{Generation: 2}) + redisAvailable.Store(true) messages <- ChangeMessage{Payload: `{"generation":2,"resource":"model"}`} waitUntil(t, time.Second, func() bool { return manager.RedisConnected() && manager.Generation() == 2 }) @@ -229,6 +247,37 @@ func TestSubscriptionMessageRestoresConnectedStateAfterPublishFailure(t *testing } } +func TestRedisHealthCheckReportsFailureAndRecovery(t *testing.T) { + store := newFakeManagerStore(1) + store.redisEnabled = true + var redisAvailable atomic.Bool + redisAvailable.Store(true) + store.pingRedis = func(context.Context) error { + if !redisAvailable.Load() { + return errors.New("redis unavailable") + } + return nil + } + manager := newTestManager(store, slog.New(slog.NewTextHandler(&safeLogBuffer{}, nil)), 10*time.Millisecond) + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { + manager.Run(ctx) + close(done) + }() + waitUntil(t, time.Second, manager.RedisConnected) + redisAvailable.Store(false) + waitUntil(t, time.Second, func() bool { return !manager.RedisConnected() }) + redisAvailable.Store(true) + waitUntil(t, time.Second, manager.RedisConnected) + cancel() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("manager did not stop") + } +} + func TestRedisCanBeDisabled(t *testing.T) { store := newFakeManagerStore(3) manager := newTestManager(store, slog.New(slog.NewTextHandler(&bytes.Buffer{}, nil)), 10*time.Millisecond) diff --git a/internal/controlplane/mutations.go b/internal/controlplane/mutations.go index 9f81d6e..aba4d36 100644 --- a/internal/controlplane/mutations.go +++ b/internal/controlplane/mutations.go @@ -85,8 +85,8 @@ func (s *Store) CreateAPIKey(ctx context.Context, input CreateAPIKeyInput) (Crea if input.TenantID == "" || input.ProjectID == "" || input.Name == "" { return CreatedAPIKey{}, 0, errors.New("API key requires tenant_id, project_id, and name") } - if len(input.Name) > 120 || input.MonthlySpendMicros < 0 { - return CreatedAPIKey{}, 0, errors.New("API key name or monthly spend limit is invalid") + if len(input.Name) > 120 || input.MonthlySpendMicros < 0 || input.DailySpendMicros < 0 || input.RequestsPerMinute < 0 || input.TokensPerMinute < 0 { + return CreatedAPIKey{}, 0, errors.New("API key name, spend limits, or rate limits are invalid") } if input.ExpiresAt != nil && !input.ExpiresAt.After(time.Now()) { return CreatedAPIKey{}, 0, errors.New("API key expiry must be in the future") @@ -107,13 +107,10 @@ func (s *Store) CreateAPIKey(ctx context.Context, input CreateAPIKeyInput) (Crea } scopesJSON, _ := json.Marshal(scopes) tagsJSON, _ := json.Marshal(tags) - random := make([]byte, 32) - if _, err := rand.Read(random); err != nil { - return CreatedAPIKey{}, 0, fmt.Errorf("generate API key: %w", err) + rawKey, prefix, suffix, hash, err := generateAPIKeySecret() + if err != nil { + return CreatedAPIKey{}, 0, err } - rawKey := "sk-aigw-" + base64.RawURLEncoding.EncodeToString(random) - hash := sha256.Sum256([]byte(rawKey)) - prefix := rawKey[:min(18, len(rawKey))] + "..." tx, err := s.db.Begin(ctx) if err != nil { @@ -122,14 +119,17 @@ func (s *Store) CreateAPIKey(ctx context.Context, input CreateAPIKeyInput) (Crea defer tx.Rollback(ctx) var result CreatedAPIKey err = tx.QueryRow(ctx, ` - INSERT INTO api_keys (tenant_id, project_id, name, key_prefix, key_hash, scopes, tags, monthly_spend_micros, expires_at) - VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9) - RETURNING id::text, tenant_id::text, project_id::text, name, key_prefix, scopes, tags, - monthly_spend_micros, status, expires_at, last_used_at, created_at`, - input.TenantID, input.ProjectID, input.Name, prefix, hash[:], scopesJSON, tagsJSON, - input.MonthlySpendMicros, input.ExpiresAt, - ).Scan(&result.ID, &result.TenantID, &result.ProjectID, &result.Name, &result.KeyPrefix, - &scopesJSON, &tagsJSON, &result.MonthlySpendMicros, &result.Status, &result.ExpiresAt, + INSERT INTO api_keys (tenant_id, project_id, name, key_prefix, key_suffix, key_hash, scopes, tags, monthly_spend_micros, + daily_spend_micros, requests_per_minute, tokens_per_minute, expires_at) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13) + RETURNING id::text, tenant_id::text, project_id::text, name, key_prefix, key_suffix, scopes, tags, + monthly_spend_micros, daily_spend_micros, requests_per_minute, tokens_per_minute, + status, expires_at, last_used_at, created_at`, + input.TenantID, input.ProjectID, input.Name, prefix, suffix, hash[:], scopesJSON, tagsJSON, + input.MonthlySpendMicros, input.DailySpendMicros, input.RequestsPerMinute, input.TokensPerMinute, input.ExpiresAt, + ).Scan(&result.ID, &result.TenantID, &result.ProjectID, &result.Name, &result.KeyPrefix, &result.KeySuffix, + &scopesJSON, &tagsJSON, &result.MonthlySpendMicros, &result.DailySpendMicros, + &result.RequestsPerMinute, &result.TokensPerMinute, &result.Status, &result.ExpiresAt, &result.LastUsedAt, &result.CreatedAt) if err != nil { return CreatedAPIKey{}, 0, fmt.Errorf("create API key: %w", err) @@ -163,6 +163,90 @@ func (s *Store) RevokeAPIKey(ctx context.Context, id string) (int64, error) { return s.toggle(ctx, `UPDATE api_keys SET status = 'revoked', revoked_at = now() WHERE id = $1 AND status <> 'revoked'`, id) } +func (s *Store) DisableAPIKey(ctx context.Context, id string) (int64, error) { + return s.toggle(ctx, `UPDATE api_keys SET status='disabled', disabled_at=now() WHERE id=$1 AND status='active'`, id) +} + +func (s *Store) EnableAPIKey(ctx context.Context, id string) (int64, error) { + return s.toggle(ctx, `UPDATE api_keys SET status='active', disabled_at=NULL WHERE id=$1 AND status='disabled' AND (expires_at IS NULL OR expires_at > now())`, id) +} + +func (s *Store) RotateAPIKey(ctx context.Context, id string) (CreatedAPIKey, int64, error) { + tx, err := s.db.Begin(ctx) + if err != nil { + return CreatedAPIKey{}, 0, err + } + defer tx.Rollback(ctx) + + var tenantID, projectID, name, status string + var scopesJSON, tagsJSON []byte + var monthlySpendMicros, dailySpendMicros, requestsPerMinute, tokensPerMinute int64 + var expiresAt *time.Time + err = tx.QueryRow(ctx, `SELECT tenant_id::text,project_id::text,name,scopes,tags,monthly_spend_micros, + daily_spend_micros,requests_per_minute,tokens_per_minute,expires_at,status + FROM api_keys WHERE id=$1 FOR UPDATE`, id).Scan(&tenantID, &projectID, &name, &scopesJSON, &tagsJSON, + &monthlySpendMicros, &dailySpendMicros, &requestsPerMinute, &tokensPerMinute, &expiresAt, &status) + if errors.Is(err, pgx.ErrNoRows) { + return CreatedAPIKey{}, 0, ErrNotFound + } + if err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("lock API key for rotation: %w", err) + } + if status == "revoked" { + return CreatedAPIKey{}, 0, errors.New("revoked API key cannot be rotated") + } + if expiresAt != nil && !expiresAt.After(time.Now()) { + return CreatedAPIKey{}, 0, errors.New("expired API key cannot be rotated") + } + rawKey, prefix, suffix, hash, err := generateAPIKeySecret() + if err != nil { + return CreatedAPIKey{}, 0, err + } + var result CreatedAPIKey + err = tx.QueryRow(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_suffix,key_hash,scopes,tags, + monthly_spend_micros,daily_spend_micros,requests_per_minute,tokens_per_minute,expires_at,rotated_from_id) + VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14) + RETURNING id::text,tenant_id::text,project_id::text,name,key_prefix,key_suffix,scopes,tags,monthly_spend_micros, + daily_spend_micros,requests_per_minute,tokens_per_minute,status,expires_at,last_used_at,created_at`, + tenantID, projectID, name, prefix, suffix, hash[:], scopesJSON, tagsJSON, monthlySpendMicros, dailySpendMicros, + requestsPerMinute, tokensPerMinute, expiresAt, id).Scan(&result.ID, &result.TenantID, &result.ProjectID, + &result.Name, &result.KeyPrefix, &result.KeySuffix, &scopesJSON, &tagsJSON, &result.MonthlySpendMicros, &result.DailySpendMicros, + &result.RequestsPerMinute, &result.TokensPerMinute, &result.Status, &result.ExpiresAt, &result.LastUsedAt, &result.CreatedAt) + if err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("create rotated API key: %w", err) + } + if _, err := tx.Exec(ctx, `INSERT INTO api_key_model_restrictions (api_key_id,model_id) + SELECT $1,model_id FROM api_key_model_restrictions WHERE api_key_id=$2`, result.ID, id); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("copy rotated API key model restrictions: %w", err) + } + var allowedModelsJSON []byte + if err := tx.QueryRow(ctx, `SELECT COALESCE(jsonb_agg(m.public_id ORDER BY m.public_id),'[]'::jsonb) + FROM api_key_model_restrictions r JOIN models m ON m.id=r.model_id WHERE r.api_key_id=$1`, result.ID).Scan(&allowedModelsJSON); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("read rotated API key model restrictions: %w", err) + } + if _, err := tx.Exec(ctx, `UPDATE api_keys SET status='revoked',revoked_at=now() WHERE id=$1`, id); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("revoke rotated API key: %w", err) + } + if err := json.Unmarshal(scopesJSON, &result.Scopes); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("decode rotated API key scopes: %w", err) + } + if err := json.Unmarshal(tagsJSON, &result.Tags); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("decode rotated API key tags: %w", err) + } + if err := json.Unmarshal(allowedModelsJSON, &result.AllowedModels); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("decode rotated API key model restrictions: %w", err) + } + result.Key = rawKey + generation, err := bumpGeneration(ctx, tx) + if err != nil { + return CreatedAPIKey{}, 0, err + } + if err := tx.Commit(ctx); err != nil { + return CreatedAPIKey{}, 0, err + } + return result, generation, nil +} + func (s *Store) CreateProvider(ctx context.Context, input CreateProviderInput) (Provider, int64, error) { input.Slug = strings.ToLower(strings.TrimSpace(input.Slug)) input.Name = strings.TrimSpace(input.Name) @@ -184,7 +268,7 @@ func (s *Store) CreateProvider(ctx context.Context, input CreateProviderInput) ( if !slugPattern.MatchString(input.Slug) || input.Name == "" || input.APIKey == "" || (input.Protocol != "openai" && input.Protocol != "anthropic") { return Provider{}, 0, errors.New("provider requires a unique 3-64 character lowercase slug, name, protocol openai|anthropic, base_url, and api_key") } - if (input.Protocol == "openai" && input.WireAPI != "chat_completions" && input.WireAPI != "responses") || + if (input.Protocol == "openai" && input.WireAPI != "chat_completions" && input.WireAPI != "responses" && input.WireAPI != "embeddings") || (input.Protocol == "anthropic" && input.WireAPI != "messages") { return Provider{}, 0, errors.New("provider wire_api is incompatible with protocol") } @@ -474,3 +558,15 @@ func uniqueStrings(values []string) []string { } return result } + +func generateAPIKeySecret() (string, string, string, [sha256.Size]byte, error) { + random := make([]byte, 32) + if _, err := rand.Read(random); err != nil { + return "", "", "", [sha256.Size]byte{}, fmt.Errorf("generate API key: %w", err) + } + rawKey := "sk-aigw-" + base64.RawURLEncoding.EncodeToString(random) + hash := sha256.Sum256([]byte(rawKey)) + prefix := rawKey[:min(18, len(rawKey))] + "..." + suffix := rawKey[max(0, len(rawKey)-6):] + return rawKey, prefix, suffix, hash, nil +} diff --git a/internal/controlplane/preferences.go b/internal/controlplane/preferences.go index 73bcce6..2f80ad5 100644 --- a/internal/controlplane/preferences.go +++ b/internal/controlplane/preferences.go @@ -10,22 +10,37 @@ import ( "github.com/jackc/pgx/v5" ) -const maxLowBalanceThresholdMicros int64 = 1_000_000_000_000_000 +const maxAlertThresholdMicros int64 = 1_000_000_000_000_000 -func (s *Store) GetTenantPreferences(ctx context.Context, tenantID string, defaultThresholdMicros int64) (TenantPreferences, error) { - if defaultThresholdMicros <= 0 { - defaultThresholdMicros = 5_000_000 +func normalizeBillingPreferenceDefaults(defaults BillingPreferenceDefaults) BillingPreferenceDefaults { + if defaults.LowBalanceThresholdMicros <= 0 { + defaults.LowBalanceThresholdMicros = 5_000_000 } - result := TenantPreferences{TenantID: strings.TrimSpace(tenantID), LowBalanceEnabled: true, LowBalanceThresholdMicros: defaultThresholdMicros} + if defaults.SpendAnomalyMultiplier < 2 { + defaults.SpendAnomalyMultiplier = 3 + } + if defaults.SpendAnomalyMinMicros < 0 { + defaults.SpendAnomalyMinMicros = 10_000_000 + } + return defaults +} + +func (s *Store) GetTenantPreferences(ctx context.Context, tenantID string, defaults BillingPreferenceDefaults) (TenantPreferences, error) { + defaults = normalizeBillingPreferenceDefaults(defaults) + result := TenantPreferences{TenantID: strings.TrimSpace(tenantID), LowBalanceEnabled: true, + LowBalanceThresholdMicros: defaults.LowBalanceThresholdMicros, SpendAnomalyEnabled: true, + SpendAnomalyMultiplier: defaults.SpendAnomalyMultiplier, SpendAnomalyMinMicros: defaults.SpendAnomalyMinMicros} if result.TenantID == "" { return result, nil } var defaultModel, fallbackModel *string var updatedAt time.Time err := s.db.QueryRow(ctx, ` - SELECT default_model, fallback_model, low_balance_enabled, low_balance_threshold_micros, updated_at - FROM tenant_preferences WHERE tenant_id=$1`, result.TenantID).Scan( - &defaultModel, &fallbackModel, &result.LowBalanceEnabled, &result.LowBalanceThresholdMicros, &updatedAt) + SELECT default_model, fallback_model, low_balance_enabled, low_balance_threshold_micros, + spend_anomaly_enabled, COALESCE(spend_anomaly_multiplier,$2), COALESCE(spend_anomaly_min_micros,$3), updated_at + FROM tenant_preferences WHERE tenant_id=$1`, result.TenantID, defaults.SpendAnomalyMultiplier, defaults.SpendAnomalyMinMicros).Scan( + &defaultModel, &fallbackModel, &result.LowBalanceEnabled, &result.LowBalanceThresholdMicros, + &result.SpendAnomalyEnabled, &result.SpendAnomalyMultiplier, &result.SpendAnomalyMinMicros, &updatedAt) if errors.Is(err, pgx.ErrNoRows) { return result, nil } @@ -42,10 +57,11 @@ func (s *Store) GetTenantPreferences(ctx context.Context, tenantID string, defau return result, nil } -func (s *Store) SetDeveloperPreferences(ctx context.Context, input SetDeveloperPreferencesInput) (TenantPreferences, error) { +func (s *Store) SetDeveloperPreferences(ctx context.Context, input SetDeveloperPreferencesInput, defaults BillingPreferenceDefaults) (TenantPreferences, error) { input.TenantID = strings.TrimSpace(input.TenantID) input.DefaultModel = strings.TrimSpace(input.DefaultModel) input.FallbackModel = strings.TrimSpace(input.FallbackModel) + defaults = normalizeBillingPreferenceDefaults(defaults) if input.TenantID == "" { return TenantPreferences{}, errors.New("tenant_id is required") } @@ -87,9 +103,12 @@ func (s *Store) SetDeveloperPreferences(ctx context.Context, input SetDeveloperP ON CONFLICT (tenant_id) DO UPDATE SET default_model=EXCLUDED.default_model, fallback_model=EXCLUDED.fallback_model, updated_at=now() RETURNING tenant_id::text, default_model, fallback_model, low_balance_enabled, - low_balance_threshold_micros, updated_at`, input.TenantID, input.DefaultModel, input.FallbackModel).Scan( + low_balance_threshold_micros, spend_anomaly_enabled, COALESCE(spend_anomaly_multiplier,$4), + COALESCE(spend_anomaly_min_micros,$5), updated_at`, input.TenantID, input.DefaultModel, input.FallbackModel, + defaults.SpendAnomalyMultiplier, defaults.SpendAnomalyMinMicros).Scan( &result.TenantID, &defaultModel, &fallbackModel, &result.LowBalanceEnabled, - &result.LowBalanceThresholdMicros, &updatedAt); err != nil { + &result.LowBalanceThresholdMicros, &result.SpendAnomalyEnabled, &result.SpendAnomalyMultiplier, + &result.SpendAnomalyMinMicros, &updatedAt); err != nil { return TenantPreferences{}, fmt.Errorf("save developer preferences: %w", err) } if defaultModel != nil { @@ -105,17 +124,21 @@ func (s *Store) SetDeveloperPreferences(ctx context.Context, input SetDeveloperP return result, nil } -func (s *Store) SetBillingPreferences(ctx context.Context, input SetBillingPreferencesInput, defaultThresholdMicros int64) (TenantPreferences, error) { +func (s *Store) SetBillingPreferences(ctx context.Context, input SetBillingPreferencesInput, defaults BillingPreferenceDefaults) (TenantPreferences, error) { input.TenantID = strings.TrimSpace(input.TenantID) if input.TenantID == "" { return TenantPreferences{}, errors.New("tenant_id is required") } - if defaultThresholdMicros <= 0 { - defaultThresholdMicros = 5_000_000 - } - if input.LowBalanceThresholdMicros != nil && (*input.LowBalanceThresholdMicros < 0 || *input.LowBalanceThresholdMicros > maxLowBalanceThresholdMicros) { + defaults = normalizeBillingPreferenceDefaults(defaults) + if input.LowBalanceThresholdMicros != nil && (*input.LowBalanceThresholdMicros < 0 || *input.LowBalanceThresholdMicros > maxAlertThresholdMicros) { return TenantPreferences{}, errors.New("low balance threshold is outside the supported range") } + if input.SpendAnomalyMultiplier != nil && (*input.SpendAnomalyMultiplier < 2 || *input.SpendAnomalyMultiplier > 1000) { + return TenantPreferences{}, errors.New("spend anomaly multiplier must be between 2 and 1000") + } + if input.SpendAnomalyMinMicros != nil && (*input.SpendAnomalyMinMicros < 0 || *input.SpendAnomalyMinMicros > maxAlertThresholdMicros) { + return TenantPreferences{}, errors.New("spend anomaly minimum is outside the supported range") + } tx, err := s.db.Begin(ctx) if err != nil { return TenantPreferences{}, err @@ -131,15 +154,25 @@ func (s *Store) SetBillingPreferences(ctx context.Context, input SetBillingPrefe var defaultModel, fallbackModel *string var updatedAt time.Time if err := tx.QueryRow(ctx, ` - INSERT INTO tenant_preferences (tenant_id, low_balance_enabled, low_balance_threshold_micros) - VALUES ($1, COALESCE($2::boolean, TRUE), COALESCE($3::bigint, $4::bigint)) + INSERT INTO tenant_preferences (tenant_id, low_balance_enabled, low_balance_threshold_micros, + spend_anomaly_enabled, spend_anomaly_multiplier, spend_anomaly_min_micros) + VALUES ($1, COALESCE($2::boolean, TRUE), COALESCE($3::bigint, $7::bigint), + COALESCE($4::boolean, TRUE), $5::bigint, $6::bigint) ON CONFLICT (tenant_id) DO UPDATE SET low_balance_enabled=COALESCE($2::boolean, tenant_preferences.low_balance_enabled), - low_balance_threshold_micros=COALESCE($3::bigint, tenant_preferences.low_balance_threshold_micros), updated_at=now() + low_balance_threshold_micros=COALESCE($3::bigint, tenant_preferences.low_balance_threshold_micros), + spend_anomaly_enabled=COALESCE($4::boolean, tenant_preferences.spend_anomaly_enabled), + spend_anomaly_multiplier=COALESCE($5::bigint, tenant_preferences.spend_anomaly_multiplier), + spend_anomaly_min_micros=COALESCE($6::bigint, tenant_preferences.spend_anomaly_min_micros), updated_at=now() RETURNING tenant_id::text, default_model, fallback_model, low_balance_enabled, - low_balance_threshold_micros, updated_at`, input.TenantID, input.LowBalanceEnabled, input.LowBalanceThresholdMicros, defaultThresholdMicros).Scan( + low_balance_threshold_micros, spend_anomaly_enabled, + COALESCE(spend_anomaly_multiplier,$8), COALESCE(spend_anomaly_min_micros,$9), updated_at`, + input.TenantID, input.LowBalanceEnabled, input.LowBalanceThresholdMicros, input.SpendAnomalyEnabled, + input.SpendAnomalyMultiplier, input.SpendAnomalyMinMicros, defaults.LowBalanceThresholdMicros, + defaults.SpendAnomalyMultiplier, defaults.SpendAnomalyMinMicros).Scan( &result.TenantID, &defaultModel, &fallbackModel, &result.LowBalanceEnabled, - &result.LowBalanceThresholdMicros, &updatedAt); err != nil { + &result.LowBalanceThresholdMicros, &result.SpendAnomalyEnabled, &result.SpendAnomalyMultiplier, + &result.SpendAnomalyMinMicros, &updatedAt); err != nil { return TenantPreferences{}, fmt.Errorf("save billing preferences: %w", err) } if defaultModel != nil { diff --git a/internal/controlplane/preferences_test.go b/internal/controlplane/preferences_test.go index 88a3d50..2d390d2 100644 --- a/internal/controlplane/preferences_test.go +++ b/internal/controlplane/preferences_test.go @@ -6,11 +6,13 @@ import ( ) func TestGetTenantPreferencesWithoutTenantUsesConfiguredDefault(t *testing.T) { - result, err := (&Store{}).GetTenantPreferences(context.Background(), "", 12_500_000) + defaults := BillingPreferenceDefaults{LowBalanceThresholdMicros: 12_500_000, SpendAnomalyMultiplier: 7, SpendAnomalyMinMicros: 8_500_000} + result, err := (&Store{}).GetTenantPreferences(context.Background(), "", defaults) if err != nil { t.Fatal(err) } - if !result.LowBalanceEnabled || result.LowBalanceThresholdMicros != 12_500_000 { + if !result.LowBalanceEnabled || result.LowBalanceThresholdMicros != 12_500_000 || !result.SpendAnomalyEnabled || + result.SpendAnomalyMultiplier != 7 || result.SpendAnomalyMinMicros != 8_500_000 { t.Fatalf("unexpected defaults: %+v", result) } } @@ -19,13 +21,19 @@ func TestPreferenceValidationRejectsUnsafeValuesBeforeDatabaseAccess(t *testing. store := &Store{} if _, err := store.SetDeveloperPreferences(context.Background(), SetDeveloperPreferencesInput{ TenantID: "tenant", DefaultModel: "same", FallbackModel: "same", - }); err == nil { + }, BillingPreferenceDefaults{}); err == nil { t.Fatal("expected identical default and fallback models to fail") } - threshold := maxLowBalanceThresholdMicros + 1 + threshold := maxAlertThresholdMicros + 1 if _, err := store.SetBillingPreferences(context.Background(), SetBillingPreferencesInput{ TenantID: "tenant", LowBalanceThresholdMicros: &threshold, - }, 5_000_000); err == nil { + }, BillingPreferenceDefaults{}); err == nil { t.Fatal("expected excessive low balance threshold to fail") } + multiplier := int64(1) + if _, err := store.SetBillingPreferences(context.Background(), SetBillingPreferencesInput{ + TenantID: "tenant", SpendAnomalyMultiplier: &multiplier, + }, BillingPreferenceDefaults{}); err == nil { + t.Fatal("expected invalid spend anomaly multiplier to fail") + } } diff --git a/internal/controlplane/queries.go b/internal/controlplane/queries.go index 48610c7..6a5ff0b 100644 --- a/internal/controlplane/queries.go +++ b/internal/controlplane/queries.go @@ -85,12 +85,17 @@ func (s *Store) ListAPIKeys(ctx context.Context) ([]APIKey, error) { } func (s *Store) ListAPIKeysFor(ctx context.Context, tenantID string) ([]APIKey, error) { - periodStart := time.Date(time.Now().UTC().Year(), time.Now().UTC().Month(), 1, 0, 0, 0, 0, time.UTC) + now := time.Now().UTC() + periodStart := time.Date(now.Year(), now.Month(), 1, 0, 0, 0, 0, time.UTC) periodEnd := periodStart.AddDate(0, 1, 0) + dayStart := time.Date(now.Year(), now.Month(), now.Day(), 0, 0, 0, 0, time.UTC) + dayEnd := dayStart.AddDate(0, 0, 1) query := ` - SELECT k.id::text, k.tenant_id::text, k.project_id::text, k.name, k.key_prefix, k.scopes, - k.tags, k.monthly_spend_micros, k.status, k.expires_at, k.last_used_at, k.created_at, + SELECT k.id::text, k.tenant_id::text, k.project_id::text, k.name, k.key_prefix, k.key_suffix, k.scopes, + k.tags, k.monthly_spend_micros, k.daily_spend_micros, k.requests_per_minute, k.tokens_per_minute, + k.status, k.expires_at, k.last_used_at, k.created_at, usage.month_spend, usage.month_requests, pending.month_reserved, + daily.day_spend, daily.day_requests, daily_pending.day_reserved, COALESCE(( SELECT jsonb_agg(m.public_id ORDER BY m.public_id) FROM api_key_model_restrictions r @@ -107,10 +112,20 @@ func (s *Store) ListAPIKeysFor(ctx context.Context, tenantID string) ([]APIKey, FROM billing_reservations b WHERE b.key_id = k.id AND b.status IN ('pending', 'metering_failed') AND b.created_at >= $1 AND b.created_at < $2 - ) pending` - args := []any{periodStart, periodEnd} + ) pending + CROSS JOIN LATERAL ( + SELECT COALESCE(SUM(u.cost_micros), 0)::bigint AS day_spend, COUNT(*)::bigint AS day_requests + FROM usage_events u WHERE u.key_id = k.id AND u.started_at >= $3 AND u.started_at < $4 + ) daily + CROSS JOIN LATERAL ( + SELECT COALESCE(SUM(b.reserved_micros), 0)::bigint AS day_reserved + FROM billing_reservations b + WHERE b.key_id = k.id AND b.status IN ('pending', 'metering_failed') + AND b.created_at >= $3 AND b.created_at < $4 + ) daily_pending` + args := []any{periodStart, periodEnd, dayStart, dayEnd} if tenantID != "" { - query += ` WHERE k.tenant_id=$3` + query += ` WHERE k.tenant_id=$5` args = append(args, tenantID) } query += ` ORDER BY k.created_at DESC` @@ -123,10 +138,12 @@ func (s *Store) ListAPIKeysFor(ctx context.Context, tenantID string) ([]APIKey, for rows.Next() { var item APIKey var scopesJSON, tagsJSON, allowedModelsJSON []byte - if err := rows.Scan(&item.ID, &item.TenantID, &item.ProjectID, &item.Name, &item.KeyPrefix, - &scopesJSON, &tagsJSON, &item.MonthlySpendMicros, &item.Status, &item.ExpiresAt, + if err := rows.Scan(&item.ID, &item.TenantID, &item.ProjectID, &item.Name, &item.KeyPrefix, &item.KeySuffix, + &scopesJSON, &tagsJSON, &item.MonthlySpendMicros, &item.DailySpendMicros, &item.RequestsPerMinute, + &item.TokensPerMinute, &item.Status, &item.ExpiresAt, &item.LastUsedAt, &item.CreatedAt, &item.CurrentMonthSpendMicros, &item.CurrentMonthRequests, - &item.CurrentMonthReservedMicros, &allowedModelsJSON); err != nil { + &item.CurrentMonthReservedMicros, &item.CurrentDaySpendMicros, &item.CurrentDayRequests, + &item.CurrentDayReservedMicros, &allowedModelsJSON); err != nil { return nil, fmt.Errorf("scan API key: %w", err) } if err := json.Unmarshal(scopesJSON, &item.Scopes); err != nil { diff --git a/internal/controlplane/schema.sql b/internal/controlplane/schema.sql index ef1ccdd..d71f37f 100644 --- a/internal/controlplane/schema.sql +++ b/internal/controlplane/schema.sql @@ -41,24 +41,35 @@ CREATE TABLE IF NOT EXISTS api_keys ( project_id UUID NOT NULL, name TEXT NOT NULL, key_prefix TEXT NOT NULL, + key_suffix TEXT NOT NULL DEFAULT '', key_hash BYTEA NOT NULL UNIQUE CHECK (octet_length(key_hash) = 32), scopes JSONB NOT NULL DEFAULT '["inference"]'::jsonb, - status TEXT NOT NULL DEFAULT 'active' CHECK (status IN ('active', 'revoked')), + status TEXT NOT NULL DEFAULT 'active' CHECK (status IN ('active', 'disabled', 'revoked')), last_used_at TIMESTAMPTZ, created_at TIMESTAMPTZ NOT NULL DEFAULT now(), revoked_at TIMESTAMPTZ, FOREIGN KEY (project_id, tenant_id) REFERENCES projects(id, tenant_id) ON DELETE CASCADE ); ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS monthly_spend_micros BIGINT NOT NULL DEFAULT 0 CHECK (monthly_spend_micros >= 0); +ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS daily_spend_micros BIGINT NOT NULL DEFAULT 0 CHECK (daily_spend_micros >= 0); +ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS requests_per_minute BIGINT NOT NULL DEFAULT 0 CHECK (requests_per_minute >= 0); +ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS tokens_per_minute BIGINT NOT NULL DEFAULT 0 CHECK (tokens_per_minute >= 0); ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS expires_at TIMESTAMPTZ; ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS tags JSONB NOT NULL DEFAULT '[]'::jsonb; +ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS disabled_at TIMESTAMPTZ; +ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS rotated_from_id UUID REFERENCES api_keys(id) ON DELETE SET NULL; +ALTER TABLE api_keys ADD COLUMN IF NOT EXISTS key_suffix TEXT NOT NULL DEFAULT ''; +ALTER TABLE api_keys DROP CONSTRAINT IF EXISTS api_keys_key_suffix_check; +ALTER TABLE api_keys ADD CONSTRAINT api_keys_key_suffix_check CHECK (key_suffix = '' OR key_suffix ~ '^[A-Za-z0-9_-]{6}$'); +ALTER TABLE api_keys DROP CONSTRAINT IF EXISTS api_keys_status_check; +ALTER TABLE api_keys ADD CONSTRAINT api_keys_status_check CHECK (status IN ('active', 'disabled', 'revoked')); CREATE TABLE IF NOT EXISTS providers ( id UUID PRIMARY KEY DEFAULT gen_random_uuid(), slug TEXT NOT NULL UNIQUE CHECK (slug ~ '^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$'), name TEXT NOT NULL UNIQUE, protocol TEXT NOT NULL CHECK (protocol IN ('openai', 'anthropic')), - wire_api TEXT NOT NULL DEFAULT 'chat_completions' CHECK (wire_api IN ('chat_completions', 'responses', 'messages')), + wire_api TEXT NOT NULL DEFAULT 'chat_completions' CHECK (wire_api IN ('chat_completions', 'responses', 'embeddings', 'messages')), base_url TEXT NOT NULL, api_key_ciphertext BYTEA NOT NULL, enabled BOOLEAN NOT NULL DEFAULT TRUE, @@ -90,10 +101,10 @@ ALTER TABLE providers ADD CONSTRAINT providers_slug_check CHECK (slug ~ '^[a-z0- ALTER TABLE providers ADD COLUMN IF NOT EXISTS wire_api TEXT NOT NULL DEFAULT 'chat_completions'; UPDATE providers SET wire_api = 'messages' WHERE protocol = 'anthropic' AND wire_api = 'chat_completions'; ALTER TABLE providers DROP CONSTRAINT IF EXISTS providers_wire_api_check; -ALTER TABLE providers ADD CONSTRAINT providers_wire_api_check CHECK (wire_api IN ('chat_completions', 'responses', 'messages')); +ALTER TABLE providers ADD CONSTRAINT providers_wire_api_check CHECK (wire_api IN ('chat_completions', 'responses', 'embeddings', 'messages')); ALTER TABLE providers DROP CONSTRAINT IF EXISTS providers_protocol_wire_api_check; ALTER TABLE providers ADD CONSTRAINT providers_protocol_wire_api_check CHECK ( - (protocol = 'openai' AND wire_api IN ('chat_completions', 'responses')) OR + (protocol = 'openai' AND wire_api IN ('chat_completions', 'responses', 'embeddings')) OR (protocol = 'anthropic' AND wire_api = 'messages') ); @@ -204,6 +215,7 @@ CREATE TABLE IF NOT EXISTS model_routes ( ); CREATE INDEX IF NOT EXISTS api_keys_active_hash_idx ON api_keys (key_hash) WHERE status = 'active'; +CREATE INDEX IF NOT EXISTS api_keys_rotated_from_idx ON api_keys (rotated_from_id) WHERE rotated_from_id IS NOT NULL; CREATE INDEX IF NOT EXISTS projects_tenant_idx ON projects (tenant_id); CREATE INDEX IF NOT EXISTS model_routes_model_idx ON model_routes (model_id) WHERE enabled; CREATE INDEX IF NOT EXISTS model_routes_provider_idx ON model_routes (provider_id) WHERE enabled; @@ -224,9 +236,21 @@ CREATE TABLE IF NOT EXISTS tenant_preferences ( fallback_model TEXT REFERENCES models(public_id) ON DELETE SET NULL, low_balance_enabled BOOLEAN NOT NULL DEFAULT TRUE, low_balance_threshold_micros BIGINT NOT NULL DEFAULT 5000000 CHECK (low_balance_threshold_micros >= 0), + spend_anomaly_enabled BOOLEAN NOT NULL DEFAULT TRUE, + spend_anomaly_multiplier BIGINT, + spend_anomaly_min_micros BIGINT, updated_at TIMESTAMPTZ NOT NULL DEFAULT now(), CHECK (default_model IS NULL OR fallback_model IS NULL OR default_model <> fallback_model) ); +ALTER TABLE tenant_preferences ADD COLUMN IF NOT EXISTS spend_anomaly_enabled BOOLEAN NOT NULL DEFAULT TRUE; +ALTER TABLE tenant_preferences ADD COLUMN IF NOT EXISTS spend_anomaly_multiplier BIGINT; +ALTER TABLE tenant_preferences ADD COLUMN IF NOT EXISTS spend_anomaly_min_micros BIGINT; +ALTER TABLE tenant_preferences DROP CONSTRAINT IF EXISTS tenant_preferences_spend_anomaly_multiplier_check; +ALTER TABLE tenant_preferences ADD CONSTRAINT tenant_preferences_spend_anomaly_multiplier_check + CHECK (spend_anomaly_multiplier IS NULL OR spend_anomaly_multiplier BETWEEN 2 AND 1000); +ALTER TABLE tenant_preferences DROP CONSTRAINT IF EXISTS tenant_preferences_spend_anomaly_min_check; +ALTER TABLE tenant_preferences ADD CONSTRAINT tenant_preferences_spend_anomaly_min_check + CHECK (spend_anomaly_min_micros IS NULL OR spend_anomaly_min_micros BETWEEN 0 AND 1000000000000000); CREATE TABLE IF NOT EXISTS billing_reservations ( request_id TEXT PRIMARY KEY, @@ -289,6 +313,7 @@ CREATE TABLE IF NOT EXISTS usage_events ( attempts INTEGER NOT NULL DEFAULT 0, started_at TIMESTAMPTZ NOT NULL, duration_ms BIGINT NOT NULL DEFAULT 0, + ttft_ms BIGINT NOT NULL DEFAULT 0 CHECK (ttft_ms >= 0), input_tokens BIGINT NOT NULL DEFAULT 0, output_tokens BIGINT NOT NULL DEFAULT 0, total_tokens BIGINT NOT NULL DEFAULT 0, @@ -299,15 +324,18 @@ CREATE TABLE IF NOT EXISTS usage_events ( uncollected_micros BIGINT NOT NULL DEFAULT 0, usage_reported BOOLEAN NOT NULL DEFAULT FALSE, metering_status TEXT NOT NULL DEFAULT 'not_billable' - CHECK (metering_status IN ('not_billable','reported','missing','upstream_failed')), + CHECK (metering_status IN ('not_billable','reported','missing','released_unmetered','upstream_failed')), created_at TIMESTAMPTZ NOT NULL DEFAULT now(), FOREIGN KEY (project_id, tenant_id) REFERENCES projects(id, tenant_id) ON DELETE RESTRICT ); ALTER TABLE usage_events ADD COLUMN IF NOT EXISTS usage_reported BOOLEAN NOT NULL DEFAULT FALSE; ALTER TABLE usage_events ADD COLUMN IF NOT EXISTS metering_status TEXT NOT NULL DEFAULT 'not_billable'; +ALTER TABLE usage_events ADD COLUMN IF NOT EXISTS ttft_ms BIGINT NOT NULL DEFAULT 0; +ALTER TABLE usage_events DROP CONSTRAINT IF EXISTS usage_events_ttft_ms_check; +ALTER TABLE usage_events ADD CONSTRAINT usage_events_ttft_ms_check CHECK (ttft_ms >= 0); ALTER TABLE usage_events DROP CONSTRAINT IF EXISTS usage_events_metering_status_check; ALTER TABLE usage_events ADD CONSTRAINT usage_events_metering_status_check - CHECK (metering_status IN ('not_billable','reported','missing','upstream_failed')); + CHECK (metering_status IN ('not_billable','reported','missing','released_unmetered','upstream_failed')); -- Usage persistence is independent from billing. Older installations created this -- foreign key, which prevented recording requests when prepaid billing was disabled. diff --git a/internal/controlplane/snapshot.go b/internal/controlplane/snapshot.go index 75a3a80..48be436 100644 --- a/internal/controlplane/snapshot.go +++ b/internal/controlplane/snapshot.go @@ -246,7 +246,7 @@ func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Prov func loadAPIKeys(ctx context.Context, tx pgx.Tx) ([]auth.HashedKeyRecord, error) { rows, err := tx.Query(ctx, ` SELECT k.id::text, k.key_hash, k.tenant_id::text, k.project_id::text, k.scopes, - k.monthly_spend_micros, k.expires_at, + k.monthly_spend_micros, k.daily_spend_micros, k.requests_per_minute, k.tokens_per_minute, k.expires_at, COALESCE(( SELECT jsonb_agg(m.public_id ORDER BY m.public_id) FROM api_key_model_restrictions r @@ -265,10 +265,10 @@ func loadAPIKeys(ctx context.Context, tx pgx.Tx) ([]auth.HashedKeyRecord, error) for rows.Next() { var keyID, tenantID, projectID string var hashBytes, scopesJSON, allowedModelsJSON []byte - var monthlySpendMicros int64 + var monthlySpendMicros, dailySpendMicros, requestsPerMinute, tokensPerMinute int64 var expiresAt *time.Time if err := rows.Scan(&keyID, &hashBytes, &tenantID, &projectID, &scopesJSON, - &monthlySpendMicros, &expiresAt, &allowedModelsJSON); err != nil { + &monthlySpendMicros, &dailySpendMicros, &requestsPerMinute, &tokensPerMinute, &expiresAt, &allowedModelsJSON); err != nil { return nil, fmt.Errorf("scan API key: %w", err) } if len(hashBytes) != sha256.Size { @@ -290,7 +290,8 @@ func loadAPIKeys(ctx context.Context, tx pgx.Tx) ([]auth.HashedKeyRecord, error) } records = append(records, auth.HashedKeyRecord{Hash: hash, Principal: domain.Principal{ KeyID: keyID, TenantID: tenantID, ProjectID: projectID, Scopes: scopes, - AllowedModels: allowedModels, MonthlySpendMicros: monthlySpendMicros, ExpiresAt: expiresAt, + AllowedModels: allowedModels, MonthlySpendMicros: monthlySpendMicros, DailySpendMicros: dailySpendMicros, + RequestsPerMinute: requestsPerMinute, TokensPerMinute: tokensPerMinute, ExpiresAt: expiresAt, }}) } if err := rows.Err(); err != nil { diff --git a/internal/controlplane/store.go b/internal/controlplane/store.go index c4d7016..82c87ae 100644 --- a/internal/controlplane/store.go +++ b/internal/controlplane/store.go @@ -23,7 +23,10 @@ var schemaSQL string var ErrRedisDisabled = errors.New("Redis propagation is disabled") -const migrationVersion int64 = 2026080605 +const ( + migrationVersion int64 = 2026080610 + migrationLockID int64 = 0x41494757 // "AIGW"; stable across migration versions. +) type Options struct { DatabaseURL string @@ -82,6 +85,13 @@ func (s *Store) RedisEnabled() bool { return s.redis != nil } +func (s *Store) PingRedis(ctx context.Context) error { + if s.redis == nil { + return ErrRedisDisabled + } + return s.redis.Ping(ctx).Err() +} + func (s *Store) Ping(ctx context.Context) error { return s.db.Ping(ctx) } func (s *Store) Migrate(ctx context.Context) error { @@ -112,28 +122,40 @@ func MigrationStatusDatabase(ctx context.Context, databaseURL string) (Migration } func applySchema(ctx context.Context, db *pgxpool.Pool) error { + hash := sha256.Sum256([]byte(schemaSQL)) + checksum := hex.EncodeToString(hash[:]) tx, err := db.Begin(ctx) if err != nil { return fmt.Errorf("begin migration: %w", err) } defer tx.Rollback(ctx) - if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock($1)`, migrationVersion); err != nil { + if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock($1)`, migrationLockID); err != nil { return err } - if _, err := tx.Exec(ctx, schemaSQL); err != nil { - return fmt.Errorf("apply control-plane schema: %w", err) + var migrationsExist bool + if err := tx.QueryRow(ctx, `SELECT to_regclass('schema_migrations') IS NOT NULL`).Scan(&migrationsExist); err != nil { + return fmt.Errorf("inspect migration table: %w", err) } - hash := sha256.Sum256([]byte(schemaSQL)) - checksum := hex.EncodeToString(hash[:]) - var existing string - err = tx.QueryRow(ctx, `SELECT checksum FROM schema_migrations WHERE version=$1`, migrationVersion).Scan(&existing) - if err == nil && existing != checksum { - return fmt.Errorf("migration %d checksum changed; deploy an explicit new migration version", migrationVersion) + if migrationsExist { + var existing string + err = tx.QueryRow(ctx, `SELECT checksum FROM schema_migrations WHERE version=$1`, migrationVersion).Scan(&existing) + if err == nil { + if existing != checksum { + return fmt.Errorf("migration %d checksum changed; deploy an explicit new migration version", migrationVersion) + } + if err := tx.Commit(ctx); err != nil { + return fmt.Errorf("commit migration check: %w", err) + } + return nil + } + if !errors.Is(err, pgx.ErrNoRows) { + return fmt.Errorf("read migration checksum: %w", err) + } } - if !errors.Is(err, pgx.ErrNoRows) && err != nil { - return err + if _, err := tx.Exec(ctx, schemaSQL); err != nil { + return fmt.Errorf("apply control-plane schema: %w", err) } - if _, err := tx.Exec(ctx, `INSERT INTO schema_migrations(version,name,checksum) VALUES ($1,$2,$3) ON CONFLICT DO NOTHING`, migrationVersion, "tenant-billing-profiles", checksum); err != nil { + if _, err := tx.Exec(ctx, `INSERT INTO schema_migrations(version,name,checksum) VALUES ($1,$2,$3) ON CONFLICT DO NOTHING`, migrationVersion, "tenant-spend-alert-preferences", checksum); err != nil { return err } if err := tx.Commit(ctx); err != nil { diff --git a/internal/controlplane/store_integration_test.go b/internal/controlplane/store_integration_test.go index 9bf093c..b0038c1 100644 --- a/internal/controlplane/store_integration_test.go +++ b/internal/controlplane/store_integration_test.go @@ -6,6 +6,7 @@ import ( "fmt" "net/url" "os" + "strings" "testing" "time" @@ -54,25 +55,30 @@ func TestAPIKeyRestrictionsRoundTripIntoRuntimeSnapshotPostgres(t *testing.T) { if err != nil { t.Fatal(err) } - provider, _, err := store.CreateProvider(ctx, CreateProviderInput{Name: "key-control-provider", Protocol: "openai", WireAPI: "responses", BaseURL: "https://example.invalid/v1", APIKey: "provider-secret"}) + provider, _, err := store.CreateProvider(ctx, CreateProviderInput{Name: "key-control-provider", Protocol: "openai", WireAPI: "embeddings", BaseURL: "https://example.invalid/v1", APIKey: "provider-secret"}) if err != nil { t.Fatal(err) } if provider.Slug != "key-control-provider" { t.Fatalf("derived provider slug = %q, want key-control-provider", provider.Slug) } - model, _, err := store.CreateModel(ctx, CreateModelInput{PublicID: "model/key-control", DisplayName: "Key Control", PriceCurrency: "usd", InputPriceMicrosPerMillion: 100_000, OutputPriceMicrosPerMillion: 200_000, Routes: []RouteInput{{ProviderID: provider.ID, UpstreamModel: "upstream-key-control", Weight: 1}}}) + if _, _, err := store.CreateProvider(ctx, CreateProviderInput{Name: "invalid-anthropic-embeddings", Protocol: "anthropic", WireAPI: "embeddings", BaseURL: "https://example.invalid/v1", APIKey: "provider-secret"}); err == nil { + t.Fatal("Anthropic provider accepted the OpenAI Embeddings wire API") + } + model, _, err := store.CreateModel(ctx, CreateModelInput{PublicID: "model/key-control", DisplayName: "Key Control", Capabilities: []string{"embeddings"}, PriceCurrency: "usd", InputPriceMicrosPerMillion: 100_000, OutputPriceMicrosPerMillion: 200_000, Routes: []RouteInput{{ProviderID: provider.ID, UpstreamModel: "upstream-key-control", Weight: 1}}}) if err != nil { t.Fatal(err) } expiresAt := time.Now().Add(24 * time.Hour).UTC().Truncate(time.Microsecond) created, _, err := store.CreateAPIKey(ctx, CreateAPIKeyInput{TenantID: tenant.ID, ProjectID: project.ID, Name: "production backend", Scopes: []string{"inference"}, Tags: []string{"production", "backend"}, - AllowedModels: []string{model.PublicID}, MonthlySpendMicros: 25_000_000, ExpiresAt: &expiresAt}) + AllowedModels: []string{model.PublicID}, MonthlySpendMicros: 25_000_000, DailySpendMicros: 5_000_000, + RequestsPerMinute: 12, TokensPerMinute: 34_000, ExpiresAt: &expiresAt}) if err != nil { t.Fatal(err) } - if created.Key == "" || created.MonthlySpendMicros != 25_000_000 || len(created.AllowedModels) != 1 || len(created.Tags) != 2 { + if created.Key == "" || len(created.KeySuffix) != 6 || !strings.HasSuffix(created.Key, created.KeySuffix) || + created.MonthlySpendMicros != 25_000_000 || created.DailySpendMicros != 5_000_000 || created.RequestsPerMinute != 12 || created.TokensPerMinute != 34_000 || len(created.AllowedModels) != 1 || len(created.Tags) != 2 { t.Fatalf("unexpected created key: %+v", created.APIKey) } if _, err := store.db.Exec(ctx, `INSERT INTO usage_events @@ -93,12 +99,15 @@ func TestAPIKeyRestrictionsRoundTripIntoRuntimeSnapshotPostgres(t *testing.T) { if err != nil { t.Fatal(err) } - if len(keys) != 1 || keys[0].AllowedModels[0] != model.PublicID || keys[0].ExpiresAt == nil || !keys[0].ExpiresAt.Equal(expiresAt) { + if len(keys) != 1 || keys[0].KeySuffix != created.KeySuffix || keys[0].AllowedModels[0] != model.PublicID || keys[0].ExpiresAt == nil || !keys[0].ExpiresAt.Equal(expiresAt) { t.Fatalf("key restrictions did not round trip: %+v", keys) } if keys[0].CurrentMonthSpendMicros != 42_000 || keys[0].CurrentMonthReservedMicros != 9_000 || keys[0].CurrentMonthRequests != 1 { t.Fatalf("key month activity is incorrect: %+v", keys[0]) } + if keys[0].CurrentDaySpendMicros != 42_000 || keys[0].CurrentDayReservedMicros != 9_000 || keys[0].CurrentDayRequests != 1 { + t.Fatalf("key daily activity is incorrect: %+v", keys[0]) + } snapshot, err := store.LoadSnapshot(ctx) if err != nil { t.Fatal(err) @@ -106,8 +115,11 @@ func TestAPIKeyRestrictionsRoundTripIntoRuntimeSnapshotPostgres(t *testing.T) { if len(snapshot.APIKeys) != 1 { t.Fatalf("snapshot API keys = %d, want 1", len(snapshot.APIKeys)) } + if len(snapshot.Models) != 1 || len(snapshot.Models[0].Routes) != 1 || snapshot.Models[0].Routes[0].Provider.EffectiveWireAPI() != "embeddings" || len(snapshot.Models[0].Capabilities) != 1 || snapshot.Models[0].Capabilities[0] != "embeddings" { + t.Fatalf("Embeddings model did not round trip into runtime snapshot: %+v", snapshot.Models) + } principal := snapshot.APIKeys[0].Principal - if principal.MonthlySpendMicros != 25_000_000 || principal.ExpiresAt == nil { + if principal.MonthlySpendMicros != 25_000_000 || principal.DailySpendMicros != 5_000_000 || principal.RequestsPerMinute != 12 || principal.TokensPerMinute != 34_000 || principal.ExpiresAt == nil { t.Fatalf("snapshot lost API key controls: %+v", principal) } if _, ok := principal.AllowedModels[model.PublicID]; !ok { @@ -132,6 +144,38 @@ func TestAPIKeyRestrictionsRoundTripIntoRuntimeSnapshotPostgres(t *testing.T) { if err != nil || len(keys) != 1 { t.Fatalf("failed key transaction leaked a row: keys=%d err=%v", len(keys), err) } + if _, err := store.DisableAPIKey(ctx, created.ID); err != nil { + t.Fatal(err) + } + snapshot, err = store.LoadSnapshot(ctx) + if err != nil { + t.Fatal(err) + } + if len(snapshot.APIKeys) != 0 { + t.Fatalf("disabled key remained in runtime snapshot: %+v", snapshot.APIKeys) + } + if _, err := store.EnableAPIKey(ctx, created.ID); err != nil { + t.Fatal(err) + } + rotated, _, err := store.RotateAPIKey(ctx, created.ID) + if err != nil { + t.Fatal(err) + } + if rotated.ID == created.ID || rotated.Key == "" || rotated.Key == created.Key || len(rotated.KeySuffix) != 6 || + !strings.HasSuffix(rotated.Key, rotated.KeySuffix) || rotated.KeySuffix == created.KeySuffix || rotated.DailySpendMicros != created.DailySpendMicros || rotated.RequestsPerMinute != created.RequestsPerMinute || len(rotated.AllowedModels) != 1 || rotated.AllowedModels[0] != model.PublicID { + t.Fatalf("rotated key did not preserve controls: old=%+v new=%+v", created.APIKey, rotated.APIKey) + } + snapshot, err = store.LoadSnapshot(ctx) + if err != nil { + t.Fatal(err) + } + if len(snapshot.APIKeys) != 1 || snapshot.APIKeys[0].Principal.KeyID != rotated.ID { + t.Fatalf("rotation snapshot = %+v, want only new key %s", snapshot.APIKeys, rotated.ID) + } + keys, err = store.ListAPIKeysFor(ctx, tenant.ID) + if err != nil || len(keys) != 2 || keys[0].Status != "active" || keys[1].Status != "revoked" { + t.Fatalf("rotation lifecycle rows are incorrect: keys=%+v err=%v", keys, err) + } } func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) { @@ -192,10 +236,21 @@ func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) { t.Fatal(err) } defer scopedDB.Close() - var tenantID, providerID, modelID string + var tenantID, projectID, providerID, modelID string if err := scopedDB.QueryRow(ctx, `INSERT INTO tenants (slug,name) VALUES ('preference-test','Preference Test') RETURNING id::text`).Scan(&tenantID); err != nil { t.Fatal(err) } + if err := scopedDB.QueryRow(ctx, `INSERT INTO projects (tenant_id,slug,name) VALUES ($1,'default','Default') RETURNING id::text`, tenantID).Scan(&projectID); err != nil { + t.Fatal(err) + } + if _, err := scopedDB.Exec(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_hash) + VALUES ($1,$2,'Legacy key','sk-aigw-legacy...',decode(repeat('01',32),'hex'))`, tenantID, projectID); err != nil { + t.Fatalf("legacy key without suffix did not retain its empty default: %v", err) + } + if _, err := scopedDB.Exec(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_suffix,key_hash) + VALUES ($1,$2,'Invalid display suffix','sk-aigw-invalid...','too-long',decode(repeat('02',32),'hex'))`, tenantID, projectID); err == nil { + t.Fatal("invalid API key display suffix was accepted") + } if err := scopedDB.QueryRow(ctx, `INSERT INTO providers (slug,name,protocol,wire_api,base_url,api_key_ciphertext) VALUES ('preference-provider','Preference provider','openai','responses','https://example.invalid',decode('00','hex')) RETURNING id::text`).Scan(&providerID); err != nil { t.Fatal(err) @@ -210,7 +265,8 @@ func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) { t.Fatal(err) } store := &Store{db: scopedDB} - prefs, err := store.SetDeveloperPreferences(ctx, SetDeveloperPreferencesInput{TenantID: tenantID, DefaultModel: "model/preference-test"}) + preferenceDefaults := BillingPreferenceDefaults{LowBalanceThresholdMicros: 5_000_000, SpendAnomalyMultiplier: 3, SpendAnomalyMinMicros: 10_000_000} + prefs, err := store.SetDeveloperPreferences(ctx, SetDeveloperPreferencesInput{TenantID: tenantID, DefaultModel: "model/preference-test"}, preferenceDefaults) if err != nil { t.Fatal(err) } @@ -219,15 +275,21 @@ func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(t *testing.T) { } enabled := false threshold := int64(9_750_000) + anomalyEnabled := false + anomalyMultiplier := int64(7) + anomalyMinimum := int64(8_250_000) if _, err := store.SetBillingPreferences(ctx, SetBillingPreferencesInput{TenantID: tenantID, - LowBalanceEnabled: &enabled, LowBalanceThresholdMicros: &threshold}, 5_000_000); err != nil { + LowBalanceEnabled: &enabled, LowBalanceThresholdMicros: &threshold, SpendAnomalyEnabled: &anomalyEnabled, + SpendAnomalyMultiplier: &anomalyMultiplier, SpendAnomalyMinMicros: &anomalyMinimum}, preferenceDefaults); err != nil { t.Fatal(err) } - prefs, err = store.GetTenantPreferences(ctx, tenantID, 5_000_000) + prefs, err = store.GetTenantPreferences(ctx, tenantID, preferenceDefaults) if err != nil { t.Fatal(err) } - if prefs.LowBalanceEnabled || prefs.LowBalanceThresholdMicros != threshold || prefs.DefaultModel != "model/preference-test" { + if prefs.LowBalanceEnabled || prefs.LowBalanceThresholdMicros != threshold || prefs.SpendAnomalyEnabled || + prefs.SpendAnomalyMultiplier != anomalyMultiplier || prefs.SpendAnomalyMinMicros != anomalyMinimum || + prefs.DefaultModel != "model/preference-test" { t.Fatalf("preferences did not round trip: %+v", prefs) } } diff --git a/internal/controlplane/types.go b/internal/controlplane/types.go index f6402d8..995fc8d 100644 --- a/internal/controlplane/types.go +++ b/internal/controlplane/types.go @@ -36,13 +36,20 @@ type APIKey struct { ProjectID string `json:"project_id"` Name string `json:"name"` KeyPrefix string `json:"key_prefix"` + KeySuffix string `json:"key_suffix"` Scopes []string `json:"scopes"` Tags []string `json:"tags"` AllowedModels []string `json:"allowed_models"` MonthlySpendMicros int64 `json:"monthly_spend_micros"` + DailySpendMicros int64 `json:"daily_spend_micros"` + RequestsPerMinute int64 `json:"requests_per_minute"` + TokensPerMinute int64 `json:"tokens_per_minute"` CurrentMonthSpendMicros int64 `json:"current_month_spend_micros"` CurrentMonthReservedMicros int64 `json:"current_month_reserved_micros"` CurrentMonthRequests int64 `json:"current_month_requests"` + CurrentDaySpendMicros int64 `json:"current_day_spend_micros"` + CurrentDayReservedMicros int64 `json:"current_day_reserved_micros"` + CurrentDayRequests int64 `json:"current_day_requests"` Status string `json:"status"` ExpiresAt *time.Time `json:"expires_at,omitempty"` LastUsedAt *time.Time `json:"last_used_at,omitempty"` @@ -117,11 +124,17 @@ type DeveloperProviderHealth struct { WireAPI string `json:"wire_api"` State string `json:"state"` Attempts uint64 `json:"attempts"` + ActiveProbes uint64 `json:"active_probes"` RecentSamples int `json:"recent_samples"` AvailabilityPercent float64 `json:"availability_percent"` HeaderLatencyEWMA int64 `json:"header_latency_ewma_ms"` + TTFTSamples uint64 `json:"ttft_samples"` + TTFTEWMA int64 `json:"ttft_ewma_ms"` + SharedAttempts uint64 `json:"shared_attempts"` + SharedTTFTSamples uint64 `json:"shared_ttft_samples"` ConsecutiveFailures uint64 `json:"consecutive_failures"` LastObservedAt *time.Time `json:"last_observed_at,omitempty"` + LastProbeAt *time.Time `json:"last_probe_at,omitempty"` CircuitOpenUntil *time.Time `json:"circuit_open_until,omitempty"` } @@ -160,9 +173,18 @@ type TenantPreferences struct { FallbackModel string `json:"fallback_model,omitempty"` LowBalanceEnabled bool `json:"low_balance_enabled"` LowBalanceThresholdMicros int64 `json:"low_balance_threshold_micros"` + SpendAnomalyEnabled bool `json:"spend_anomaly_enabled"` + SpendAnomalyMultiplier int64 `json:"spend_anomaly_multiplier"` + SpendAnomalyMinMicros int64 `json:"spend_anomaly_min_micros"` UpdatedAt *time.Time `json:"updated_at,omitempty"` } +type BillingPreferenceDefaults struct { + LowBalanceThresholdMicros int64 + SpendAnomalyMultiplier int64 + SpendAnomalyMinMicros int64 +} + type SetDeveloperPreferencesInput struct { TenantID string `json:"tenant_id"` DefaultModel string `json:"default_model"` @@ -173,6 +195,9 @@ type SetBillingPreferencesInput struct { TenantID string `json:"tenant_id"` LowBalanceEnabled *bool `json:"low_balance_enabled"` LowBalanceThresholdMicros *int64 `json:"low_balance_threshold_micros"` + SpendAnomalyEnabled *bool `json:"spend_anomaly_enabled"` + SpendAnomalyMultiplier *int64 `json:"spend_anomaly_multiplier"` + SpendAnomalyMinMicros *int64 `json:"spend_anomaly_min_micros"` } type Model struct { @@ -260,6 +285,9 @@ type CreateAPIKeyInput struct { Tags []string `json:"tags"` AllowedModels []string `json:"allowed_models"` MonthlySpendMicros int64 `json:"monthly_spend_micros"` + DailySpendMicros int64 `json:"daily_spend_micros"` + RequestsPerMinute int64 `json:"requests_per_minute"` + TokensPerMinute int64 `json:"tokens_per_minute"` ExpiresAt *time.Time `json:"expires_at"` } @@ -469,6 +497,7 @@ type UsageRecord struct { Attempts int `json:"attempts"` StartedAt time.Time `json:"started_at"` DurationMS int64 `json:"duration_ms"` + TTFTMS int64 `json:"ttft_ms"` InputTokens int64 `json:"input_tokens"` OutputTokens int64 `json:"output_tokens"` TotalTokens int64 `json:"total_tokens"` @@ -481,6 +510,11 @@ type UsageRecord struct { MeteringStatus string `json:"metering_status"` } +type UsagePage struct { + Data []UsageRecord `json:"data"` + NextCursor string `json:"next_cursor,omitempty"` +} + type UsageDailyPoint struct { Day time.Time `json:"day"` RequestCount int64 `json:"request_count"` @@ -491,7 +525,10 @@ type UsageDailyPoint struct { ChargedMicros int64 `json:"charged_micros"` UncollectedMicros int64 `json:"uncollected_micros"` AverageDurationMS int64 `json:"average_duration_ms"` + P50DurationMS int64 `json:"p50_duration_ms"` P95DurationMS int64 `json:"p95_duration_ms"` + P50TTFTMS int64 `json:"p50_ttft_ms"` + P95TTFTMS int64 `json:"p95_ttft_ms"` } // UsageAnalytics is a persisted-ledger aggregation used by the developer @@ -501,6 +538,7 @@ type UsageAnalytics struct { RangeEnd time.Time `json:"range_end"` Models []UsageModelAnalytics `json:"models"` Providers []UsageProviderAnalytics `json:"providers"` + Keys []UsageKeyAnalytics `json:"keys"` } type UsageModelAnalytics struct { @@ -518,7 +556,10 @@ type UsageModelAnalytics struct { UncollectedMicros int64 `json:"uncollected_micros"` MissingUsageRequests int64 `json:"missing_usage_requests"` AverageDurationMS int64 `json:"average_duration_ms"` + P50DurationMS int64 `json:"p50_duration_ms"` P95DurationMS int64 `json:"p95_duration_ms"` + P50TTFTMS int64 `json:"p50_ttft_ms"` + P95TTFTMS int64 `json:"p95_ttft_ms"` PreviousChargedMicros int64 `json:"previous_charged_micros"` ChargeChangePercent *float64 `json:"charge_change_percent,omitempty"` } @@ -540,11 +581,34 @@ type UsageProviderAnalytics struct { UncollectedMicros int64 `json:"uncollected_micros"` MissingUsageRequests int64 `json:"missing_usage_requests"` AverageDurationMS int64 `json:"average_duration_ms"` + P50DurationMS int64 `json:"p50_duration_ms"` P95DurationMS int64 `json:"p95_duration_ms"` + P50TTFTMS int64 `json:"p50_ttft_ms"` + P95TTFTMS int64 `json:"p95_ttft_ms"` PreviousChargedMicros int64 `json:"previous_charged_micros"` ChargeChangePercent *float64 `json:"charge_change_percent,omitempty"` } +type UsageKeyAnalytics struct { + KeyID string `json:"key_id"` + KeyName string `json:"key_name"` + RequestCount int64 `json:"request_count"` + SuccessfulRequests int64 `json:"successful_requests"` + ErrorCount int64 `json:"error_count"` + ModelCount int64 `json:"model_count"` + InputTokens int64 `json:"input_tokens"` + OutputTokens int64 `json:"output_tokens"` + TotalTokens int64 `json:"total_tokens"` + ChargedMicros int64 `json:"charged_micros"` + UncollectedMicros int64 `json:"uncollected_micros"` + MissingUsageRequests int64 `json:"missing_usage_requests"` + AverageDurationMS int64 `json:"average_duration_ms"` + P50DurationMS int64 `json:"p50_duration_ms"` + P95DurationMS int64 `json:"p95_duration_ms"` + P50TTFTMS int64 `json:"p50_ttft_ms"` + P95TTFTMS int64 `json:"p95_ttft_ms"` +} + type UsageSummary struct { PeriodStart time.Time `json:"period_start"` TenantID string `json:"tenant_id"` diff --git a/internal/controlplane/usage.go b/internal/controlplane/usage.go index ebabaf6..05bc155 100644 --- a/internal/controlplane/usage.go +++ b/internal/controlplane/usage.go @@ -2,6 +2,9 @@ package controlplane import ( "context" + "encoding/base64" + "encoding/json" + "errors" "fmt" "strings" "time" @@ -25,6 +28,14 @@ type UsageQuery struct { From time.Time To time.Time Limit int + Cursor string +} + +var ErrInvalidUsageCursor = errors.New("invalid usage cursor") + +type usageCursor struct { + StartedAt time.Time `json:"started_at"` + RequestID string `json:"request_id"` } func (s *Store) RecordUsage(ctx context.Context, event domain.UsageEvent) error { @@ -37,13 +48,13 @@ func (s *Store) RecordUsage(ctx context.Context, event domain.UsageEvent) error INSERT INTO usage_events ( request_id, tenant_id, project_id, key_id, public_model, provider_id, upstream_model, protocol, stream, status_code, success, error_type, attempts, started_at, duration_ms, - input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, + ttft_ms, input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, cost_micros, charged_micros, uncollected_micros, usage_reported, metering_status) - VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,0,0,0,$21,$22) + VALUES ($1,$2,$3,$4,$5,NULLIF($6,''),NULLIF($7,''),$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,$21,0,0,0,$22,$23) ON CONFLICT (request_id) DO NOTHING`, event.RequestID, event.TenantID, event.ProjectID, event.KeyID, event.PublicModel, event.ProviderID, event.UpstreamModel, string(event.Protocol), event.Stream, event.StatusCode, event.Success, event.ErrorType, event.Attempts, event.StartedAt, event.DurationMS, - event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, + event.TTFTMS, event.Usage.InputTokens, event.Usage.OutputTokens, event.Usage.TotalTokens, event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens, event.UsageReported, usageMeteringStatus(event)) if err != nil { return fmt.Errorf("persist usage event: %w", err) @@ -104,13 +115,13 @@ func boolInt(value bool) int { return 0 } -func (s *Store) ListUsage(ctx context.Context, query UsageQuery) ([]UsageRecord, error) { +func (s *Store) ListUsage(ctx context.Context, query UsageQuery) (UsagePage, error) { limit := query.Limit if limit < 1 || limit > 1000 { limit = 200 } where := []string{"1=1"} - args := make([]any, 0, 13) + args := make([]any, 0, 16) index := 1 for _, item := range []struct{ value, clause string }{{query.TenantID, "tenant_id=$"}, {query.ProjectID, "project_id=$"}, {query.KeyID, "key_id=$"}, {query.Model, "public_model=$"}, {query.RequestID, "request_id=$"}} { if strings.TrimSpace(item.value) != "" { @@ -151,31 +162,68 @@ func (s *Store) ListUsage(ctx context.Context, query UsageQuery) ([]UsageRecord, args = append(args, query.To) index++ } - args = append(args, limit) + if strings.TrimSpace(query.Cursor) != "" { + cursor, err := decodeUsageCursor(query.Cursor) + if err != nil { + return UsagePage{}, err + } + where = append(where, "(started_at, request_id) < ($"+fmt.Sprint(index)+",$"+fmt.Sprint(index+1)+")") + args = append(args, cursor.StartedAt, cursor.RequestID) + index += 2 + } + args = append(args, limit+1) rows, err := s.db.Query(ctx, `SELECT request_id, tenant_id::text, project_id::text, COALESCE((SELECT name FROM projects p WHERE p.id=usage_events.project_id),''), key_id::text, COALESCE((SELECT name FROM api_keys k WHERE k.id=usage_events.key_id),''), public_model, COALESCE(provider_id,''), COALESCE((SELECT name FROM providers p WHERE p.id::text=usage_events.provider_id),''), COALESCE(upstream_model,''), protocol, stream, status_code, success, error_type, - attempts, started_at, duration_ms, input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, + attempts, started_at, duration_ms, ttft_ms, input_tokens, output_tokens, total_tokens, cache_creation_input_tokens, cache_read_input_tokens, cost_micros, charged_micros, uncollected_micros, usage_reported, metering_status FROM usage_events WHERE `+ - strings.Join(where, " AND ")+` ORDER BY created_at DESC LIMIT $`+fmt.Sprint(index), args...) + strings.Join(where, " AND ")+` ORDER BY started_at DESC, request_id DESC LIMIT $`+fmt.Sprint(index), args...) if err != nil { - return nil, fmt.Errorf("query usage events: %w", err) + return UsagePage{}, fmt.Errorf("query usage events: %w", err) } defer rows.Close() - result := make([]UsageRecord, 0) + result := make([]UsageRecord, 0, limit+1) for rows.Next() { var item UsageRecord if err := rows.Scan(&item.RequestID, &item.TenantID, &item.ProjectID, &item.ProjectName, &item.KeyID, &item.KeyName, &item.PublicModel, &item.ProviderID, &item.ProviderName, &item.UpstreamModel, &item.Protocol, &item.Stream, &item.StatusCode, &item.Success, &item.ErrorType, &item.Attempts, &item.StartedAt, &item.DurationMS, - &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.CacheCreationInputTokens, &item.CacheReadInputTokens, + &item.TTFTMS, &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.CacheCreationInputTokens, &item.CacheReadInputTokens, &item.CostMicros, &item.ChargedMicros, &item.UncollectedMicros, &item.UsageReported, &item.MeteringStatus); err != nil { - return nil, fmt.Errorf("scan usage event: %w", err) + return UsagePage{}, fmt.Errorf("scan usage event: %w", err) } result = append(result, item) } - return result, rows.Err() + if err := rows.Err(); err != nil { + return UsagePage{}, err + } + page := UsagePage{Data: result} + if len(result) > limit { + page.Data = result[:limit] + page.NextCursor = encodeUsageCursor(page.Data[len(page.Data)-1]) + } + return page, nil +} + +func encodeUsageCursor(record UsageRecord) string { + payload, _ := json.Marshal(usageCursor{StartedAt: record.StartedAt.UTC(), RequestID: record.RequestID}) + return base64.RawURLEncoding.EncodeToString(payload) +} + +func decodeUsageCursor(raw string) (usageCursor, error) { + if len(raw) > 2048 { + return usageCursor{}, ErrInvalidUsageCursor + } + payload, err := base64.RawURLEncoding.DecodeString(raw) + if err != nil { + return usageCursor{}, ErrInvalidUsageCursor + } + var cursor usageCursor + if json.Unmarshal(payload, &cursor) != nil || cursor.StartedAt.IsZero() || strings.TrimSpace(cursor.RequestID) == "" { + return usageCursor{}, ErrInvalidUsageCursor + } + return cursor, nil } func (s *Store) UsageDaily(ctx context.Context, query UsageQuery) ([]UsageDailyPoint, error) { @@ -225,7 +273,10 @@ func (s *Store) UsageDaily(ctx context.Context, query UsageQuery) ([]UsageDailyP count(*) FILTER (WHERE success), COALESCE(sum(input_tokens),0), COALESCE(sum(output_tokens),0), COALESCE(sum(total_tokens),0), COALESCE(sum(charged_micros),0), COALESCE(sum(uncollected_micros),0), COALESCE(round(avg(duration_ms)),0)::bigint, - COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY duration_ms)),0)::bigint + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY ttft_ms) FILTER (WHERE ttft_ms > 0)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY ttft_ms) FILTER (WHERE ttft_ms > 0)),0)::bigint FROM usage_events WHERE `+strings.Join(where, " AND ")+` GROUP BY 1 ORDER BY 1`, args...) if err != nil { return nil, fmt.Errorf("query daily usage: %w", err) @@ -235,7 +286,8 @@ func (s *Store) UsageDaily(ctx context.Context, query UsageQuery) ([]UsageDailyP for rows.Next() { var item UsageDailyPoint if err := rows.Scan(&item.Day, &item.RequestCount, &item.SuccessfulRequests, &item.InputTokens, &item.OutputTokens, - &item.TotalTokens, &item.ChargedMicros, &item.UncollectedMicros, &item.AverageDurationMS, &item.P95DurationMS); err != nil { + &item.TotalTokens, &item.ChargedMicros, &item.UncollectedMicros, &item.AverageDurationMS, &item.P50DurationMS, + &item.P95DurationMS, &item.P50TTFTMS, &item.P95TTFTMS); err != nil { return nil, fmt.Errorf("scan daily usage: %w", err) } result = append(result, item) diff --git a/internal/controlplane/usage_analytics.go b/internal/controlplane/usage_analytics.go index 7cc042f..1bc6228 100644 --- a/internal/controlplane/usage_analytics.go +++ b/internal/controlplane/usage_analytics.go @@ -22,7 +22,13 @@ func (s *Store) UsageAnalytics(ctx context.Context, query UsageQuery) (UsageAnal if !to.After(from) { return UsageAnalytics{}, fmt.Errorf("usage analytics range must be positive") } - result := UsageAnalytics{RangeStart: from, RangeEnd: to, Models: make([]UsageModelAnalytics, 0), Providers: make([]UsageProviderAnalytics, 0)} + result := UsageAnalytics{ + RangeStart: from, + RangeEnd: to, + Models: make([]UsageModelAnalytics, 0), + Providers: make([]UsageProviderAnalytics, 0), + Keys: make([]UsageKeyAnalytics, 0), + } modelPrevious, err := s.usageModelCharges(ctx, query, from.Add(-to.Sub(from)), from) if err != nil { @@ -49,8 +55,13 @@ func (s *Store) UsageAnalytics(ctx context.Context, query UsageQuery) (UsageAnal providers[index].PreviousChargedMicros = providerPrevious[providers[index].ProviderID] providers[index].ChargeChangePercent = chargeChange(providers[index].ChargedMicros, providers[index].PreviousChargedMicros) } + keys, err := s.usageKeyAnalytics(ctx, query, from, to) + if err != nil { + return UsageAnalytics{}, err + } result.Models = models result.Providers = providers + result.Keys = keys return result, nil } @@ -128,7 +139,11 @@ func (s *Store) usageModelAnalytics(ctx context.Context, query UsageQuery, from, count(DISTINCT NULLIF(e.provider_id,'')), COALESCE(sum(e.input_tokens),0), COALESCE(sum(e.output_tokens),0), COALESCE(sum(e.total_tokens),0), COALESCE(sum(e.cache_read_input_tokens),0), COALESCE(sum(e.cache_creation_input_tokens),0), COALESCE(sum(e.charged_micros),0), COALESCE(sum(e.uncollected_micros),0), count(*) FILTER (WHERE e.metering_status='missing'), - COALESCE(round(avg(e.duration_ms)),0)::bigint, COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint + COALESCE(round(avg(e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY e.ttft_ms) FILTER (WHERE e.ttft_ms > 0)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.ttft_ms) FILTER (WHERE e.ttft_ms > 0)),0)::bigint FROM usage_events e WHERE `+where+` GROUP BY e.public_model ORDER BY sum(e.charged_micros) DESC, e.public_model`, args...) if err != nil { return nil, fmt.Errorf("query usage model analytics: %w", err) @@ -139,7 +154,8 @@ func (s *Store) usageModelAnalytics(ctx context.Context, query UsageQuery, from, var item UsageModelAnalytics if err := rows.Scan(&item.PublicModel, &item.RequestCount, &item.SuccessfulRequests, &item.ErrorCount, &item.ProviderCount, &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.CacheReadInputTokens, &item.CacheCreationInputTokens, - &item.ChargedMicros, &item.UncollectedMicros, &item.MissingUsageRequests, &item.AverageDurationMS, &item.P95DurationMS); err != nil { + &item.ChargedMicros, &item.UncollectedMicros, &item.MissingUsageRequests, &item.AverageDurationMS, + &item.P50DurationMS, &item.P95DurationMS, &item.P50TTFTMS, &item.P95TTFTMS); err != nil { return nil, fmt.Errorf("scan usage model analytics: %w", err) } result = append(result, item) @@ -174,7 +190,10 @@ func (s *Store) usageProviderAnalytics(ctx context.Context, query UsageQuery, fr COALESCE(sum(e.input_tokens),0), COALESCE(sum(e.output_tokens),0), COALESCE(sum(e.total_tokens),0), COALESCE(sum(e.cache_read_input_tokens),0), COALESCE(sum(e.cache_creation_input_tokens),0), COALESCE(sum(e.charged_micros),0), COALESCE(sum(e.uncollected_micros),0), count(*) FILTER (WHERE e.metering_status='missing'), COALESCE(round(avg(e.duration_ms)),0)::bigint, - COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY e.ttft_ms) FILTER (WHERE e.ttft_ms > 0)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.ttft_ms) FILTER (WHERE e.ttft_ms > 0)),0)::bigint FROM usage_events e LEFT JOIN providers p ON p.id::text=e.provider_id WHERE `+where+` GROUP BY e.provider_id, p.name, p.wire_api ORDER BY sum(e.charged_micros) DESC, COALESCE(NULLIF(p.name,''),'Unassigned')`, args...) if err != nil { @@ -186,7 +205,8 @@ func (s *Store) usageProviderAnalytics(ctx context.Context, query UsageQuery, fr var item UsageProviderAnalytics if err := rows.Scan(&item.ProviderID, &item.ProviderName, &item.WireAPI, &item.RequestCount, &item.SuccessfulRequests, &item.ErrorCount, &item.ModelCount, &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.CacheReadInputTokens, &item.CacheCreationInputTokens, - &item.ChargedMicros, &item.UncollectedMicros, &item.MissingUsageRequests, &item.AverageDurationMS, &item.P95DurationMS); err != nil { + &item.ChargedMicros, &item.UncollectedMicros, &item.MissingUsageRequests, &item.AverageDurationMS, + &item.P50DurationMS, &item.P95DurationMS, &item.P50TTFTMS, &item.P95TTFTMS); err != nil { return nil, fmt.Errorf("scan usage provider analytics: %w", err) } result = append(result, item) @@ -194,6 +214,37 @@ func (s *Store) usageProviderAnalytics(ctx context.Context, query UsageQuery, fr return result, rows.Err() } +func (s *Store) usageKeyAnalytics(ctx context.Context, query UsageQuery, from, to time.Time) ([]UsageKeyAnalytics, error) { + where, args := analyticsUsageWhere(query, from, to) + rows, err := s.db.Query(ctx, `SELECT e.key_id::text, COALESCE(NULLIF(k.name,''),'Deleted key'), + count(*), count(*) FILTER (WHERE e.success), count(*) FILTER (WHERE NOT e.success), count(DISTINCT e.public_model), + COALESCE(sum(e.input_tokens),0), COALESCE(sum(e.output_tokens),0), COALESCE(sum(e.total_tokens),0), + COALESCE(sum(e.charged_micros),0), COALESCE(sum(e.uncollected_micros),0), count(*) FILTER (WHERE e.metering_status='missing'), + COALESCE(round(avg(e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.duration_ms)),0)::bigint, + COALESCE(round(percentile_cont(0.50) WITHIN GROUP (ORDER BY e.ttft_ms) FILTER (WHERE e.ttft_ms > 0)),0)::bigint, + COALESCE(round(percentile_cont(0.95) WITHIN GROUP (ORDER BY e.ttft_ms) FILTER (WHERE e.ttft_ms > 0)),0)::bigint + FROM usage_events e LEFT JOIN api_keys k ON k.id=e.key_id WHERE `+where+` + GROUP BY e.key_id, k.name ORDER BY sum(e.charged_micros) DESC, COALESCE(NULLIF(k.name,''),'Deleted key')`, args...) + if err != nil { + return nil, fmt.Errorf("query usage API key analytics: %w", err) + } + defer rows.Close() + result := make([]UsageKeyAnalytics, 0) + for rows.Next() { + var item UsageKeyAnalytics + if err := rows.Scan(&item.KeyID, &item.KeyName, &item.RequestCount, &item.SuccessfulRequests, &item.ErrorCount, + &item.ModelCount, &item.InputTokens, &item.OutputTokens, &item.TotalTokens, &item.ChargedMicros, + &item.UncollectedMicros, &item.MissingUsageRequests, &item.AverageDurationMS, &item.P50DurationMS, + &item.P95DurationMS, &item.P50TTFTMS, &item.P95TTFTMS); err != nil { + return nil, fmt.Errorf("scan usage API key analytics: %w", err) + } + result = append(result, item) + } + return result, rows.Err() +} + func (s *Store) usageProviderCharges(ctx context.Context, query UsageQuery, from, to time.Time) (map[string]int64, error) { where, args := analyticsUsageWhere(query, from, to) rows, err := s.db.Query(ctx, `SELECT COALESCE(e.provider_id,''), COALESCE(sum(e.charged_micros),0) diff --git a/internal/controlplane/usage_integration_test.go b/internal/controlplane/usage_integration_test.go index 2545c1d..10305bf 100644 --- a/internal/controlplane/usage_integration_test.go +++ b/internal/controlplane/usage_integration_test.go @@ -78,7 +78,7 @@ func TestUsageFiltersAndDailyAggregationPostgres(t *testing.T) { started := time.Now().UTC().Add(-time.Hour).Truncate(time.Second) events := []domain.UsageEvent{ - {RequestID: fmt.Sprintf("req_usage_ok_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started, DurationMS: 120, Usage: domain.Usage{InputTokens: 10, OutputTokens: 4, TotalTokens: 14, CacheReadInputTokens: 5}}, + {RequestID: fmt.Sprintf("req_usage_ok_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started, DurationMS: 120, TTFTMS: 47, Usage: domain.Usage{InputTokens: 10, OutputTokens: 4, TotalTokens: 14, CacheReadInputTokens: 5}}, {RequestID: fmt.Sprintf("req_usage_error_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAIResponses, Stream: true, StatusCode: 502, Success: false, ErrorType: "provider_error", StartedAt: started.Add(time.Minute), DurationMS: 350}, } for _, event := range events { @@ -94,7 +94,7 @@ func TestUsageFiltersAndDailyAggregationPostgres(t *testing.T) { if err != nil { t.Fatal(err) } - if len(records) != 1 || records[0].RequestID != events[0].RequestID || records[0].ProjectName != "Production" || records[0].KeyName != "Production key" || records[0].MeteringStatus != "reported" { + if len(records.Data) != 1 || records.Data[0].RequestID != events[0].RequestID || records.Data[0].ProjectName != "Production" || records.Data[0].KeyName != "Production key" || records.Data[0].MeteringStatus != "reported" || records.Data[0].TTFTMS != 47 { t.Fatalf("unexpected filtered usage: %+v", records) } streaming := true @@ -103,7 +103,7 @@ func TestUsageFiltersAndDailyAggregationPostgres(t *testing.T) { if err != nil { t.Fatal(err) } - if len(records) != 1 || records[0].RequestID != events[1].RequestID { + if len(records.Data) != 1 || records.Data[0].RequestID != events[1].RequestID { t.Fatalf("unexpected provider/protocol/transport usage filter: %+v", records) } filteredPoints, err := store.UsageDaily(ctx, UsageQuery{TenantID: tenantID, Provider: providerSlug, Protocol: string(domain.ProtocolOpenAIResponses), ErrorType: "provider_error", Stream: &streaming, From: started.Add(-time.Minute), To: started.Add(time.Hour)}) @@ -117,11 +117,11 @@ func TestUsageFiltersAndDailyAggregationPostgres(t *testing.T) { if err != nil { t.Fatal(err) } - if len(points) != 1 || points[0].RequestCount != 2 || points[0].SuccessfulRequests != 1 || points[0].TotalTokens != 14 || points[0].ChargedMicros != 125 || points[0].P95DurationMS < 120 { + if len(points) != 1 || points[0].RequestCount != 2 || points[0].SuccessfulRequests != 1 || points[0].TotalTokens != 14 || points[0].ChargedMicros != 125 || points[0].P50DurationMS != 235 || points[0].P95DurationMS < 120 || points[0].P50TTFTMS != 47 || points[0].P95TTFTMS != 47 { t.Fatalf("unexpected daily usage: %+v", points) } - previous := domain.UsageEvent{RequestID: fmt.Sprintf("req_usage_previous_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started.Add(-5 * time.Minute), DurationMS: 80, Usage: domain.Usage{InputTokens: 3, OutputTokens: 2, TotalTokens: 5}} + previous := domain.UsageEvent{RequestID: fmt.Sprintf("req_usage_previous_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started.Add(-5 * time.Minute), DurationMS: 80, TTFTMS: 30, Usage: domain.Usage{InputTokens: 3, OutputTokens: 2, TotalTokens: 5}} missing := domain.UsageEvent{RequestID: fmt.Sprintf("req_usage_missing_%d", suffix), TenantID: tenantID, ProjectID: projectID, KeyID: keyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, StartedAt: started.Add(2 * time.Minute), DurationMS: 200} otherTenantEvent := domain.UsageEvent{RequestID: fmt.Sprintf("req_usage_other_%d", suffix), TenantID: otherTenantID, ProjectID: otherProjectID, KeyID: otherKeyID, PublicModel: "openai/test", ProviderID: providerID, Protocol: domain.ProtocolOpenAI, StatusCode: 200, Success: true, UsageReported: true, StartedAt: started.Add(3 * time.Minute), DurationMS: 900, Usage: domain.Usage{InputTokens: 1000, OutputTokens: 1000, TotalTokens: 2000}} for _, event := range []domain.UsageEvent{previous, missing, otherTenantEvent} { @@ -129,6 +129,34 @@ func TestUsageFiltersAndDailyAggregationPostgres(t *testing.T) { t.Fatal(err) } } + requestDetail, err := store.ListUsage(ctx, UsageQuery{TenantID: tenantID, RequestID: events[0].RequestID, Limit: 1}) + if err != nil { + t.Fatal(err) + } + if len(requestDetail.Data) != 1 || requestDetail.Data[0].RequestID != events[0].RequestID { + t.Fatalf("tenant could not retrieve its ledger-linked request: %+v", requestDetail) + } + crossTenantDetail, err := store.ListUsage(ctx, UsageQuery{TenantID: tenantID, RequestID: otherTenantEvent.RequestID, Limit: 1}) + if err != nil { + t.Fatal(err) + } + if len(crossTenantDetail.Data) != 0 { + t.Fatalf("tenant retrieved another tenant's request by request id: %+v", crossTenantDetail) + } + firstPage, err := store.ListUsage(ctx, UsageQuery{TenantID: tenantID, From: started.Add(-time.Hour), To: started.Add(10 * time.Minute), Limit: 1}) + if err != nil { + t.Fatal(err) + } + if len(firstPage.Data) != 1 || firstPage.NextCursor == "" { + t.Fatalf("expected a cursor for the first usage page: %+v", firstPage) + } + secondPage, err := store.ListUsage(ctx, UsageQuery{TenantID: tenantID, From: started.Add(-time.Hour), To: started.Add(10 * time.Minute), Limit: 1, Cursor: firstPage.NextCursor}) + if err != nil { + t.Fatal(err) + } + if len(secondPage.Data) != 1 || secondPage.Data[0].RequestID == firstPage.Data[0].RequestID { + t.Fatalf("cursor did not advance usage page: first=%+v second=%+v", firstPage, secondPage) + } if _, err := db.Exec(ctx, `UPDATE usage_events SET charged_micros=25,cost_micros=25 WHERE request_id=$1`, previous.RequestID); err != nil { t.Fatal(err) } @@ -136,9 +164,12 @@ func TestUsageFiltersAndDailyAggregationPostgres(t *testing.T) { if err != nil { t.Fatal(err) } - if len(analytics.Models) != 1 || analytics.Models[0].RequestCount != 3 || analytics.Models[0].SuccessfulRequests != 2 || analytics.Models[0].ProviderCount != 1 || analytics.Models[0].CacheReadInputTokens != 5 || analytics.Models[0].MissingUsageRequests != 1 || analytics.Models[0].ChargedMicros != 125 || analytics.Models[0].PreviousChargedMicros != 25 || analytics.Models[0].ChargeChangePercent == nil || *analytics.Models[0].ChargeChangePercent != 400 { + if len(analytics.Models) != 1 || analytics.Models[0].RequestCount != 3 || analytics.Models[0].SuccessfulRequests != 2 || analytics.Models[0].ProviderCount != 1 || analytics.Models[0].CacheReadInputTokens != 5 || analytics.Models[0].MissingUsageRequests != 1 || analytics.Models[0].ChargedMicros != 125 || analytics.Models[0].P50DurationMS != 200 || analytics.Models[0].P95DurationMS < 200 || analytics.Models[0].P50TTFTMS != 47 || analytics.Models[0].P95TTFTMS != 47 || analytics.Models[0].PreviousChargedMicros != 25 || analytics.Models[0].ChargeChangePercent == nil || *analytics.Models[0].ChargeChangePercent != 400 { t.Fatalf("unexpected model analytics: %+v", analytics.Models) } + if len(analytics.Keys) != 1 || analytics.Keys[0].KeyID != keyID || analytics.Keys[0].RequestCount != 3 || analytics.Keys[0].ChargedMicros != 125 { + t.Fatalf("unexpected key analytics: %+v", analytics.Keys) + } if len(analytics.Providers) != 1 || analytics.Providers[0].ProviderID != providerID || analytics.Providers[0].WireAPI != "responses" || analytics.Providers[0].RequestCount != 3 || analytics.Providers[0].MissingUsageRequests != 1 || analytics.Providers[0].P95DurationMS < 200 { t.Fatalf("unexpected provider analytics: %+v", analytics.Providers) } diff --git a/internal/controlplane/usage_test.go b/internal/controlplane/usage_test.go new file mode 100644 index 0000000..d47ce16 --- /dev/null +++ b/internal/controlplane/usage_test.go @@ -0,0 +1,26 @@ +package controlplane + +import ( + "errors" + "testing" + "time" +) + +func TestUsageCursorRoundTrip(t *testing.T) { + record := UsageRecord{RequestID: "req_cursor_test", StartedAt: time.Date(2026, 8, 6, 2, 3, 4, 567, time.UTC)} + cursor, err := decodeUsageCursor(encodeUsageCursor(record)) + if err != nil { + t.Fatal(err) + } + if cursor.RequestID != record.RequestID || !cursor.StartedAt.Equal(record.StartedAt) { + t.Fatalf("cursor = %+v, want request %s at %v", cursor, record.RequestID, record.StartedAt) + } +} + +func TestUsageCursorRejectsInvalidInput(t *testing.T) { + for _, raw := range []string{"not-base64!", "e30", string(make([]byte, 2049))} { + if _, err := decodeUsageCursor(raw); !errors.Is(err, ErrInvalidUsageCursor) { + t.Fatalf("decodeUsageCursor(%q) error = %v", raw, err) + } + } +} |
