diff options
| author | Chia <Chia@93.nz> | 2026-08-04 19:58:52 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-04 20:43:23 +1200 |
| commit | 5b651488b081b65fda8a323f228e139adb79a35d (patch) | |
| tree | 08baf40efb8fe103b32721cd991ff712323e3173 /internal/controlplane/snapshot.go | |
Build AI gateway control plane and admin UI
Diffstat (limited to 'internal/controlplane/snapshot.go')
| -rw-r--r-- | internal/controlplane/snapshot.go | 150 |
1 files changed, 150 insertions, 0 deletions
diff --git a/internal/controlplane/snapshot.go b/internal/controlplane/snapshot.go new file mode 100644 index 0000000..c8931ab --- /dev/null +++ b/internal/controlplane/snapshot.go @@ -0,0 +1,150 @@ +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 +} |
