summaryrefslogtreecommitdiff
path: root/internal/controlplane
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-05 22:01:29 +1200
committerChia <Chia@93.nz>2026-08-05 22:07:50 +1200
commiteadb2ffe85c43cf6fc741c9823cd28eedb4a844c (patch)
tree1aba2536d57360da403aa35c9ced58b615c7064e /internal/controlplane
parentcd0dd91ab93653631904f2ea0e574ccde6d60339 (diff)
feat: harden prepaid billing and commercial operations
Diffstat (limited to 'internal/controlplane')
-rw-r--r--internal/controlplane/mail_operations.go213
-rw-r--r--internal/controlplane/mail_operations_test.go29
-rw-r--r--internal/controlplane/manager.go44
-rw-r--r--internal/controlplane/mutations.go143
-rw-r--r--internal/controlplane/outbox.go19
-rw-r--r--internal/controlplane/queries.go73
-rw-r--r--internal/controlplane/retention.go57
-rw-r--r--internal/controlplane/rotation.go69
-rw-r--r--internal/controlplane/schema.sql265
-rw-r--r--internal/controlplane/snapshot.go108
-rw-r--r--internal/controlplane/store.go82
-rw-r--r--internal/controlplane/store_integration_test.go66
-rw-r--r--internal/controlplane/types.go72
-rw-r--r--internal/controlplane/usage.go16
14 files changed, 1198 insertions, 58 deletions
diff --git a/internal/controlplane/mail_operations.go b/internal/controlplane/mail_operations.go
new file mode 100644
index 0000000..aaaf2c5
--- /dev/null
+++ b/internal/controlplane/mail_operations.go
@@ -0,0 +1,213 @@
+package controlplane
+
+import (
+ "context"
+ "crypto/hmac"
+ "crypto/sha256"
+ "encoding/hex"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "log/slog"
+ "net/http"
+ "strconv"
+ "strings"
+ "time"
+
+ "github.com/jackc/pgx/v5"
+)
+
+const maxMailFeedbackBytes = 64 << 10
+
+type MailNotificationConfig struct {
+ LowBalanceMicros int64
+ SpendAnomalyMultiplier int64
+ SpendAnomalyMinMicros int64
+ Interval time.Duration
+}
+
+type mailFeedback struct {
+ EventID string `json:"event_id"`
+ EventType string `json:"event_type"`
+ Recipient string `json:"recipient"`
+ Provider string `json:"provider"`
+ Detail string `json:"detail"`
+}
+
+// MailFeedbackHandler accepts a provider-neutral normalized callback. An edge
+// adapter maps the provider's native event into this payload and signs
+// "<unix timestamp>.<raw body>" with HMAC-SHA256.
+func (s *Store) MailFeedbackHandler(secret string) http.Handler {
+ return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
+ if r.Method != http.MethodPost {
+ w.Header().Set("Allow", http.MethodPost)
+ http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
+ return
+ }
+ body, err := io.ReadAll(io.LimitReader(r.Body, maxMailFeedbackBytes+1))
+ if err != nil || len(body) > maxMailFeedbackBytes || !validMailFeedbackSignature(secret, r.Header.Get("X-AIGW-Mail-Timestamp"), r.Header.Get("X-AIGW-Mail-Signature"), body, time.Now().UTC()) {
+ http.Error(w, "invalid feedback signature", http.StatusBadRequest)
+ return
+ }
+ var event mailFeedback
+ decoder := json.NewDecoder(strings.NewReader(string(body)))
+ decoder.DisallowUnknownFields()
+ if decoder.Decode(&event) != nil || strings.TrimSpace(event.EventID) == "" ||
+ !strings.Contains(event.Recipient, "@") || (event.EventType != "delivered" && event.EventType != "bounce" && event.EventType != "complaint") {
+ http.Error(w, "invalid feedback event", http.StatusBadRequest)
+ return
+ }
+ if err := s.applyMailFeedback(r.Context(), event); err != nil {
+ http.Error(w, "feedback persistence failed", http.StatusInternalServerError)
+ return
+ }
+ w.Header().Set("Content-Type", "application/json")
+ _, _ = io.WriteString(w, "{\"received\":true}\n")
+ })
+}
+
+func validMailFeedbackSignature(secret, timestamp, signature string, body []byte, now time.Time) bool {
+ if secret == "" || timestamp == "" || !strings.HasPrefix(signature, "sha256=") {
+ return false
+ }
+ seconds, err := strconv.ParseInt(timestamp, 10, 64)
+ if err != nil || now.Sub(time.Unix(seconds, 0)).Abs() > 5*time.Minute {
+ return false
+ }
+ provided, err := hex.DecodeString(strings.TrimPrefix(signature, "sha256="))
+ if err != nil {
+ return false
+ }
+ mac := hmac.New(sha256.New, []byte(secret))
+ _, _ = mac.Write([]byte(timestamp))
+ _, _ = mac.Write([]byte("."))
+ _, _ = mac.Write(body)
+ return hmac.Equal(provided, mac.Sum(nil))
+}
+
+func (s *Store) applyMailFeedback(ctx context.Context, event mailFeedback) error {
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return err
+ }
+ defer tx.Rollback(ctx)
+ tag, err := tx.Exec(ctx, `INSERT INTO mail_feedback_events (event_id,event_type,recipient,provider)
+ VALUES ($1,$2,lower($3),$4) ON CONFLICT DO NOTHING`, event.EventID, event.EventType, event.Recipient, event.Provider)
+ if err != nil || tag.RowsAffected() == 0 {
+ if err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+ }
+ if event.EventType == "bounce" || event.EventType == "complaint" {
+ detail := event.Detail
+ if len(detail) > 1000 {
+ detail = detail[:1000]
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO mail_suppressions
+ (recipient,reason,provider,provider_event_id,detail) VALUES (lower($1),$2,$3,$4,$5)
+ ON CONFLICT (recipient) DO UPDATE SET reason=EXCLUDED.reason,provider=EXCLUDED.provider,
+ provider_event_id=EXCLUDED.provider_event_id,detail=EXCLUDED.detail,updated_at=now()`,
+ event.Recipient, event.EventType, event.Provider, event.EventID, detail); err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE console_mail_outbox SET status='suppressed',last_error=$2,claimed_at=NULL
+ WHERE lower(recipient)=lower($1) AND status IN ('pending','retry','sending')`, event.Recipient, "recipient suppressed after "+event.EventType); err != nil {
+ return err
+ }
+ }
+ return tx.Commit(ctx)
+}
+
+func (s *Store) RunMailNotificationWorker(ctx context.Context, config MailNotificationConfig, logger *slog.Logger) {
+ if config.Interval <= 0 {
+ config.Interval = 5 * time.Minute
+ }
+ if logger == nil {
+ logger = slog.Default()
+ }
+ ticker := time.NewTicker(config.Interval)
+ defer ticker.Stop()
+ for {
+ if err := s.queueBillingNotifications(ctx, config); err != nil && !errors.Is(err, context.Canceled) {
+ logger.Warn("billing_notification_scan_failed", "error", err)
+ }
+ select {
+ case <-ctx.Done():
+ return
+ case <-ticker.C:
+ }
+ }
+}
+
+func (s *Store) queueBillingNotifications(ctx context.Context, config MailNotificationConfig) error {
+ rows, err := s.db.Query(ctx, `WITH spend AS (
+ SELECT tenant_id,
+ COALESCE(sum(-amount_micros) FILTER (WHERE kind='usage' AND created_at>=date_trunc('day',now())),0)::bigint today,
+ (COALESCE(sum(-amount_micros) FILTER (WHERE kind='usage' AND created_at>=date_trunc('day',now())-interval '7 days' AND created_at<date_trunc('day',now())),0)/7)::bigint baseline
+ FROM billing_ledger GROUP BY tenant_id)
+ SELECT w.tenant_id::text,w.currency,w.balance_micros-w.reserved_micros,u.email,u.display_name,
+ COALESCE(spend.today,0),COALESCE(spend.baseline,0)
+ 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
+ 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))`)
+ if err != nil {
+ return err
+ }
+ defer rows.Close()
+ for rows.Next() {
+ var tenantID, currency, email, name string
+ var available, today, baseline int64
+ if err := rows.Scan(&tenantID, &currency, &available, &email, &name, &today, &baseline); err != nil {
+ return err
+ }
+ day := time.Now().UTC().Format("2006-01-02")
+ if available <= config.LowBalanceMicros {
+ body := fmt.Sprintf("Hi %s,\n\nYour AIGW prepaid balance is low: %.6f %s remains available. Add funds to avoid interrupted API access.\n", displayName(name), float64(available)/1_000_000, strings.ToUpper(currency))
+ if err := s.queueNotification(ctx, tenantID, email, "low_balance", day, "AIGW balance is low", body); err != nil {
+ return err
+ }
+ }
+ if baseline > 0 && today >= config.SpendAnomalyMinMicros && today >= baseline*config.SpendAnomalyMultiplier {
+ 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
+ }
+ }
+ }
+ return rows.Err()
+}
+
+func (s *Store) queueNotification(ctx context.Context, tenantID, recipient, kind, dedupe, subject, body string) error {
+ tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.ReadCommitted})
+ if err != nil {
+ return err
+ }
+ defer tx.Rollback(ctx)
+ tag, err := tx.Exec(ctx, `INSERT INTO mail_notification_events (tenant_id,recipient,notification_type,dedupe_key)
+ VALUES ($1,lower($2),$3,$4) ON CONFLICT DO NOTHING`, tenantID, recipient, kind, dedupe)
+ if err != nil || tag.RowsAffected() == 0 {
+ if err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+ }
+ ciphertext, err := s.cipher.Encrypt(body)
+ if err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO console_mail_outbox (recipient,template,subject,body_ciphertext)
+ VALUES (lower($1),$2,$3,$4)`, recipient, kind, subject, ciphertext); err != nil {
+ return err
+ }
+ return tx.Commit(ctx)
+}
+
+func displayName(value string) string {
+ if value = strings.TrimSpace(value); value != "" {
+ return value
+ }
+ return "there"
+}
diff --git a/internal/controlplane/mail_operations_test.go b/internal/controlplane/mail_operations_test.go
new file mode 100644
index 0000000..0b47cdb
--- /dev/null
+++ b/internal/controlplane/mail_operations_test.go
@@ -0,0 +1,29 @@
+package controlplane
+
+import (
+ "crypto/hmac"
+ "crypto/sha256"
+ "encoding/hex"
+ "strconv"
+ "testing"
+ "time"
+)
+
+func TestMailFeedbackSignature(t *testing.T) {
+ now := time.Unix(1_800_000_000, 0).UTC()
+ timestamp := strconv.FormatInt(now.Unix(), 10)
+ body := []byte(`{"event_id":"evt_1","event_type":"bounce","recipient":"test@example.com","provider":"test"}`)
+ mac := hmac.New(sha256.New, []byte("a-production-length-feedback-secret"))
+ _, _ = mac.Write([]byte(timestamp + "."))
+ _, _ = mac.Write(body)
+ signature := "sha256=" + hex.EncodeToString(mac.Sum(nil))
+ if !validMailFeedbackSignature("a-production-length-feedback-secret", timestamp, signature, body, now) {
+ t.Fatal("valid signature was rejected")
+ }
+ if validMailFeedbackSignature("a-production-length-feedback-secret", timestamp, signature, []byte(`{}`), now) {
+ t.Fatal("signature must bind the raw body")
+ }
+ if validMailFeedbackSignature("a-production-length-feedback-secret", timestamp, signature, body, now.Add(6*time.Minute)) {
+ t.Fatal("stale signature was accepted")
+ }
+}
diff --git a/internal/controlplane/manager.go b/internal/controlplane/manager.go
index b6be748..212963b 100644
--- a/internal/controlplane/manager.go
+++ b/internal/controlplane/manager.go
@@ -28,16 +28,19 @@ type policyReplacer interface {
}
type Manager struct {
- store managerStore
- catalog *catalog.Catalog
- authenticator *auth.StaticAuthenticator
- logger *slog.Logger
- pollInterval time.Duration
- generation atomic.Int64
- redisConnected atomic.Bool
- reloadMu sync.Mutex
- broadcasts chan ChangeEvent
- policyTarget policyReplacer
+ store managerStore
+ catalog *catalog.Catalog
+ authenticator *auth.StaticAuthenticator
+ logger *slog.Logger
+ pollInterval time.Duration
+ generation atomic.Int64
+ redisConnected atomic.Bool
+ loaded atomic.Bool
+ lastReloadUnix atomic.Int64
+ lastHealthyUnix atomic.Int64
+ reloadMu sync.Mutex
+ broadcasts chan ChangeEvent
+ policyTarget policyReplacer
}
func NewManager(store managerStore, modelCatalog *catalog.Catalog, authenticator *auth.StaticAuthenticator, logger *slog.Logger, pollInterval time.Duration, policyTargets ...policyReplacer) *Manager {
@@ -67,10 +70,30 @@ func (m *Manager) Reload(ctx context.Context) (int64, error) {
m.policyTarget.ReplacePolicies(snapshot.Limits)
}
m.generation.Store(snapshot.Generation)
+ m.loaded.Store(true)
+ m.lastReloadUnix.Store(time.Now().UTC().Unix())
+ m.lastHealthyUnix.Store(time.Now().UTC().Unix())
m.logger.Info("control_plane_reloaded", "generation", snapshot.Generation, "models", len(snapshot.Models), "api_keys", len(snapshot.APIKeys))
return snapshot.Generation, nil
}
+func (m *Manager) Loaded() bool { return m.loaded.Load() }
+func (m *Manager) LastReloadAt() time.Time {
+ value := m.lastReloadUnix.Load()
+ if value == 0 {
+ return time.Time{}
+ }
+ return time.Unix(value, 0).UTC()
+}
+
+func (m *Manager) LastHealthyAt() time.Time {
+ value := m.lastHealthyUnix.Load()
+ if value == 0 {
+ return time.Time{}
+ }
+ return time.Unix(value, 0).UTC()
+}
+
func (m *Manager) AfterMutation(ctx context.Context, generation int64, resource, id string) error {
loadedGeneration, err := m.Reload(ctx)
if err != nil {
@@ -135,6 +158,7 @@ func (m *Manager) runPolling(ctx context.Context) {
m.logger.Warn("control_plane_generation_check_failed", "error", err)
continue
}
+ m.lastHealthyUnix.Store(time.Now().UTC().Unix())
if generation > m.generation.Load() {
if _, err := m.Reload(ctx); err != nil {
m.logger.Error("control_plane_reload_failed", "source", "postgres", "error", err)
diff --git a/internal/controlplane/mutations.go b/internal/controlplane/mutations.go
index 3cedf70..c2cb7d8 100644
--- a/internal/controlplane/mutations.go
+++ b/internal/controlplane/mutations.go
@@ -11,6 +11,7 @@ import (
"net/url"
"regexp"
"strings"
+ "time"
"github.com/jackc/pgx/v5"
)
@@ -172,10 +173,45 @@ func (s *Store) SetProviderEnabled(ctx context.Context, id string, enabled bool)
func (s *Store) CreateModel(ctx context.Context, input CreateModelInput) (Model, int64, error) {
input.PublicID = strings.TrimSpace(input.PublicID)
+ input.DisplayName = strings.TrimSpace(input.DisplayName)
+ input.Description = strings.TrimSpace(input.Description)
input.OwnedBy = strings.TrimSpace(input.OwnedBy)
+ input.Lifecycle = strings.ToLower(strings.TrimSpace(input.Lifecycle))
+ input.PriceCurrency = strings.ToLower(strings.TrimSpace(input.PriceCurrency))
+ if input.DisplayName == "" {
+ input.DisplayName = input.PublicID
+ }
+ if input.Lifecycle == "" {
+ input.Lifecycle = "active"
+ }
+ if input.PriceCurrency == "" {
+ input.PriceCurrency = "usd"
+ }
+ if len(input.InputModalities) == 0 {
+ input.InputModalities = []string{"text"}
+ }
+ if len(input.OutputModalities) == 0 {
+ input.OutputModalities = []string{"text"}
+ }
+ if len(input.Capabilities) == 0 {
+ input.Capabilities = []string{"chat", "streaming"}
+ }
+ input.InputModalities = uniqueStrings(input.InputModalities)
+ input.OutputModalities = uniqueStrings(input.OutputModalities)
+ input.Capabilities = uniqueStrings(input.Capabilities)
+ input.Regions = uniqueStrings(input.Regions)
+ input.Aliases = uniqueStrings(input.Aliases)
+ input.AllowedTenantIDs = uniqueStrings(input.AllowedTenantIDs)
+ input.AllowedKeyIDs = uniqueStrings(input.AllowedKeyIDs)
if input.PublicID == "" || len(input.Routes) == 0 {
return Model{}, 0, errors.New("model requires public_id and at least one route")
}
+ if input.ContextWindow < 0 || input.MaxOutputTokens < 0 || len(input.PriceCurrency) != 3 {
+ return Model{}, 0, errors.New("model context, output limit, or price currency is invalid")
+ }
+ if input.Lifecycle != "preview" && input.Lifecycle != "active" && input.Lifecycle != "deprecated" && input.Lifecycle != "retired" {
+ return Model{}, 0, errors.New("model lifecycle must be preview, active, deprecated, or retired")
+ }
if input.InputPriceMicrosPerMillion < 0 || input.OutputPriceMicrosPerMillion < 0 || input.CacheReadPriceMicrosPerMillion < 0 || input.CacheWritePriceMicrosPerMillion < 0 {
return Model{}, 0, errors.New("model prices cannot be negative")
}
@@ -194,21 +230,63 @@ func (s *Store) CreateModel(ctx context.Context, input CreateModelInput) (Model,
return Model{}, 0, err
}
defer tx.Rollback(ctx)
+ inputModalitiesJSON, _ := json.Marshal(input.InputModalities)
+ outputModalitiesJSON, _ := json.Marshal(input.OutputModalities)
+ capabilitiesJSON, _ := json.Marshal(input.Capabilities)
+ regionsJSON, _ := json.Marshal(input.Regions)
var result Model
err = tx.QueryRow(ctx, `
- INSERT INTO models (public_id, owned_by, input_price_micros_per_million, output_price_micros_per_million,
- cache_read_price_micros_per_million, cache_write_price_micros_per_million)
- VALUES ($1, $2, $3, $4, $5, $6)
- RETURNING id::text, public_id, owned_by, input_price_micros_per_million, output_price_micros_per_million,
+ INSERT INTO models (public_id, display_name, description, owned_by, input_modalities, output_modalities,
+ context_window, max_output_tokens, capabilities, regions, lifecycle, released_at,
+ deprecated_at, retired_at, replacement_model, input_price_micros_per_million,
+ output_price_micros_per_million, cache_read_price_micros_per_million, cache_write_price_micros_per_million)
+ VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,NULLIF($15,''),$16,$17,$18,$19)
+ RETURNING id::text, public_id, display_name, description, owned_by,
+ input_price_micros_per_million, output_price_micros_per_million,
cache_read_price_micros_per_million, cache_write_price_micros_per_million, enabled, created_at`,
- input.PublicID, input.OwnedBy, input.InputPriceMicrosPerMillion, input.OutputPriceMicrosPerMillion,
+ input.PublicID, input.DisplayName, input.Description, input.OwnedBy, inputModalitiesJSON, outputModalitiesJSON,
+ input.ContextWindow, input.MaxOutputTokens, capabilitiesJSON, regionsJSON, input.Lifecycle,
+ input.ReleasedAt, input.DeprecatedAt, input.RetiredAt, input.ReplacementModel,
+ input.InputPriceMicrosPerMillion, input.OutputPriceMicrosPerMillion,
input.CacheReadPriceMicrosPerMillion, input.CacheWritePriceMicrosPerMillion,
- ).Scan(&result.ID, &result.PublicID, &result.OwnedBy, &result.InputPriceMicrosPerMillion,
+ ).Scan(&result.ID, &result.PublicID, &result.DisplayName, &result.Description, &result.OwnedBy, &result.InputPriceMicrosPerMillion,
&result.OutputPriceMicrosPerMillion, &result.CacheReadPriceMicrosPerMillion,
&result.CacheWritePriceMicrosPerMillion, &result.Enabled, &result.CreatedAt)
if err != nil {
return Model{}, 0, fmt.Errorf("create model: %w", err)
}
+ result.InputModalities, result.OutputModalities = input.InputModalities, input.OutputModalities
+ result.ContextWindow, result.MaxOutputTokens = input.ContextWindow, input.MaxOutputTokens
+ result.Capabilities, result.Regions, result.Lifecycle = input.Capabilities, input.Regions, input.Lifecycle
+ result.ReleasedAt, result.DeprecatedAt, result.RetiredAt = input.ReleasedAt, input.DeprecatedAt, input.RetiredAt
+ result.ReplacementModel, result.Aliases = input.ReplacementModel, input.Aliases
+ result.AllowedTenantIDs, result.AllowedKeyIDs = input.AllowedTenantIDs, input.AllowedKeyIDs
+ result.PriceCurrency, result.PriceVersion = input.PriceCurrency, 1
+ err = tx.QueryRow(ctx, `INSERT INTO model_price_versions (model_id,version,currency,
+ input_price_micros_per_million,output_price_micros_per_million,
+ cache_read_price_micros_per_million,cache_write_price_micros_per_million)
+ VALUES ($1,1,$2,$3,$4,$5,$6) RETURNING id::text,effective_from`, result.ID, input.PriceCurrency,
+ input.InputPriceMicrosPerMillion, input.OutputPriceMicrosPerMillion,
+ input.CacheReadPriceMicrosPerMillion, input.CacheWritePriceMicrosPerMillion,
+ ).Scan(&result.PriceVersionID, &result.PriceEffectiveFrom)
+ if err != nil {
+ return Model{}, 0, fmt.Errorf("create model price version: %w", err)
+ }
+ for _, alias := range input.Aliases {
+ if _, err := tx.Exec(ctx, `INSERT INTO model_aliases (alias,model_id,deprecated) VALUES ($1,$2,true)`, alias, result.ID); err != nil {
+ return Model{}, 0, fmt.Errorf("create model alias: %w", err)
+ }
+ }
+ for _, tenantID := range input.AllowedTenantIDs {
+ if _, err := tx.Exec(ctx, `INSERT INTO model_tenant_allowlist (model_id,tenant_id) VALUES ($1,$2)`, result.ID, tenantID); err != nil {
+ return Model{}, 0, fmt.Errorf("create tenant model allowlist: %w", err)
+ }
+ }
+ for _, keyID := range input.AllowedKeyIDs {
+ if _, err := tx.Exec(ctx, `INSERT INTO api_key_model_allowlist (api_key_id,model_id) VALUES ($1,$2)`, keyID, result.ID); err != nil {
+ return Model{}, 0, fmt.Errorf("create key model allowlist: %w", err)
+ }
+ }
result.Routes = make([]Route, 0, len(input.Routes))
for _, route := range input.Routes {
var created Route
@@ -233,6 +311,59 @@ func (s *Store) CreateModel(ctx context.Context, input CreateModelInput) (Model,
return result, generation, nil
}
+func (s *Store) CreateModelPriceVersion(ctx context.Context, modelID string, input CreatePriceVersionInput) (int64, error) {
+ input.Currency = strings.ToLower(strings.TrimSpace(input.Currency))
+ if len(input.Currency) != 3 || input.InputPriceMicrosPerMillion < 0 || input.OutputPriceMicrosPerMillion < 0 ||
+ input.CacheReadPriceMicrosPerMillion < 0 || input.CacheWritePriceMicrosPerMillion < 0 {
+ return 0, errors.New("price version has invalid currency or negative price")
+ }
+ if input.EffectiveFrom.IsZero() {
+ input.EffectiveFrom = time.Now().UTC()
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return 0, err
+ }
+ defer tx.Rollback(ctx)
+ // Lock the model row before allocating the next version. PostgreSQL does
+ // not allow FOR UPDATE on an aggregate query.
+ var modelExists bool
+ if err := tx.QueryRow(ctx, `SELECT true FROM models WHERE id=$1 FOR UPDATE`, modelID).Scan(&modelExists); errors.Is(err, pgx.ErrNoRows) {
+ return 0, ErrNotFound
+ } else if err != nil {
+ return 0, err
+ }
+ var currentEffectiveFrom time.Time
+ err = tx.QueryRow(ctx, `SELECT effective_from FROM model_price_versions WHERE model_id=$1 AND effective_to IS NULL`, modelID).Scan(&currentEffectiveFrom)
+ if err != nil && !errors.Is(err, pgx.ErrNoRows) {
+ return 0, err
+ }
+ if err == nil && !input.EffectiveFrom.After(currentEffectiveFrom) {
+ return 0, errors.New("new price version must become effective after the current open version")
+ }
+ var version int
+ if err := tx.QueryRow(ctx, `SELECT COALESCE(max(version),0)+1 FROM model_price_versions WHERE model_id=$1`, modelID).Scan(&version); err != nil {
+ return 0, err
+ }
+ if _, err := tx.Exec(ctx, `UPDATE model_price_versions SET effective_to=$2 WHERE model_id=$1 AND effective_to IS NULL AND effective_from < $2`, modelID, input.EffectiveFrom); err != nil {
+ return 0, err
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO model_price_versions (model_id,version,currency,input_price_micros_per_million,
+ output_price_micros_per_million,cache_read_price_micros_per_million,cache_write_price_micros_per_million,effective_from)
+ VALUES ($1,$2,$3,$4,$5,$6,$7,$8)`, modelID, version, input.Currency, input.InputPriceMicrosPerMillion,
+ input.OutputPriceMicrosPerMillion, input.CacheReadPriceMicrosPerMillion, input.CacheWritePriceMicrosPerMillion, input.EffectiveFrom); err != nil {
+ return 0, err
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return 0, err
+ }
+ return generation, nil
+}
+
func (s *Store) SetModelEnabled(ctx context.Context, id string, enabled bool) (int64, error) {
return s.toggle(ctx, `UPDATE models SET enabled = $2, updated_at = now() WHERE id = $1`, id, enabled)
}
diff --git a/internal/controlplane/outbox.go b/internal/controlplane/outbox.go
index b54a064..b708a90 100644
--- a/internal/controlplane/outbox.go
+++ b/internal/controlplane/outbox.go
@@ -109,8 +109,9 @@ func (s *Store) ClaimMail(ctx context.Context) (mailer.Message, bool, error) {
var ciphertext []byte
err = tx.QueryRow(ctx, `WITH candidate AS (
SELECT id FROM console_mail_outbox
- WHERE ((status IN ('pending','failed') AND available_at <= now())
+ WHERE ((status IN ('pending','retry') AND available_at <= now())
OR (status='sending' AND claimed_at < now()-interval '5 minutes'))
+ AND NOT EXISTS (SELECT 1 FROM mail_suppressions s WHERE lower(s.recipient)=lower(console_mail_outbox.recipient))
ORDER BY available_at,created_at FOR UPDATE SKIP LOCKED LIMIT 1
) UPDATE console_mail_outbox o SET status='sending',claimed_at=now(),attempts=attempts+1,last_error=''
FROM candidate WHERE o.id=candidate.id
@@ -146,7 +147,19 @@ func (s *Store) MarkMailFailed(ctx context.Context, id string, deliveryErr error
if len(message) > 1000 {
message = message[:1000]
}
- _, err := s.db.Exec(ctx, `UPDATE console_mail_outbox SET status='failed',claimed_at=NULL,last_error=$2,
- available_at=now()+make_interval(secs => LEAST(300, 5 * attempts)) WHERE id=$1 AND status='sending'`, id, message)
+ _, err := s.db.Exec(ctx, `UPDATE console_mail_outbox SET
+ status=CASE WHEN attempts>=10 THEN 'dead' ELSE 'retry' END,claimed_at=NULL,last_error=$2,
+ available_at=now()+make_interval(secs => LEAST(3600, 5 * power(2,LEAST(attempts,9))::int))
+ WHERE id=$1 AND status='sending'`, id, message)
return err
}
+
+func (s *Store) MailQueueStatus(ctx context.Context) (MailQueueStatus, error) {
+ var result MailQueueStatus
+ err := s.db.QueryRow(ctx, `SELECT
+ count(*) FILTER (WHERE status IN ('pending','sending','retry')),
+ count(*) FILTER (WHERE status='dead'),
+ min(created_at) FILTER (WHERE status IN ('pending','sending','retry'))
+ FROM console_mail_outbox`).Scan(&result.Backlog, &result.Failed, &result.OldestPending)
+ return result, err
+}
diff --git a/internal/controlplane/queries.go b/internal/controlplane/queries.go
index 9b76fad..73d869f 100644
--- a/internal/controlplane/queries.go
+++ b/internal/controlplane/queries.go
@@ -173,10 +173,17 @@ func (s *Store) ListProviders(ctx context.Context) ([]Provider, error) {
func (s *Store) ListModels(ctx context.Context) ([]Model, error) {
rows, err := s.db.Query(ctx, `
- SELECT id::text, public_id, owned_by, input_price_micros_per_million,
- output_price_micros_per_million, cache_read_price_micros_per_million,
- cache_write_price_micros_per_million, enabled, created_at
- FROM models ORDER BY public_id`)
+ SELECT m.id::text, m.public_id, m.display_name, m.description, m.owned_by,
+ m.input_modalities, m.output_modalities, m.context_window, m.max_output_tokens,
+ m.capabilities, m.regions, m.lifecycle, m.released_at, m.deprecated_at, m.retired_at,
+ COALESCE(m.replacement_model,''), pv.id::text, pv.version, pv.currency, pv.effective_from,
+ pv.input_price_micros_per_million, pv.output_price_micros_per_million,
+ pv.cache_read_price_micros_per_million, pv.cache_write_price_micros_per_million,
+ m.enabled, m.created_at
+ FROM models m JOIN LATERAL (
+ SELECT * FROM model_price_versions v WHERE v.model_id=m.id
+ ORDER BY v.effective_from DESC LIMIT 1
+ ) pv ON TRUE ORDER BY m.public_id`)
if err != nil {
return nil, fmt.Errorf("query models: %w", err)
}
@@ -184,13 +191,23 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) {
positions := make(map[string]int)
for rows.Next() {
var item Model
- if err := rows.Scan(&item.ID, &item.PublicID, &item.OwnedBy, &item.InputPriceMicrosPerMillion,
+ var inputModalitiesJSON, outputModalitiesJSON, capabilitiesJSON, regionsJSON []byte
+ if err := rows.Scan(&item.ID, &item.PublicID, &item.DisplayName, &item.Description, &item.OwnedBy,
+ &inputModalitiesJSON, &outputModalitiesJSON, &item.ContextWindow, &item.MaxOutputTokens,
+ &capabilitiesJSON, &regionsJSON, &item.Lifecycle, &item.ReleasedAt, &item.DeprecatedAt,
+ &item.RetiredAt, &item.ReplacementModel, &item.PriceVersionID, &item.PriceVersion,
+ &item.PriceCurrency, &item.PriceEffectiveFrom, &item.InputPriceMicrosPerMillion,
&item.OutputPriceMicrosPerMillion, &item.CacheReadPriceMicrosPerMillion,
&item.CacheWritePriceMicrosPerMillion, &item.Enabled, &item.CreatedAt); err != nil {
rows.Close()
return nil, fmt.Errorf("scan model: %w", err)
}
+ _ = json.Unmarshal(inputModalitiesJSON, &item.InputModalities)
+ _ = json.Unmarshal(outputModalitiesJSON, &item.OutputModalities)
+ _ = json.Unmarshal(capabilitiesJSON, &item.Capabilities)
+ _ = json.Unmarshal(regionsJSON, &item.Regions)
item.Routes = []Route{}
+ item.Aliases, item.AllowedTenantIDs, item.AllowedKeyIDs = []string{}, []string{}, []string{}
positions[item.ID] = len(models)
models = append(models, item)
}
@@ -200,6 +217,52 @@ func (s *Store) ListModels(ctx context.Context) ([]Model, error) {
}
rows.Close()
+ aliasRows, err := s.db.Query(ctx, `SELECT model_id::text, alias FROM model_aliases ORDER BY alias`)
+ if err != nil {
+ return nil, fmt.Errorf("query model aliases: %w", err)
+ }
+ for aliasRows.Next() {
+ var modelID, alias string
+ if err := aliasRows.Scan(&modelID, &alias); err != nil {
+ aliasRows.Close()
+ return nil, err
+ }
+ if p, ok := positions[modelID]; ok {
+ models[p].Aliases = append(models[p].Aliases, alias)
+ }
+ }
+ aliasRows.Close()
+ tenantRows, err := s.db.Query(ctx, `SELECT model_id::text, tenant_id::text FROM model_tenant_allowlist`)
+ if err != nil {
+ return nil, fmt.Errorf("query tenant model allowlist: %w", err)
+ }
+ for tenantRows.Next() {
+ var modelID, tenantID string
+ if err := tenantRows.Scan(&modelID, &tenantID); err != nil {
+ tenantRows.Close()
+ return nil, err
+ }
+ if p, ok := positions[modelID]; ok {
+ models[p].AllowedTenantIDs = append(models[p].AllowedTenantIDs, tenantID)
+ }
+ }
+ tenantRows.Close()
+ keyRows, err := s.db.Query(ctx, `SELECT model_id::text, api_key_id::text FROM api_key_model_allowlist`)
+ if err != nil {
+ return nil, fmt.Errorf("query key model allowlist: %w", err)
+ }
+ for keyRows.Next() {
+ var modelID, keyID string
+ if err := keyRows.Scan(&modelID, &keyID); err != nil {
+ keyRows.Close()
+ return nil, err
+ }
+ if p, ok := positions[modelID]; ok {
+ models[p].AllowedKeyIDs = append(models[p].AllowedKeyIDs, keyID)
+ }
+ }
+ keyRows.Close()
+
routeRows, err := s.db.Query(ctx, `
SELECT r.id::text, r.model_id::text, r.provider_id::text, p.name, p.protocol,
r.upstream_model, r.priority, r.weight, r.enabled
diff --git a/internal/controlplane/retention.go b/internal/controlplane/retention.go
new file mode 100644
index 0000000..4b20611
--- /dev/null
+++ b/internal/controlplane/retention.go
@@ -0,0 +1,57 @@
+package controlplane
+
+import (
+ "context"
+ "log/slog"
+ "time"
+)
+
+func (s *Store) RunRetentionWorker(ctx context.Context, auditDays, securityDays int, logger *slog.Logger) {
+ run := func() {
+ cleanupCtx, cancel := context.WithTimeout(ctx, 30*time.Second)
+ defer cancel()
+ if err := s.pruneExpiredSecurityData(cleanupCtx, auditDays, securityDays); err != nil {
+ logger.Error("retention_cleanup_failed", "error", err)
+ }
+ }
+ run()
+ ticker := time.NewTicker(24 * time.Hour)
+ defer ticker.Stop()
+ for {
+ select {
+ case <-ctx.Done():
+ return
+ case <-ticker.C:
+ run()
+ }
+ }
+}
+
+func (s *Store) pruneExpiredSecurityData(ctx context.Context, auditDays, securityDays int) error {
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return err
+ }
+ defer tx.Rollback(ctx)
+ auditCutoff := time.Now().UTC().AddDate(0, 0, -auditDays)
+ securityCutoff := time.Now().UTC().AddDate(0, 0, -securityDays)
+ queries := []struct {
+ query string
+ cutoff time.Time
+ }{
+ {`DELETE FROM audit_logs WHERE created_at < $1`, auditCutoff},
+ {`DELETE FROM console_sessions WHERE expires_at < $1 OR (revoked_at IS NOT NULL AND revoked_at < $1)`, securityCutoff},
+ {`DELETE FROM console_action_tokens WHERE expires_at < $1 OR (consumed_at IS NOT NULL AND consumed_at < $1)`, securityCutoff},
+ {`DELETE FROM console_auth_challenges WHERE expires_at < $1 OR (consumed_at IS NOT NULL AND consumed_at < $1)`, securityCutoff},
+ {`DELETE FROM console_webauthn_challenges WHERE expires_at < $1 OR (consumed_at IS NOT NULL AND consumed_at < $1)`, securityCutoff},
+ {`DELETE FROM console_login_throttles WHERE updated_at < $1`, securityCutoff},
+ {`DELETE FROM console_rate_limits WHERE updated_at < $1`, securityCutoff},
+ {`DELETE FROM console_mail_outbox WHERE status='sent' AND sent_at < $1`, securityCutoff},
+ }
+ for _, item := range queries {
+ if _, err := tx.Exec(ctx, item.query, item.cutoff); err != nil {
+ return err
+ }
+ }
+ return tx.Commit(ctx)
+}
diff --git a/internal/controlplane/rotation.go b/internal/controlplane/rotation.go
new file mode 100644
index 0000000..1fc8084
--- /dev/null
+++ b/internal/controlplane/rotation.go
@@ -0,0 +1,69 @@
+package controlplane
+
+import (
+ "context"
+ "fmt"
+)
+
+type encryptedColumn struct{ table, key, column string }
+
+// RotateCredentials re-encrypts every control-plane secret with the primary
+// key in the configured keyring. Run all gateway instances with both new and
+// previous keys before invoking this operation.
+func (s *Store) RotateCredentials(ctx context.Context) (int, error) {
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return 0, err
+ }
+ defer tx.Rollback(ctx)
+ columns := []encryptedColumn{
+ {"providers", "id", "api_key_ciphertext"},
+ {"console_mail_outbox", "id", "body_ciphertext"},
+ {"console_totp_credentials", "user_id", "secret_ciphertext"},
+ {"console_passkeys", "id", "credential_ciphertext"},
+ {"console_webauthn_challenges", "id", "session_ciphertext"},
+ }
+ total := 0
+ for _, item := range columns {
+ rows, queryErr := tx.Query(ctx, fmt.Sprintf(`SELECT %s::text,%s FROM %s`, item.key, item.column, item.table))
+ if queryErr != nil {
+ return total, queryErr
+ }
+ type record struct {
+ id string
+ ciphertext []byte
+ }
+ records := []record{}
+ for rows.Next() {
+ var value record
+ if err := rows.Scan(&value.id, &value.ciphertext); err != nil {
+ rows.Close()
+ return total, err
+ }
+ records = append(records, value)
+ }
+ if err := rows.Err(); err != nil {
+ rows.Close()
+ return total, err
+ }
+ rows.Close()
+ for _, value := range records {
+ plaintext, err := s.cipher.Decrypt(value.ciphertext)
+ if err != nil {
+ return total, fmt.Errorf("decrypt %s %s: %w", item.table, value.id, err)
+ }
+ ciphertext, err := s.cipher.Encrypt(plaintext)
+ if err != nil {
+ return total, err
+ }
+ if _, err := tx.Exec(ctx, fmt.Sprintf(`UPDATE %s SET %s=$2 WHERE %s=$1`, item.table, item.column, item.key), value.id, ciphertext); err != nil {
+ return total, err
+ }
+ total++
+ }
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return total, err
+ }
+ return total, nil
+}
diff --git a/internal/controlplane/schema.sql b/internal/controlplane/schema.sql
index 8a5a605..88db49b 100644
--- a/internal/controlplane/schema.sql
+++ b/internal/controlplane/schema.sql
@@ -1,5 +1,12 @@
CREATE EXTENSION IF NOT EXISTS pgcrypto;
+CREATE TABLE IF NOT EXISTS schema_migrations (
+ version BIGINT PRIMARY KEY,
+ name TEXT NOT NULL,
+ checksum TEXT NOT NULL,
+ applied_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
CREATE TABLE IF NOT EXISTS control_state (
singleton BOOLEAN PRIMARY KEY DEFAULT TRUE CHECK (singleton),
generation BIGINT NOT NULL DEFAULT 0,
@@ -71,6 +78,72 @@ ALTER TABLE models ADD COLUMN IF NOT EXISTS input_price_micros_per_million BIGIN
ALTER TABLE models ADD COLUMN IF NOT EXISTS output_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (output_price_micros_per_million >= 0);
ALTER TABLE models ADD COLUMN IF NOT EXISTS cache_read_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (cache_read_price_micros_per_million >= 0);
ALTER TABLE models ADD COLUMN IF NOT EXISTS cache_write_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (cache_write_price_micros_per_million >= 0);
+ALTER TABLE models ADD COLUMN IF NOT EXISTS display_name TEXT NOT NULL DEFAULT '';
+ALTER TABLE models ADD COLUMN IF NOT EXISTS description TEXT NOT NULL DEFAULT '';
+ALTER TABLE models ADD COLUMN IF NOT EXISTS input_modalities JSONB NOT NULL DEFAULT '["text"]'::jsonb;
+ALTER TABLE models ADD COLUMN IF NOT EXISTS output_modalities JSONB NOT NULL DEFAULT '["text"]'::jsonb;
+ALTER TABLE models ADD COLUMN IF NOT EXISTS context_window BIGINT NOT NULL DEFAULT 0 CHECK (context_window >= 0);
+ALTER TABLE models ADD COLUMN IF NOT EXISTS max_output_tokens BIGINT NOT NULL DEFAULT 0 CHECK (max_output_tokens >= 0);
+ALTER TABLE models ADD COLUMN IF NOT EXISTS capabilities JSONB NOT NULL DEFAULT '["chat","streaming"]'::jsonb;
+ALTER TABLE models ADD COLUMN IF NOT EXISTS regions JSONB NOT NULL DEFAULT '[]'::jsonb;
+ALTER TABLE models ADD COLUMN IF NOT EXISTS lifecycle TEXT NOT NULL DEFAULT 'active';
+ALTER TABLE models ADD COLUMN IF NOT EXISTS released_at TIMESTAMPTZ;
+ALTER TABLE models ADD COLUMN IF NOT EXISTS deprecated_at TIMESTAMPTZ;
+ALTER TABLE models ADD COLUMN IF NOT EXISTS retired_at TIMESTAMPTZ;
+ALTER TABLE models ADD COLUMN IF NOT EXISTS replacement_model TEXT;
+ALTER TABLE models DROP CONSTRAINT IF EXISTS models_lifecycle_check;
+ALTER TABLE models ADD CONSTRAINT models_lifecycle_check CHECK (lifecycle IN ('preview','active','deprecated','retired'));
+
+CREATE TABLE IF NOT EXISTS model_price_versions (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE,
+ version INTEGER NOT NULL CHECK (version > 0),
+ currency TEXT NOT NULL CHECK (currency = lower(currency) AND length(currency) = 3),
+ input_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (input_price_micros_per_million >= 0),
+ output_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (output_price_micros_per_million >= 0),
+ cache_read_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (cache_read_price_micros_per_million >= 0),
+ cache_write_price_micros_per_million BIGINT NOT NULL DEFAULT 0 CHECK (cache_write_price_micros_per_million >= 0),
+ effective_from TIMESTAMPTZ NOT NULL DEFAULT now(),
+ effective_to TIMESTAMPTZ,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ UNIQUE (model_id, version),
+ CHECK (effective_to IS NULL OR effective_to > effective_from)
+);
+CREATE UNIQUE INDEX IF NOT EXISTS model_price_versions_one_open_idx
+ ON model_price_versions (model_id) WHERE effective_to IS NULL;
+CREATE INDEX IF NOT EXISTS model_price_versions_effective_idx
+ ON model_price_versions (model_id, effective_from DESC);
+
+INSERT INTO model_price_versions (
+ model_id, version, currency, input_price_micros_per_million,
+ output_price_micros_per_million, cache_read_price_micros_per_million,
+ cache_write_price_micros_per_million, effective_from)
+SELECT id, 1, 'usd', input_price_micros_per_million, output_price_micros_per_million,
+ cache_read_price_micros_per_million, cache_write_price_micros_per_million, created_at
+FROM models
+ON CONFLICT (model_id, version) DO NOTHING;
+
+CREATE TABLE IF NOT EXISTS model_aliases (
+ alias TEXT PRIMARY KEY,
+ model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE,
+ deprecated BOOLEAN NOT NULL DEFAULT FALSE,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+CREATE INDEX IF NOT EXISTS model_aliases_model_idx ON model_aliases (model_id);
+
+CREATE TABLE IF NOT EXISTS model_tenant_allowlist (
+ model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE,
+ tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE CASCADE,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ PRIMARY KEY (model_id, tenant_id)
+);
+
+CREATE TABLE IF NOT EXISTS api_key_model_allowlist (
+ api_key_id UUID NOT NULL REFERENCES api_keys(id) ON DELETE CASCADE,
+ model_id UUID NOT NULL REFERENCES models(id) ON DELETE CASCADE,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ PRIMARY KEY (api_key_id, model_id)
+);
CREATE TABLE IF NOT EXISTS model_routes (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
@@ -106,7 +179,7 @@ CREATE TABLE IF NOT EXISTS billing_reservations (
public_model TEXT NOT NULL,
currency TEXT NOT NULL,
reserved_micros BIGINT NOT NULL CHECK (reserved_micros >= 0),
- status TEXT NOT NULL DEFAULT 'pending' CHECK (status IN ('pending', 'settled', 'released')),
+ status TEXT NOT NULL DEFAULT 'pending' CHECK (status IN ('pending', 'settled', 'released', 'metering_failed')),
input_price_micros_per_million BIGINT NOT NULL,
output_price_micros_per_million BIGINT NOT NULL,
cache_read_price_micros_per_million BIGINT NOT NULL,
@@ -118,6 +191,30 @@ CREATE TABLE IF NOT EXISTS billing_reservations (
settled_at TIMESTAMPTZ,
FOREIGN KEY (project_id, tenant_id) REFERENCES projects(id, tenant_id) ON DELETE CASCADE
);
+ALTER TABLE billing_reservations ADD COLUMN IF NOT EXISTS price_version_id UUID REFERENCES model_price_versions(id) ON DELETE SET NULL;
+ALTER TABLE billing_reservations DROP CONSTRAINT IF EXISTS billing_reservations_status_check;
+ALTER TABLE billing_reservations ADD CONSTRAINT billing_reservations_status_check
+ CHECK (status IN ('pending','settled','released','metering_failed'));
+
+CREATE TABLE IF NOT EXISTS billing_settlement_jobs (
+ request_id TEXT PRIMARY KEY REFERENCES billing_reservations(request_id) ON DELETE CASCADE,
+ event JSONB,
+ status TEXT NOT NULL DEFAULT 'awaiting_event'
+ CHECK (status IN ('awaiting_event', 'pending', 'processing', 'retry', 'done')),
+ attempts INTEGER NOT NULL DEFAULT 0 CHECK (attempts >= 0),
+ available_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ locked_at TIMESTAMPTZ,
+ last_error TEXT NOT NULL DEFAULT '',
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ completed_at TIMESTAMPTZ
+);
+CREATE INDEX IF NOT EXISTS billing_settlement_jobs_ready_idx
+ ON billing_settlement_jobs (available_at, created_at)
+ WHERE status IN ('pending', 'retry');
+CREATE INDEX IF NOT EXISTS billing_settlement_jobs_stale_idx
+ ON billing_settlement_jobs (created_at)
+ WHERE status IN ('awaiting_event', 'processing', 'retry');
CREATE TABLE IF NOT EXISTS usage_events (
request_id TEXT PRIMARY KEY,
@@ -143,9 +240,17 @@ CREATE TABLE IF NOT EXISTS usage_events (
cost_micros BIGINT NOT NULL DEFAULT 0,
charged_micros BIGINT NOT NULL DEFAULT 0,
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')),
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 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'));
-- Usage persistence is independent from billing. Older installations created this
-- foreign key, which prevented recording requests when prepaid billing was disabled.
@@ -165,6 +270,9 @@ CREATE TABLE IF NOT EXISTS billing_ledger (
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
UNIQUE (source_type, source_id)
);
+ALTER TABLE billing_ledger DROP CONSTRAINT IF EXISTS billing_ledger_kind_check;
+ALTER TABLE billing_ledger ADD CONSTRAINT billing_ledger_kind_check
+ CHECK (kind IN ('topup','usage','adjustment','refund','release','dispute','dispute_reversal'));
CREATE TABLE IF NOT EXISTS topup_orders (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
@@ -178,13 +286,127 @@ CREATE TABLE IF NOT EXISTS topup_orders (
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
paid_at TIMESTAMPTZ
);
+ALTER TABLE topup_orders DROP CONSTRAINT IF EXISTS topup_orders_status_check;
+ALTER TABLE topup_orders ADD CONSTRAINT topup_orders_status_check
+ CHECK (status IN ('pending','paid','failed','expired','partially_refunded','refunded','disputed','reversed'));
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS stripe_customer_id TEXT;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS stripe_payment_intent_id TEXT;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS stripe_charge_id TEXT;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS stripe_invoice_id TEXT;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS invoice_url TEXT;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS invoice_pdf_url TEXT;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS receipt_url TEXT;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS refunded_micros BIGINT NOT NULL DEFAULT 0 CHECK (refunded_micros >= 0);
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS disputed_micros BIGINT NOT NULL DEFAULT 0 CHECK (disputed_micros >= 0);
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS reconciliation_status TEXT NOT NULL DEFAULT 'unknown';
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS reconciled_at TIMESTAMPTZ;
+ALTER TABLE topup_orders ADD COLUMN IF NOT EXISTS reconciliation_error TEXT NOT NULL DEFAULT '';
+ALTER TABLE topup_orders DROP CONSTRAINT IF EXISTS topup_orders_reconciliation_status_check;
+ALTER TABLE topup_orders ADD CONSTRAINT topup_orders_reconciliation_status_check
+ CHECK (reconciliation_status IN ('unknown','ok','repaired','missing','mismatch','resolved'));
+CREATE INDEX IF NOT EXISTS topup_orders_payment_intent_idx ON topup_orders (stripe_payment_intent_id) WHERE stripe_payment_intent_id IS NOT NULL;
+CREATE INDEX IF NOT EXISTS topup_orders_customer_idx ON topup_orders (stripe_customer_id) WHERE stripe_customer_id IS NOT NULL;
+
+CREATE TABLE IF NOT EXISTS billing_reconciliation_resolutions (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ topup_order_id UUID NOT NULL UNIQUE REFERENCES topup_orders(id) ON DELETE RESTRICT,
+ tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE RESTRICT,
+ actor_id TEXT NOT NULL DEFAULT '',
+ actor_type TEXT NOT NULL CHECK (actor_type IN ('console_user','bootstrap','maintenance')),
+ reason TEXT NOT NULL,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE TABLE IF NOT EXISTS stripe_customers (
+ tenant_id UUID PRIMARY KEY REFERENCES tenants(id) ON DELETE CASCADE,
+ stripe_customer_id TEXT NOT NULL UNIQUE,
+ email TEXT NOT NULL DEFAULT '',
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE TABLE IF NOT EXISTS stripe_refunds (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE RESTRICT,
+ topup_order_id UUID NOT NULL REFERENCES topup_orders(id) ON DELETE RESTRICT,
+ stripe_refund_id TEXT UNIQUE,
+ amount_minor BIGINT NOT NULL CHECK (amount_minor > 0),
+ amount_micros BIGINT NOT NULL CHECK (amount_micros > 0),
+ held_micros BIGINT NOT NULL DEFAULT 0 CHECK (held_micros >= 0),
+ uncollected_micros BIGINT NOT NULL DEFAULT 0 CHECK (uncollected_micros >= 0),
+ currency TEXT NOT NULL,
+ reason TEXT NOT NULL DEFAULT 'requested_by_customer',
+ status TEXT NOT NULL DEFAULT 'queued'
+ CHECK (status IN ('queued','submitting','pending','requires_action','succeeded','failed','canceled')),
+ attempts INTEGER NOT NULL DEFAULT 0,
+ available_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ last_error TEXT NOT NULL DEFAULT '',
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ completed_at TIMESTAMPTZ
+);
+CREATE INDEX IF NOT EXISTS stripe_refunds_ready_idx ON stripe_refunds (available_at, created_at)
+ WHERE status IN ('queued','submitting');
+
+CREATE TABLE IF NOT EXISTS stripe_disputes (
+ stripe_dispute_id TEXT PRIMARY KEY,
+ tenant_id UUID REFERENCES tenants(id) ON DELETE SET NULL,
+ topup_order_id UUID REFERENCES topup_orders(id) ON DELETE SET NULL,
+ stripe_payment_intent_id TEXT,
+ amount_minor BIGINT NOT NULL CHECK (amount_minor >= 0),
+ amount_micros BIGINT NOT NULL CHECK (amount_micros >= 0),
+ currency TEXT NOT NULL,
+ status TEXT NOT NULL,
+ reason TEXT NOT NULL DEFAULT '',
+ debited_micros BIGINT NOT NULL DEFAULT 0 CHECK (debited_micros >= 0),
+ uncollected_micros BIGINT NOT NULL DEFAULT 0 CHECK (uncollected_micros >= 0),
+ due_by TIMESTAMPTZ,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ closed_at TIMESTAMPTZ
+);
+
+CREATE TABLE IF NOT EXISTS stripe_invoices (
+ stripe_invoice_id TEXT PRIMARY KEY,
+ tenant_id UUID REFERENCES tenants(id) ON DELETE SET NULL,
+ topup_order_id UUID REFERENCES topup_orders(id) ON DELETE SET NULL,
+ stripe_customer_id TEXT,
+ status TEXT NOT NULL DEFAULT '',
+ currency TEXT NOT NULL DEFAULT '',
+ amount_due_minor BIGINT NOT NULL DEFAULT 0,
+ amount_paid_minor BIGINT NOT NULL DEFAULT 0,
+ attempt_count INTEGER NOT NULL DEFAULT 0,
+ next_payment_attempt TIMESTAMPTZ,
+ hosted_invoice_url TEXT NOT NULL DEFAULT '',
+ invoice_pdf_url TEXT NOT NULL DEFAULT '',
+ last_failure TEXT NOT NULL DEFAULT '',
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE TABLE IF NOT EXISTS billing_reconciliation_runs (
+ id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
+ status TEXT NOT NULL CHECK (status IN ('running','clean','mismatch','failed')),
+ checked_orders BIGINT NOT NULL DEFAULT 0,
+ mismatch_count BIGINT NOT NULL DEFAULT 0,
+ report JSONB NOT NULL DEFAULT '[]'::jsonb,
+ error TEXT NOT NULL DEFAULT '',
+ started_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ completed_at TIMESTAMPTZ
+);
CREATE TABLE IF NOT EXISTS stripe_webhook_events (
event_id TEXT PRIMARY KEY,
event_type TEXT NOT NULL,
processed_at TIMESTAMPTZ,
+ attempts INTEGER NOT NULL DEFAULT 0,
+ last_attempt_at TIMESTAMPTZ,
+ processing_error TEXT NOT NULL DEFAULT '',
created_at TIMESTAMPTZ NOT NULL DEFAULT now()
);
+ALTER TABLE stripe_webhook_events ADD COLUMN IF NOT EXISTS attempts INTEGER NOT NULL DEFAULT 0;
+ALTER TABLE stripe_webhook_events ADD COLUMN IF NOT EXISTS last_attempt_at TIMESTAMPTZ;
+ALTER TABLE stripe_webhook_events ADD COLUMN IF NOT EXISTS processing_error TEXT NOT NULL DEFAULT '';
CREATE TABLE IF NOT EXISTS console_users (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
@@ -279,7 +501,7 @@ CREATE INDEX IF NOT EXISTS console_action_tokens_active_idx
CREATE TABLE IF NOT EXISTS console_mail_outbox (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
recipient TEXT NOT NULL,
- template TEXT NOT NULL CHECK (template IN ('verify_email', 'password_reset', 'invite')),
+ template TEXT NOT NULL CHECK (template IN ('verify_email', 'password_reset', 'invite', 'low_balance', 'spend_anomaly')),
subject TEXT NOT NULL,
body_ciphertext BYTEA NOT NULL,
status TEXT NOT NULL DEFAULT 'pending' CHECK (status IN ('pending', 'sending', 'sent', 'failed')),
@@ -290,8 +512,43 @@ CREATE TABLE IF NOT EXISTS console_mail_outbox (
last_error TEXT NOT NULL DEFAULT '',
created_at TIMESTAMPTZ NOT NULL DEFAULT now()
);
-CREATE INDEX IF NOT EXISTS console_mail_outbox_pending_idx
- ON console_mail_outbox (available_at, created_at) WHERE status IN ('pending', 'failed');
+ALTER TABLE console_mail_outbox DROP CONSTRAINT IF EXISTS console_mail_outbox_template_check;
+ALTER TABLE console_mail_outbox ADD CONSTRAINT console_mail_outbox_template_check
+ CHECK (template IN ('verify_email','password_reset','invite','low_balance','spend_anomaly'));
+ALTER TABLE console_mail_outbox DROP CONSTRAINT IF EXISTS console_mail_outbox_status_check;
+UPDATE console_mail_outbox SET status='retry' WHERE status='failed';
+ALTER TABLE console_mail_outbox ADD CONSTRAINT console_mail_outbox_status_check
+ CHECK (status IN ('pending','sending','sent','retry','dead','suppressed'));
+DROP INDEX IF EXISTS console_mail_outbox_pending_idx;
+CREATE INDEX console_mail_outbox_pending_idx
+ ON console_mail_outbox (available_at, created_at) WHERE status IN ('pending', 'retry');
+
+CREATE TABLE IF NOT EXISTS mail_suppressions (
+ recipient TEXT PRIMARY KEY,
+ reason TEXT NOT NULL CHECK (reason IN ('bounce','complaint','manual')),
+ provider TEXT NOT NULL DEFAULT '',
+ provider_event_id TEXT NOT NULL DEFAULT '',
+ detail TEXT NOT NULL DEFAULT '',
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE TABLE IF NOT EXISTS mail_feedback_events (
+ event_id TEXT PRIMARY KEY,
+ event_type TEXT NOT NULL CHECK (event_type IN ('delivered','bounce','complaint')),
+ recipient TEXT NOT NULL,
+ provider TEXT NOT NULL DEFAULT '',
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now()
+);
+
+CREATE TABLE IF NOT EXISTS mail_notification_events (
+ tenant_id UUID NOT NULL REFERENCES tenants(id) ON DELETE CASCADE,
+ recipient TEXT NOT NULL,
+ notification_type TEXT NOT NULL CHECK (notification_type IN ('low_balance','spend_anomaly')),
+ dedupe_key TEXT NOT NULL,
+ created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
+ PRIMARY KEY (tenant_id,recipient,notification_type,dedupe_key)
+);
CREATE TABLE IF NOT EXISTS console_auth_challenges (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
diff --git a/internal/controlplane/snapshot.go b/internal/controlplane/snapshot.go
index bc04e7c..b9bddaf 100644
--- a/internal/controlplane/snapshot.go
+++ b/internal/controlplane/snapshot.go
@@ -6,6 +6,7 @@ import (
"encoding/json"
"fmt"
"strings"
+ "time"
"aigw/internal/auth"
"aigw/internal/domain"
@@ -102,13 +103,23 @@ func (s *Store) loadProviders(ctx context.Context, tx pgx.Tx) (map[string]domain
func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Provider) ([]domain.Model, error) {
rows, err := tx.Query(ctx, `
- SELECT m.public_id, m.owned_by, m.input_price_micros_per_million, m.output_price_micros_per_million,
- m.cache_read_price_micros_per_million, m.cache_write_price_micros_per_million,
+ SELECT m.id::text, m.public_id, m.display_name, m.description, m.owned_by,
+ m.input_modalities, m.output_modalities, m.context_window, m.max_output_tokens,
+ m.capabilities, m.regions, m.lifecycle, m.released_at, m.deprecated_at,
+ m.retired_at, COALESCE(m.replacement_model,''),
+ pv.id::text, pv.version, pv.currency, pv.effective_from,
+ pv.input_price_micros_per_million, pv.output_price_micros_per_million,
+ pv.cache_read_price_micros_per_million, pv.cache_write_price_micros_per_million,
r.provider_id::text, r.upstream_model, r.priority, r.weight
FROM models m
+ JOIN LATERAL (
+ SELECT * FROM model_price_versions v WHERE v.model_id=m.id
+ AND v.effective_from <= now() AND (v.effective_to IS NULL OR v.effective_to > now())
+ ORDER BY v.effective_from DESC LIMIT 1
+ ) pv ON TRUE
JOIN model_routes r ON r.model_id = m.id AND r.enabled = TRUE
JOIN providers p ON p.id = r.provider_id AND p.enabled = TRUE
- WHERE m.enabled = TRUE
+ WHERE m.enabled = TRUE AND m.lifecycle <> 'retired'
ORDER BY m.public_id, r.priority, r.created_at`)
if err != nil {
return nil, fmt.Errorf("query model routes: %w", err)
@@ -117,10 +128,21 @@ func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Prov
models := make([]domain.Model, 0)
index := make(map[string]int)
for rows.Next() {
- var publicID, ownedBy, providerID, upstreamModel string
+ var modelID, publicID, displayName, description, ownedBy, replacement, providerID, upstreamModel string
+ var inputModalitiesJSON, outputModalitiesJSON, capabilitiesJSON, regionsJSON []byte
+ var contextWindow, maxOutput int64
+ var lifecycle, priceVersionID, priceCurrency string
+ var releasedAt, deprecatedAt, retiredAt *time.Time
+ var priceVersion int
+ var priceEffectiveFrom time.Time
var inputPrice, outputPrice, cacheReadPrice, cacheWritePrice int64
var priority, weight int
- if err := rows.Scan(&publicID, &ownedBy, &inputPrice, &outputPrice, &cacheReadPrice, &cacheWritePrice, &providerID, &upstreamModel, &priority, &weight); err != nil {
+ if err := rows.Scan(&modelID, &publicID, &displayName, &description, &ownedBy,
+ &inputModalitiesJSON, &outputModalitiesJSON, &contextWindow, &maxOutput,
+ &capabilitiesJSON, &regionsJSON, &lifecycle, &releasedAt, &deprecatedAt, &retiredAt, &replacement,
+ &priceVersionID, &priceVersion, &priceCurrency, &priceEffectiveFrom,
+ &inputPrice, &outputPrice, &cacheReadPrice, &cacheWritePrice,
+ &providerID, &upstreamModel, &priority, &weight); err != nil {
return nil, fmt.Errorf("scan model route: %w", err)
}
provider, ok := providers[providerID]
@@ -129,13 +151,32 @@ func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Prov
}
position, exists := index[publicID]
if !exists {
+ var inputModalities, outputModalities, capabilities, regions []string
+ if err := json.Unmarshal(inputModalitiesJSON, &inputModalities); err != nil {
+ return nil, fmt.Errorf("decode model input modalities: %w", err)
+ }
+ if err := json.Unmarshal(outputModalitiesJSON, &outputModalities); err != nil {
+ return nil, fmt.Errorf("decode model output modalities: %w", err)
+ }
+ if err := json.Unmarshal(capabilitiesJSON, &capabilities); err != nil {
+ return nil, fmt.Errorf("decode model capabilities: %w", err)
+ }
+ if err := json.Unmarshal(regionsJSON, &regions); err != nil {
+ return nil, fmt.Errorf("decode model regions: %w", err)
+ }
position = len(models)
index[publicID] = position
models = append(models, domain.Model{
- ID: publicID, OwnedBy: ownedBy,
+ ID: publicID, DisplayName: displayName, Description: description, OwnedBy: ownedBy,
+ InputModalities: inputModalities, OutputModalities: outputModalities,
+ ContextWindow: contextWindow, MaxOutputTokens: maxOutput, Capabilities: capabilities,
+ Regions: regions, Lifecycle: lifecycle, ReleasedAt: releasedAt, DeprecatedAt: deprecatedAt,
+ RetiredAt: retiredAt, ReplacementModel: replacement, PriceVersionID: priceVersionID,
+ PriceVersion: priceVersion, PriceCurrency: priceCurrency, PriceEffectiveFrom: priceEffectiveFrom,
InputPriceMicrosPerMillion: inputPrice, OutputPriceMicrosPerMillion: outputPrice,
CacheReadPriceMicrosPerMillion: cacheReadPrice, CacheWritePriceMicrosPerMillion: cacheWritePrice,
})
+ _ = modelID
}
models[position].Routes = append(models[position].Routes, domain.Route{
Provider: provider, UpstreamModel: upstreamModel, Priority: priority, Weight: weight,
@@ -144,6 +185,61 @@ func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Prov
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("read model routes: %w", err)
}
+ byID := make(map[string]int, len(models))
+ for i := range models {
+ byID[models[i].ID] = i
+ }
+ aliasRows, err := tx.Query(ctx, `SELECT a.alias, m.public_id FROM model_aliases a JOIN models m ON m.id=a.model_id`)
+ if err != nil {
+ return nil, fmt.Errorf("query model aliases: %w", err)
+ }
+ for aliasRows.Next() {
+ var alias, publicID string
+ if err := aliasRows.Scan(&alias, &publicID); err != nil {
+ aliasRows.Close()
+ return nil, err
+ }
+ if position, ok := byID[publicID]; ok {
+ models[position].Aliases = append(models[position].Aliases, alias)
+ }
+ }
+ aliasRows.Close()
+ tenantRows, err := tx.Query(ctx, `SELECT m.public_id, a.tenant_id::text FROM model_tenant_allowlist a JOIN models m ON m.id=a.model_id`)
+ if err != nil {
+ return nil, fmt.Errorf("query model tenant allowlist: %w", err)
+ }
+ for tenantRows.Next() {
+ var publicID, tenantID string
+ if err := tenantRows.Scan(&publicID, &tenantID); err != nil {
+ tenantRows.Close()
+ return nil, err
+ }
+ if position, ok := byID[publicID]; ok {
+ if models[position].AllowedTenantIDs == nil {
+ models[position].AllowedTenantIDs = map[string]struct{}{}
+ }
+ models[position].AllowedTenantIDs[tenantID] = struct{}{}
+ }
+ }
+ tenantRows.Close()
+ keyRows, err := tx.Query(ctx, `SELECT m.public_id, a.api_key_id::text FROM api_key_model_allowlist a JOIN models m ON m.id=a.model_id`)
+ if err != nil {
+ return nil, fmt.Errorf("query model key allowlist: %w", err)
+ }
+ for keyRows.Next() {
+ var publicID, keyID string
+ if err := keyRows.Scan(&publicID, &keyID); err != nil {
+ keyRows.Close()
+ return nil, err
+ }
+ if position, ok := byID[publicID]; ok {
+ if models[position].AllowedKeyIDs == nil {
+ models[position].AllowedKeyIDs = map[string]struct{}{}
+ }
+ models[position].AllowedKeyIDs[keyID] = struct{}{}
+ }
+ }
+ keyRows.Close()
return models, nil
}
diff --git a/internal/controlplane/store.go b/internal/controlplane/store.go
index 833f846..fb48d78 100644
--- a/internal/controlplane/store.go
+++ b/internal/controlplane/store.go
@@ -2,7 +2,9 @@ package controlplane
import (
"context"
+ "crypto/sha256"
_ "embed"
+ "encoding/hex"
"encoding/json"
"errors"
"fmt"
@@ -11,6 +13,7 @@ import (
"aigw/internal/security"
+ "github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/redis/go-redis/v9"
)
@@ -20,12 +23,15 @@ var schemaSQL string
var ErrRedisDisabled = errors.New("Redis propagation is disabled")
+const migrationVersion int64 = 2026080504
+
type Options struct {
- DatabaseURL string
- RedisURL string
- CredentialKey string
- RedisChannel string
- VersionCacheKey string
+ DatabaseURL string
+ RedisURL string
+ CredentialKey string
+ PreviousCredentialKeys []string
+ RedisChannel string
+ VersionCacheKey string
}
type Store struct {
@@ -37,7 +43,7 @@ type Store struct {
}
func NewStore(ctx context.Context, options Options) (*Store, error) {
- cipher, err := security.NewCredentialCipher(options.CredentialKey)
+ cipher, err := security.NewCredentialKeyring(options.CredentialKey, options.PreviousCredentialKeys)
if err != nil {
return nil, err
}
@@ -76,11 +82,10 @@ func (s *Store) RedisEnabled() bool {
return s.redis != nil
}
+func (s *Store) Ping(ctx context.Context) error { return s.db.Ping(ctx) }
+
func (s *Store) Migrate(ctx context.Context) error {
- if _, err := s.db.Exec(ctx, schemaSQL); err != nil {
- return fmt.Errorf("apply control-plane schema: %w", err)
- }
- return nil
+ return applySchema(ctx, s.db)
}
func MigrateDatabase(ctx context.Context, databaseURL string) error {
@@ -89,12 +94,67 @@ func MigrateDatabase(ctx context.Context, databaseURL string) error {
return fmt.Errorf("configure PostgreSQL: %w", err)
}
defer db.Close()
- if _, err := db.Exec(ctx, schemaSQL); err != nil {
+ return applySchema(ctx, db)
+}
+
+func MigrationStatusDatabase(ctx context.Context, databaseURL string) (MigrationStatus, error) {
+ db, err := pgxpool.New(ctx, databaseURL)
+ if err != nil {
+ return MigrationStatus{}, fmt.Errorf("configure PostgreSQL: %w", err)
+ }
+ defer db.Close()
+ var result MigrationStatus
+ err = db.QueryRow(ctx, `SELECT version,name,checksum,applied_at FROM schema_migrations ORDER BY version DESC LIMIT 1`).Scan(&result.Version, &result.Name, &result.Checksum, &result.AppliedAt)
+ if err != nil {
+ return MigrationStatus{}, fmt.Errorf("read migration status: %w", err)
+ }
+ return result, nil
+}
+
+func applySchema(ctx context.Context, db *pgxpool.Pool) error {
+ 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 {
+ return err
+ }
+ if _, err := tx.Exec(ctx, schemaSQL); err != nil {
return fmt.Errorf("apply control-plane schema: %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 !errors.Is(err, pgx.ErrNoRows) && err != nil {
+ return err
+ }
+ if _, err := tx.Exec(ctx, `INSERT INTO schema_migrations(version,name,checksum) VALUES ($1,$2,$3) ON CONFLICT DO NOTHING`, migrationVersion, "commercial-control-plane", checksum); err != nil {
+ return err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return fmt.Errorf("commit migration: %w", err)
+ }
return nil
}
+type MigrationStatus struct {
+ Version int64 `json:"version"`
+ Name string `json:"name"`
+ Checksum string `json:"checksum"`
+ AppliedAt time.Time `json:"applied_at"`
+}
+
+func (s *Store) MigrationStatus(ctx context.Context) (MigrationStatus, error) {
+ var result MigrationStatus
+ err := s.db.QueryRow(ctx, `SELECT version,name,checksum,applied_at FROM schema_migrations ORDER BY version DESC LIMIT 1`).Scan(&result.Version, &result.Name, &result.Checksum, &result.AppliedAt)
+ return result, err
+}
+
func (s *Store) DatabaseGeneration(ctx context.Context) (int64, error) {
var generation int64
err := s.db.QueryRow(ctx, `SELECT generation FROM control_state WHERE singleton = TRUE`).Scan(&generation)
diff --git a/internal/controlplane/store_integration_test.go b/internal/controlplane/store_integration_test.go
new file mode 100644
index 0000000..4f55e10
--- /dev/null
+++ b/internal/controlplane/store_integration_test.go
@@ -0,0 +1,66 @@
+package controlplane
+
+import (
+ "context"
+ "fmt"
+ "net/url"
+ "os"
+ "testing"
+ "time"
+
+ "github.com/jackc/pgx/v5/pgxpool"
+)
+
+func TestMigrationUpgradesPreviousVersionAndIsIdempotentPostgres(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()
+ db, err := pgxpool.New(ctx, databaseURL)
+ if err != nil {
+ t.Fatal(err)
+ }
+ defer db.Close()
+ schema := fmt.Sprintf("migration_drill_%d", time.Now().UnixNano())
+ if _, err := db.Exec(ctx, "CREATE SCHEMA "+schema); err != nil {
+ t.Fatal(err)
+ }
+ t.Cleanup(func() { _, _ = db.Exec(context.Background(), "DROP SCHEMA "+schema+" CASCADE") })
+ if _, err := db.Exec(ctx, `CREATE TABLE `+schema+`.schema_migrations (
+ version BIGINT PRIMARY KEY,name TEXT NOT NULL,checksum TEXT NOT NULL,applied_at TIMESTAMPTZ NOT NULL DEFAULT now())`); err != nil {
+ t.Fatal(err)
+ }
+ if _, err := db.Exec(ctx, `INSERT INTO `+schema+`.schema_migrations(version,name,checksum) VALUES ($1,'previous-release','immutable-previous-checksum')`, migrationVersion-1); err != nil {
+ t.Fatal(err)
+ }
+ 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)
+ }
+ if err := MigrateDatabase(ctx, isolatedURL); err != nil {
+ t.Fatalf("second migration must be idempotent: %v", err)
+ }
+ status, err := MigrationStatusDatabase(ctx, isolatedURL)
+ if err != nil {
+ t.Fatal(err)
+ }
+ if status.Version != migrationVersion {
+ t.Fatalf("migration version = %d, want %d", status.Version, migrationVersion)
+ }
+ var count int
+ if err := db.QueryRow(ctx, `SELECT count(*) FROM `+schema+`.schema_migrations`).Scan(&count); err != nil {
+ t.Fatal(err)
+ }
+ if count != 2 {
+ t.Fatalf("migration history contains %d rows, want previous and current", count)
+ }
+}
diff --git a/internal/controlplane/types.go b/internal/controlplane/types.go
index f81e63b..c24dde9 100644
--- a/internal/controlplane/types.go
+++ b/internal/controlplane/types.go
@@ -7,6 +7,12 @@ import (
"aigw/internal/domain"
)
+type MailQueueStatus struct {
+ Backlog int64 `json:"backlog"`
+ Failed int64 `json:"failed"`
+ OldestPending *time.Time `json:"oldest_pending,omitempty"`
+}
+
type Tenant struct {
ID string `json:"id"`
Slug string `json:"slug"`
@@ -62,16 +68,36 @@ type Route struct {
}
type Model struct {
- ID string `json:"id"`
- PublicID string `json:"public_id"`
- OwnedBy string `json:"owned_by"`
- InputPriceMicrosPerMillion int64 `json:"input_price_micros_per_million"`
- OutputPriceMicrosPerMillion int64 `json:"output_price_micros_per_million"`
- CacheReadPriceMicrosPerMillion int64 `json:"cache_read_price_micros_per_million"`
- CacheWritePriceMicrosPerMillion int64 `json:"cache_write_price_micros_per_million"`
- Enabled bool `json:"enabled"`
- Routes []Route `json:"routes"`
- CreatedAt time.Time `json:"created_at"`
+ ID string `json:"id"`
+ PublicID string `json:"public_id"`
+ DisplayName string `json:"display_name"`
+ Description string `json:"description"`
+ OwnedBy string `json:"owned_by"`
+ InputModalities []string `json:"input_modalities"`
+ OutputModalities []string `json:"output_modalities"`
+ ContextWindow int64 `json:"context_window"`
+ MaxOutputTokens int64 `json:"max_output_tokens"`
+ Capabilities []string `json:"capabilities"`
+ Regions []string `json:"regions"`
+ Lifecycle string `json:"lifecycle"`
+ ReleasedAt *time.Time `json:"released_at,omitempty"`
+ DeprecatedAt *time.Time `json:"deprecated_at,omitempty"`
+ RetiredAt *time.Time `json:"retired_at,omitempty"`
+ ReplacementModel string `json:"replacement_model,omitempty"`
+ Aliases []string `json:"aliases"`
+ AllowedTenantIDs []string `json:"allowed_tenant_ids"`
+ AllowedKeyIDs []string `json:"allowed_key_ids"`
+ PriceVersionID string `json:"price_version_id"`
+ PriceVersion int `json:"price_version"`
+ PriceCurrency string `json:"price_currency"`
+ PriceEffectiveFrom time.Time `json:"price_effective_from"`
+ InputPriceMicrosPerMillion int64 `json:"input_price_micros_per_million"`
+ OutputPriceMicrosPerMillion int64 `json:"output_price_micros_per_million"`
+ CacheReadPriceMicrosPerMillion int64 `json:"cache_read_price_micros_per_million"`
+ CacheWritePriceMicrosPerMillion int64 `json:"cache_write_price_micros_per_million"`
+ Enabled bool `json:"enabled"`
+ Routes []Route `json:"routes"`
+ CreatedAt time.Time `json:"created_at"`
}
type Overview struct {
@@ -141,7 +167,24 @@ type RouteInput struct {
type CreateModelInput struct {
PublicID string `json:"public_id"`
+ DisplayName string `json:"display_name"`
+ Description string `json:"description"`
OwnedBy string `json:"owned_by"`
+ InputModalities []string `json:"input_modalities"`
+ OutputModalities []string `json:"output_modalities"`
+ ContextWindow int64 `json:"context_window"`
+ MaxOutputTokens int64 `json:"max_output_tokens"`
+ Capabilities []string `json:"capabilities"`
+ Regions []string `json:"regions"`
+ Lifecycle string `json:"lifecycle"`
+ ReleasedAt *time.Time `json:"released_at"`
+ DeprecatedAt *time.Time `json:"deprecated_at"`
+ RetiredAt *time.Time `json:"retired_at"`
+ ReplacementModel string `json:"replacement_model"`
+ Aliases []string `json:"aliases"`
+ AllowedTenantIDs []string `json:"allowed_tenant_ids"`
+ AllowedKeyIDs []string `json:"allowed_key_ids"`
+ PriceCurrency string `json:"price_currency"`
InputPriceMicrosPerMillion int64 `json:"input_price_micros_per_million"`
OutputPriceMicrosPerMillion int64 `json:"output_price_micros_per_million"`
CacheReadPriceMicrosPerMillion int64 `json:"cache_read_price_micros_per_million"`
@@ -149,6 +192,15 @@ type CreateModelInput struct {
Routes []RouteInput `json:"routes"`
}
+type CreatePriceVersionInput struct {
+ Currency string `json:"currency"`
+ InputPriceMicrosPerMillion int64 `json:"input_price_micros_per_million"`
+ OutputPriceMicrosPerMillion int64 `json:"output_price_micros_per_million"`
+ CacheReadPriceMicrosPerMillion int64 `json:"cache_read_price_micros_per_million"`
+ CacheWritePriceMicrosPerMillion int64 `json:"cache_write_price_micros_per_million"`
+ EffectiveFrom time.Time `json:"effective_from"`
+}
+
type ConsoleActor struct {
ID string `json:"id,omitempty"`
TenantID string `json:"tenant_id,omitempty"`
diff --git a/internal/controlplane/usage.go b/internal/controlplane/usage.go
index 6c69a9f..b436d8c 100644
--- a/internal/controlplane/usage.go
+++ b/internal/controlplane/usage.go
@@ -29,13 +29,13 @@ func (s *Store) RecordUsage(ctx context.Context, event domain.UsageEvent) error
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,
- cost_micros, charged_micros, uncollected_micros)
- 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)
+ 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)
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.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens)
+ event.Usage.CacheCreationInputTokens, event.Usage.CacheReadInputTokens, event.UsageReported, usageMeteringStatus(event))
if err != nil {
return fmt.Errorf("persist usage event: %w", err)
}
@@ -51,6 +51,16 @@ func (s *Store) RecordUsage(ctx context.Context, event domain.UsageEvent) error
return nil
}
+func usageMeteringStatus(event domain.UsageEvent) string {
+ if event.StatusCode < 200 || event.StatusCode >= 300 || !event.Success {
+ return "upstream_failed"
+ }
+ if event.UsageReported {
+ return "reported"
+ }
+ return "missing"
+}
+
func upsertUsageRollup(ctx context.Context, tx pgx.Tx, event domain.UsageEvent, cost, charged, uncollected int64) error {
period := time.Date(event.StartedAt.UTC().Year(), event.StartedAt.UTC().Month(), 1, 0, 0, 0, 0, time.UTC)
_, err := tx.Exec(ctx, `