package controlplane import ( "context" "crypto/sha256" "encoding/json" "fmt" "strings" "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 } if err := tx.Commit(ctx); err != nil { return Snapshot{}, fmt.Errorf("commit snapshot transaction: %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.public_id, m.owned_by, r.provider_id::text, r.upstream_model, r.priority, r.weight FROM models m 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 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 publicID, ownedBy, providerID, upstreamModel string var priority, weight int if err := rows.Scan(&publicID, &ownedBy, &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 { position = len(models) index[publicID] = position models = append(models, domain.Model{ID: publicID, OwnedBy: ownedBy}) } 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) } 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 }