From 5b651488b081b65fda8a323f228e139adb79a35d Mon Sep 17 00:00:00 2001 From: Chia Date: Tue, 4 Aug 2026 19:58:52 +1200 Subject: Build AI gateway control plane and admin UI --- internal/controlplane/mutations.go | 281 +++++++++++++++++++++++++++++++++++++ 1 file changed, 281 insertions(+) create mode 100644 internal/controlplane/mutations.go (limited to 'internal/controlplane/mutations.go') 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 +} -- cgit v1.2.3