summaryrefslogtreecommitdiff
path: root/internal/controlplane/snapshot.go
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-04 19:58:52 +1200
committerChia <Chia@93.nz>2026-08-04 20:43:23 +1200
commit5b651488b081b65fda8a323f228e139adb79a35d (patch)
tree08baf40efb8fe103b32721cd991ff712323e3173 /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.go150
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
+}