summaryrefslogtreecommitdiff
path: root/internal/controlplane/mutations.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/mutations.go
Build AI gateway control plane and admin UI
Diffstat (limited to '')
-rw-r--r--internal/controlplane/mutations.go281
1 files changed, 281 insertions, 0 deletions
diff --git a/internal/controlplane/mutations.go b/internal/controlplane/mutations.go
new file mode 100644
index 0000000..a781994
--- /dev/null
+++ b/internal/controlplane/mutations.go
@@ -0,0 +1,281 @@
+package controlplane
+
+import (
+ "context"
+ "crypto/rand"
+ "crypto/sha256"
+ "encoding/base64"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "net/url"
+ "regexp"
+ "strings"
+
+ "github.com/jackc/pgx/v5"
+)
+
+var (
+ ErrNotFound = errors.New("control-plane resource not found")
+ slugPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$`)
+)
+
+func (s *Store) CreateTenant(ctx context.Context, input CreateTenantInput) (Tenant, int64, error) {
+ input.Slug = strings.ToLower(strings.TrimSpace(input.Slug))
+ input.Name = strings.TrimSpace(input.Name)
+ if !slugPattern.MatchString(input.Slug) || input.Name == "" {
+ return Tenant{}, 0, errors.New("tenant requires a 3-64 character lowercase slug and a name")
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return Tenant{}, 0, err
+ }
+ defer tx.Rollback(ctx)
+ var result Tenant
+ err = tx.QueryRow(ctx, `
+ INSERT INTO tenants (slug, name) VALUES ($1, $2)
+ RETURNING id::text, slug, name, status, created_at`, input.Slug, input.Name,
+ ).Scan(&result.ID, &result.Slug, &result.Name, &result.Status, &result.CreatedAt)
+ if err != nil {
+ return Tenant{}, 0, fmt.Errorf("create tenant: %w", err)
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return Tenant{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Tenant{}, 0, err
+ }
+ return result, generation, nil
+}
+
+func (s *Store) CreateProject(ctx context.Context, input CreateProjectInput) (Project, int64, error) {
+ input.Slug = strings.ToLower(strings.TrimSpace(input.Slug))
+ input.Name = strings.TrimSpace(input.Name)
+ if input.TenantID == "" || !slugPattern.MatchString(input.Slug) || input.Name == "" {
+ return Project{}, 0, errors.New("project requires tenant_id, a 3-64 character lowercase slug, and a name")
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return Project{}, 0, err
+ }
+ defer tx.Rollback(ctx)
+ var result Project
+ err = tx.QueryRow(ctx, `
+ INSERT INTO projects (tenant_id, slug, name) VALUES ($1, $2, $3)
+ RETURNING id::text, tenant_id::text, slug, name, status, created_at`, input.TenantID, input.Slug, input.Name,
+ ).Scan(&result.ID, &result.TenantID, &result.Slug, &result.Name, &result.Status, &result.CreatedAt)
+ if err != nil {
+ return Project{}, 0, fmt.Errorf("create project: %w", err)
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return Project{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Project{}, 0, err
+ }
+ return result, generation, nil
+}
+
+func (s *Store) CreateAPIKey(ctx context.Context, input CreateAPIKeyInput) (CreatedAPIKey, int64, error) {
+ input.Name = strings.TrimSpace(input.Name)
+ if input.TenantID == "" || input.ProjectID == "" || input.Name == "" {
+ return CreatedAPIKey{}, 0, errors.New("API key requires tenant_id, project_id, and name")
+ }
+ if len(input.Scopes) == 0 {
+ input.Scopes = []string{"inference"}
+ }
+ scopes := uniqueStrings(input.Scopes)
+ scopesJSON, _ := json.Marshal(scopes)
+ random := make([]byte, 32)
+ if _, err := rand.Read(random); err != nil {
+ return CreatedAPIKey{}, 0, fmt.Errorf("generate API key: %w", err)
+ }
+ rawKey := "sk-aigw-" + base64.RawURLEncoding.EncodeToString(random)
+ hash := sha256.Sum256([]byte(rawKey))
+ prefix := rawKey[:min(18, len(rawKey))] + "..."
+
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return CreatedAPIKey{}, 0, err
+ }
+ defer tx.Rollback(ctx)
+ var result CreatedAPIKey
+ err = tx.QueryRow(ctx, `
+ INSERT INTO api_keys (tenant_id, project_id, name, key_prefix, key_hash, scopes)
+ VALUES ($1, $2, $3, $4, $5, $6)
+ RETURNING id::text, tenant_id::text, project_id::text, name, key_prefix, scopes, status, created_at`,
+ input.TenantID, input.ProjectID, input.Name, prefix, hash[:], scopesJSON,
+ ).Scan(&result.ID, &result.TenantID, &result.ProjectID, &result.Name, &result.KeyPrefix, &scopesJSON, &result.Status, &result.CreatedAt)
+ if err != nil {
+ return CreatedAPIKey{}, 0, fmt.Errorf("create API key: %w", err)
+ }
+ result.Scopes = scopes
+ result.Key = rawKey
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return CreatedAPIKey{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return CreatedAPIKey{}, 0, err
+ }
+ return result, generation, nil
+}
+
+func (s *Store) RevokeAPIKey(ctx context.Context, id string) (int64, error) {
+ return s.toggle(ctx, `UPDATE api_keys SET status = 'revoked', revoked_at = now() WHERE id = $1 AND status <> 'revoked'`, id)
+}
+
+func (s *Store) CreateProvider(ctx context.Context, input CreateProviderInput) (Provider, int64, error) {
+ input.Name = strings.TrimSpace(input.Name)
+ input.BaseURL = strings.TrimRight(strings.TrimSpace(input.BaseURL), "/")
+ if input.Name == "" || input.APIKey == "" || (input.Protocol != "openai" && input.Protocol != "anthropic") {
+ return Provider{}, 0, errors.New("provider requires name, protocol openai|anthropic, base_url, and api_key")
+ }
+ parsed, err := url.Parse(input.BaseURL)
+ if err != nil || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") {
+ return Provider{}, 0, errors.New("provider base_url must be an absolute http(s) URL")
+ }
+ ciphertext, err := s.cipher.Encrypt(input.APIKey)
+ if err != nil {
+ return Provider{}, 0, err
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return Provider{}, 0, err
+ }
+ defer tx.Rollback(ctx)
+ var result Provider
+ err = tx.QueryRow(ctx, `
+ INSERT INTO providers (name, protocol, base_url, api_key_ciphertext)
+ VALUES ($1, $2, $3, $4)
+ RETURNING id::text, name, protocol, base_url, enabled, created_at`,
+ input.Name, input.Protocol, input.BaseURL, ciphertext,
+ ).Scan(&result.ID, &result.Name, &result.Protocol, &result.BaseURL, &result.Enabled, &result.CreatedAt)
+ if err != nil {
+ return Provider{}, 0, fmt.Errorf("create provider: %w", err)
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return Provider{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Provider{}, 0, err
+ }
+ return result, generation, nil
+}
+
+func (s *Store) SetProviderEnabled(ctx context.Context, id string, enabled bool) (int64, error) {
+ return s.toggle(ctx, `UPDATE providers SET enabled = $2, updated_at = now() WHERE id = $1`, id, enabled)
+}
+
+func (s *Store) CreateModel(ctx context.Context, input CreateModelInput) (Model, int64, error) {
+ input.PublicID = strings.TrimSpace(input.PublicID)
+ input.OwnedBy = strings.TrimSpace(input.OwnedBy)
+ if input.PublicID == "" || len(input.Routes) == 0 {
+ return Model{}, 0, errors.New("model requires public_id and at least one route")
+ }
+ for i := range input.Routes {
+ input.Routes[i].ProviderID = strings.TrimSpace(input.Routes[i].ProviderID)
+ input.Routes[i].UpstreamModel = strings.TrimSpace(input.Routes[i].UpstreamModel)
+ if input.Routes[i].Weight == 0 {
+ input.Routes[i].Weight = 1
+ }
+ if input.Routes[i].ProviderID == "" || input.Routes[i].UpstreamModel == "" || input.Routes[i].Priority < 0 || input.Routes[i].Weight < 1 || input.Routes[i].Weight > 100 {
+ return Model{}, 0, fmt.Errorf("route %d has invalid provider, upstream model, priority, or weight", i+1)
+ }
+ }
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return Model{}, 0, err
+ }
+ defer tx.Rollback(ctx)
+ var result Model
+ err = tx.QueryRow(ctx, `
+ INSERT INTO models (public_id, owned_by) VALUES ($1, $2)
+ RETURNING id::text, public_id, owned_by, enabled, created_at`, input.PublicID, input.OwnedBy,
+ ).Scan(&result.ID, &result.PublicID, &result.OwnedBy, &result.Enabled, &result.CreatedAt)
+ if err != nil {
+ return Model{}, 0, fmt.Errorf("create model: %w", err)
+ }
+ result.Routes = make([]Route, 0, len(input.Routes))
+ for _, route := range input.Routes {
+ var created Route
+ err := tx.QueryRow(ctx, `
+ INSERT INTO model_routes (model_id, provider_id, upstream_model, priority, weight)
+ VALUES ($1, $2, $3, $4, $5)
+ RETURNING id::text, provider_id::text, upstream_model, priority, weight, enabled`,
+ result.ID, route.ProviderID, route.UpstreamModel, route.Priority, route.Weight,
+ ).Scan(&created.ID, &created.ProviderID, &created.UpstreamModel, &created.Priority, &created.Weight, &created.Enabled)
+ if err != nil {
+ return Model{}, 0, fmt.Errorf("create model route: %w", err)
+ }
+ result.Routes = append(result.Routes, created)
+ }
+ generation, err := bumpGeneration(ctx, tx)
+ if err != nil {
+ return Model{}, 0, err
+ }
+ if err := tx.Commit(ctx); err != nil {
+ return Model{}, 0, err
+ }
+ return result, 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)
+}
+
+func (s *Store) toggle(ctx context.Context, query, id string, args ...any) (int64, error) {
+ tx, err := s.db.Begin(ctx)
+ if err != nil {
+ return 0, err
+ }
+ defer tx.Rollback(ctx)
+ parameters := append([]any{id}, args...)
+ command, err := tx.Exec(ctx, query, parameters...)
+ if err != nil {
+ return 0, err
+ }
+ if command.RowsAffected() == 0 {
+ return 0, ErrNotFound
+ }
+ 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 bumpGeneration(ctx context.Context, tx pgx.Tx) (int64, error) {
+ var generation int64
+ err := tx.QueryRow(ctx, `
+ UPDATE control_state SET generation = generation + 1, updated_at = now()
+ WHERE singleton = TRUE RETURNING generation`,
+ ).Scan(&generation)
+ if err != nil {
+ return 0, fmt.Errorf("advance control-plane generation: %w", err)
+ }
+ return generation, nil
+}
+
+func uniqueStrings(values []string) []string {
+ seen := make(map[string]struct{}, len(values))
+ result := make([]string, 0, len(values))
+ for _, value := range values {
+ value = strings.TrimSpace(value)
+ if value == "" {
+ continue
+ }
+ if _, exists := seen[value]; exists {
+ continue
+ }
+ seen[value] = struct{}{}
+ result = append(result, value)
+ }
+ return result
+}