package controlplane import ( "context" "crypto/sha256" "encoding/json" "fmt" "strings" "time" "aigw/internal/auth" "aigw/internal/domain" "github.com/jackc/pgx/v5" ) func (s *Store) LoadSnapshot(ctx context.Context) (Snapshot, error) { tx, err := s.db.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly}) if err != nil { return Snapshot{}, fmt.Errorf("begin snapshot transaction: %w", err) } defer tx.Rollback(ctx) var result Snapshot if err := tx.QueryRow(ctx, `SELECT generation FROM control_state WHERE singleton = TRUE`).Scan(&result.Generation); err != nil { return Snapshot{}, fmt.Errorf("read snapshot generation: %w", err) } providers, err := s.loadProviders(ctx, tx) if err != nil { return Snapshot{}, err } result.Models, err = loadModels(ctx, tx, providers) if err != nil { return Snapshot{}, err } result.APIKeys, err = loadAPIKeys(ctx, tx) if err != nil { return Snapshot{}, err } result.Limits, err = loadLimitPolicies(ctx, tx) if err != nil { return Snapshot{}, err } if err := tx.Commit(ctx); err != nil { return Snapshot{}, fmt.Errorf("commit snapshot transaction: %w", err) } return result, nil } func loadLimitPolicies(ctx context.Context, tx pgx.Tx) ([]domain.LimitPolicy, error) { rows, err := tx.Query(ctx, ` SELECT tenant_id::text, project_id::text, requests_per_minute, tokens_per_minute, concurrent_requests, monthly_spend_micros FROM project_limits`) if err != nil { return nil, fmt.Errorf("query project limits: %w", err) } defer rows.Close() result := make([]domain.LimitPolicy, 0) for rows.Next() { var policy domain.LimitPolicy if err := rows.Scan(&policy.TenantID, &policy.ProjectID, &policy.RequestsPerMinute, &policy.TokensPerMinute, &policy.Concurrent, &policy.MonthlySpendMicros); err != nil { return nil, fmt.Errorf("scan project limit: %w", err) } result = append(result, policy) } if err := rows.Err(); err != nil { return nil, fmt.Errorf("read project limits: %w", err) } return result, nil } func (s *Store) loadProviders(ctx context.Context, tx pgx.Tx) (map[string]domain.Provider, error) { rows, err := tx.Query(ctx, ` SELECT id::text, name, protocol, base_url, api_key_ciphertext FROM providers WHERE enabled = TRUE ORDER BY name`) if err != nil { return nil, fmt.Errorf("query providers: %w", err) } defer rows.Close() providers := make(map[string]domain.Provider) for rows.Next() { var id, name, protocol, baseURL string var ciphertext []byte if err := rows.Scan(&id, &name, &protocol, &baseURL, &ciphertext); err != nil { return nil, fmt.Errorf("scan provider: %w", err) } apiKey, err := s.cipher.Decrypt(ciphertext) if err != nil { return nil, fmt.Errorf("decrypt provider %q credential: %w", name, err) } providers[id] = domain.Provider{ID: id, Protocol: domain.Protocol(protocol), BaseURL: strings.TrimRight(baseURL, "/"), APIKey: apiKey} } if err := rows.Err(); err != nil { return nil, fmt.Errorf("read providers: %w", err) } return providers, nil } func loadModels(ctx context.Context, tx pgx.Tx, providers map[string]domain.Provider) ([]domain.Model, error) { rows, err := tx.Query(ctx, ` 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 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) } defer rows.Close() models := make([]domain.Model, 0) index := make(map[string]int) for rows.Next() { 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(&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] if !ok { continue } 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, 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, }) } 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 } func loadAPIKeys(ctx context.Context, tx pgx.Tx) ([]auth.HashedKeyRecord, error) { rows, err := tx.Query(ctx, ` SELECT k.id::text, k.key_hash, k.tenant_id::text, k.project_id::text, k.scopes FROM api_keys k JOIN tenants t ON t.id = k.tenant_id AND t.status = 'active' JOIN projects p ON p.id = k.project_id AND p.status = 'active' WHERE k.status = 'active'`) if err != nil { return nil, fmt.Errorf("query API keys: %w", err) } defer rows.Close() records := make([]auth.HashedKeyRecord, 0) for rows.Next() { var keyID, tenantID, projectID string var hashBytes, scopesJSON []byte if err := rows.Scan(&keyID, &hashBytes, &tenantID, &projectID, &scopesJSON); err != nil { return nil, fmt.Errorf("scan API key: %w", err) } if len(hashBytes) != sha256.Size { return nil, fmt.Errorf("API key %s has invalid hash length", keyID) } var hash [sha256.Size]byte copy(hash[:], hashBytes) var scopes []string if err := json.Unmarshal(scopesJSON, &scopes); err != nil { return nil, fmt.Errorf("decode API key %s scopes: %w", keyID, err) } records = append(records, auth.HashedKeyRecord{Hash: hash, Principal: domain.Principal{ KeyID: keyID, TenantID: tenantID, ProjectID: projectID, Scopes: scopes, }}) } if err := rows.Err(); err != nil { return nil, fmt.Errorf("read API keys: %w", err) } return records, nil }