summaryrefslogtreecommitdiff
path: root/internal/controlplane
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-06 15:58:57 +1200
committerChia <Chia@93.nz>2026-08-06 15:58:57 +1200
commit3f702084d20b3c3a3ea916f3110e99b22bda60b3 (patch)
tree517f76c51025ce1ee085ea4898c60f799e5c37ea /internal/controlplane
parent41e322c53d7b4b796eb377d0df9c29ecd10ba431 (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 '')
-rw-r--r--internal/controlplane/api_key_test.go26
-rw-r--r--internal/controlplane/mail_operations.go16
-rw-r--r--internal/controlplane/mail_operations_integration_test.go144
-rw-r--r--internal/controlplane/manager.go38
-rw-r--r--internal/controlplane/manager_test.go49
-rw-r--r--internal/controlplane/mutations.go130
-rw-r--r--internal/controlplane/preferences.go75
-rw-r--r--internal/controlplane/preferences_test.go18
-rw-r--r--internal/controlplane/queries.go35
-rw-r--r--internal/controlplane/schema.sql40
-rw-r--r--internal/controlplane/snapshot.go9
-rw-r--r--internal/controlplane/store.go48
-rw-r--r--internal/controlplane/store_integration_test.go84
-rw-r--r--internal/controlplane/types.go64
-rw-r--r--internal/controlplane/usage.go82
-rw-r--r--internal/controlplane/usage_analytics.go61
-rw-r--r--internal/controlplane/usage_integration_test.go43
-rw-r--r--internal/controlplane/usage_test.go26
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, &currency, &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, &currency, &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(&notificationEvents); 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)
+ }
+ }
+}