diff options
Diffstat (limited to 'internal/controlplane')
| -rw-r--r-- | internal/controlplane/mail_operations.go | 213 | ||||
| -rw-r--r-- | internal/controlplane/mail_operations_test.go | 29 | ||||
| -rw-r--r-- | internal/controlplane/manager.go | 44 | ||||
| -rw-r--r-- | internal/controlplane/mutations.go | 143 | ||||
| -rw-r--r-- | internal/controlplane/outbox.go | 19 | ||||
| -rw-r--r-- | internal/controlplane/queries.go | 73 | ||||
| -rw-r--r-- | internal/controlplane/retention.go | 57 | ||||
| -rw-r--r-- | internal/controlplane/rotation.go | 69 | ||||
| -rw-r--r-- | internal/controlplane/schema.sql | 265 | ||||
| -rw-r--r-- | internal/controlplane/snapshot.go | 108 | ||||
| -rw-r--r-- | internal/controlplane/store.go | 82 | ||||
| -rw-r--r-- | internal/controlplane/store_integration_test.go | 66 | ||||
| -rw-r--r-- | internal/controlplane/types.go | 72 | ||||
| -rw-r--r-- | internal/controlplane/usage.go | 16 |
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, ¤cy, &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(¤tEffectiveFrom) + 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, ®ionsJSON, &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, ®ionsJSON, &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, ®ions); 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, ` |
