package controlplane import ( "context" "crypto/rand" "crypto/sha256" "encoding/base64" "encoding/json" "errors" "fmt" "net/url" "regexp" "strings" "time" "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]$`) nonSlugCharacters = regexp.MustCompile(`[^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.Name) > 120 || input.MonthlySpendMicros < 0 { return CreatedAPIKey{}, 0, errors.New("API key name or monthly spend limit is invalid") } if input.ExpiresAt != nil && !input.ExpiresAt.After(time.Now()) { return CreatedAPIKey{}, 0, errors.New("API key expiry must be in the future") } if len(input.Scopes) == 0 { input.Scopes = []string{"inference"} } scopes := uniqueStrings(input.Scopes) tags := uniqueStrings(input.Tags) allowedModels := uniqueStrings(input.AllowedModels) if len(scopes) > 20 || len(tags) > 20 || len(allowedModels) > 200 { return CreatedAPIKey{}, 0, errors.New("API key has too many scopes, tags, or model restrictions") } for _, value := range append(append(append([]string{}, scopes...), tags...), allowedModels...) { if len(value) > 160 { return CreatedAPIKey{}, 0, errors.New("API key scope, tag, or model ID is too long") } } scopesJSON, _ := json.Marshal(scopes) tagsJSON, _ := json.Marshal(tags) 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, tags, monthly_spend_micros, expires_at) VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9) RETURNING id::text, tenant_id::text, project_id::text, name, key_prefix, scopes, tags, monthly_spend_micros, status, expires_at, last_used_at, created_at`, input.TenantID, input.ProjectID, input.Name, prefix, hash[:], scopesJSON, tagsJSON, input.MonthlySpendMicros, input.ExpiresAt, ).Scan(&result.ID, &result.TenantID, &result.ProjectID, &result.Name, &result.KeyPrefix, &scopesJSON, &tagsJSON, &result.MonthlySpendMicros, &result.Status, &result.ExpiresAt, &result.LastUsedAt, &result.CreatedAt) if err != nil { return CreatedAPIKey{}, 0, fmt.Errorf("create API key: %w", err) } if len(allowedModels) > 0 { command, err := tx.Exec(ctx, ` INSERT INTO api_key_model_restrictions (api_key_id, model_id) SELECT $1, id FROM models WHERE public_id = ANY($2::text[])`, result.ID, allowedModels) if err != nil { return CreatedAPIKey{}, 0, fmt.Errorf("restrict API key models: %w", err) } if command.RowsAffected() != int64(len(allowedModels)) { return CreatedAPIKey{}, 0, errors.New("one or more allowed model IDs do not exist") } } result.Scopes = scopes result.Tags = tags result.AllowedModels = allowedModels 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.Slug = strings.ToLower(strings.TrimSpace(input.Slug)) input.Name = strings.TrimSpace(input.Name) if input.Slug == "" { input.Slug = strings.Trim(nonSlugCharacters.ReplaceAllString(strings.ToLower(input.Name), "-"), "-") if len(input.Slug) > 64 { input.Slug = strings.TrimRight(input.Slug[:64], "-") } } input.BaseURL = strings.TrimRight(strings.TrimSpace(input.BaseURL), "/") input.WireAPI = strings.TrimSpace(input.WireAPI) if input.WireAPI == "" { if input.Protocol == "anthropic" { input.WireAPI = "messages" } else { input.WireAPI = "chat_completions" } } if !slugPattern.MatchString(input.Slug) || input.Name == "" || input.APIKey == "" || (input.Protocol != "openai" && input.Protocol != "anthropic") { return Provider{}, 0, errors.New("provider requires a unique 3-64 character lowercase slug, name, protocol openai|anthropic, base_url, and api_key") } if (input.Protocol == "openai" && input.WireAPI != "chat_completions" && input.WireAPI != "responses") || (input.Protocol == "anthropic" && input.WireAPI != "messages") { return Provider{}, 0, errors.New("provider wire_api is incompatible with protocol") } 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 (slug, name, protocol, wire_api, base_url, api_key_ciphertext) VALUES ($1, $2, $3, $4, $5, $6) RETURNING id::text, slug, name, protocol, wire_api, base_url, enabled, created_at`, input.Slug, input.Name, input.Protocol, input.WireAPI, input.BaseURL, ciphertext, ).Scan(&result.ID, &result.Slug, &result.Name, &result.Protocol, &result.WireAPI, &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.DisplayName = strings.TrimSpace(input.DisplayName) input.Description = strings.TrimSpace(input.Description) input.OwnedBy = strings.TrimSpace(input.OwnedBy) input.Lifecycle = strings.ToLower(strings.TrimSpace(input.Lifecycle)) input.PriceCurrency = strings.ToLower(strings.TrimSpace(input.PriceCurrency)) if input.DisplayName == "" { input.DisplayName = input.PublicID } if input.Lifecycle == "" { input.Lifecycle = "active" } if input.PriceCurrency == "" { input.PriceCurrency = "usd" } if len(input.InputModalities) == 0 { input.InputModalities = []string{"text"} } if len(input.OutputModalities) == 0 { input.OutputModalities = []string{"text"} } if len(input.Capabilities) == 0 { input.Capabilities = []string{"chat", "streaming"} } input.InputModalities = uniqueStrings(input.InputModalities) input.OutputModalities = uniqueStrings(input.OutputModalities) input.Capabilities = uniqueStrings(input.Capabilities) input.Regions = uniqueStrings(input.Regions) input.Aliases = uniqueStrings(input.Aliases) input.AllowedTenantIDs = uniqueStrings(input.AllowedTenantIDs) input.AllowedKeyIDs = uniqueStrings(input.AllowedKeyIDs) if input.PublicID == "" || len(input.Routes) == 0 { return Model{}, 0, errors.New("model requires public_id and at least one route") } if input.ContextWindow < 0 || input.MaxOutputTokens < 0 || len(input.PriceCurrency) != 3 { return Model{}, 0, errors.New("model context, output limit, or price currency is invalid") } if input.Lifecycle != "preview" && input.Lifecycle != "active" && input.Lifecycle != "deprecated" && input.Lifecycle != "retired" { return Model{}, 0, errors.New("model lifecycle must be preview, active, deprecated, or retired") } if input.InputPriceMicrosPerMillion < 0 || input.OutputPriceMicrosPerMillion < 0 || input.CacheReadPriceMicrosPerMillion < 0 || input.CacheWritePriceMicrosPerMillion < 0 { return Model{}, 0, errors.New("model prices cannot be negative") } 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) inputModalitiesJSON, _ := json.Marshal(input.InputModalities) outputModalitiesJSON, _ := json.Marshal(input.OutputModalities) capabilitiesJSON, _ := json.Marshal(input.Capabilities) regionsJSON, _ := json.Marshal(input.Regions) var result Model err = tx.QueryRow(ctx, ` INSERT INTO models (public_id, display_name, description, owned_by, input_modalities, output_modalities, context_window, max_output_tokens, capabilities, regions, lifecycle, released_at, deprecated_at, retired_at, replacement_model, input_price_micros_per_million, output_price_micros_per_million, cache_read_price_micros_per_million, cache_write_price_micros_per_million) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,NULLIF($15,''),$16,$17,$18,$19) RETURNING id::text, public_id, display_name, description, owned_by, input_price_micros_per_million, output_price_micros_per_million, cache_read_price_micros_per_million, cache_write_price_micros_per_million, enabled, created_at`, input.PublicID, input.DisplayName, input.Description, input.OwnedBy, inputModalitiesJSON, outputModalitiesJSON, input.ContextWindow, input.MaxOutputTokens, capabilitiesJSON, regionsJSON, input.Lifecycle, input.ReleasedAt, input.DeprecatedAt, input.RetiredAt, input.ReplacementModel, input.InputPriceMicrosPerMillion, input.OutputPriceMicrosPerMillion, input.CacheReadPriceMicrosPerMillion, input.CacheWritePriceMicrosPerMillion, ).Scan(&result.ID, &result.PublicID, &result.DisplayName, &result.Description, &result.OwnedBy, &result.InputPriceMicrosPerMillion, &result.OutputPriceMicrosPerMillion, &result.CacheReadPriceMicrosPerMillion, &result.CacheWritePriceMicrosPerMillion, &result.Enabled, &result.CreatedAt) if err != nil { return Model{}, 0, fmt.Errorf("create model: %w", err) } result.InputModalities, result.OutputModalities = input.InputModalities, input.OutputModalities result.ContextWindow, result.MaxOutputTokens = input.ContextWindow, input.MaxOutputTokens result.Capabilities, result.Regions, result.Lifecycle = input.Capabilities, input.Regions, input.Lifecycle result.ReleasedAt, result.DeprecatedAt, result.RetiredAt = input.ReleasedAt, input.DeprecatedAt, input.RetiredAt result.ReplacementModel, result.Aliases = input.ReplacementModel, input.Aliases result.AllowedTenantIDs, result.AllowedKeyIDs = input.AllowedTenantIDs, input.AllowedKeyIDs result.PriceCurrency, result.PriceVersion = input.PriceCurrency, 1 err = tx.QueryRow(ctx, `INSERT INTO model_price_versions (model_id,version,currency, input_price_micros_per_million,output_price_micros_per_million, cache_read_price_micros_per_million,cache_write_price_micros_per_million) VALUES ($1,1,$2,$3,$4,$5,$6) RETURNING id::text,effective_from`, result.ID, input.PriceCurrency, input.InputPriceMicrosPerMillion, input.OutputPriceMicrosPerMillion, input.CacheReadPriceMicrosPerMillion, input.CacheWritePriceMicrosPerMillion, ).Scan(&result.PriceVersionID, &result.PriceEffectiveFrom) if err != nil { return Model{}, 0, fmt.Errorf("create model price version: %w", err) } for _, alias := range input.Aliases { if _, err := tx.Exec(ctx, `INSERT INTO model_aliases (alias,model_id,deprecated) VALUES ($1,$2,true)`, alias, result.ID); err != nil { return Model{}, 0, fmt.Errorf("create model alias: %w", err) } } for _, tenantID := range input.AllowedTenantIDs { if _, err := tx.Exec(ctx, `INSERT INTO model_tenant_allowlist (model_id,tenant_id) VALUES ($1,$2)`, result.ID, tenantID); err != nil { return Model{}, 0, fmt.Errorf("create tenant model allowlist: %w", err) } } for _, keyID := range input.AllowedKeyIDs { if _, err := tx.Exec(ctx, `INSERT INTO api_key_model_allowlist (api_key_id,model_id) VALUES ($1,$2)`, keyID, result.ID); err != nil { return Model{}, 0, fmt.Errorf("create key model allowlist: %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) CreateModelPriceVersion(ctx context.Context, modelID string, input CreatePriceVersionInput) (int64, error) { input.Currency = strings.ToLower(strings.TrimSpace(input.Currency)) if len(input.Currency) != 3 || input.InputPriceMicrosPerMillion < 0 || input.OutputPriceMicrosPerMillion < 0 || input.CacheReadPriceMicrosPerMillion < 0 || input.CacheWritePriceMicrosPerMillion < 0 { return 0, errors.New("price version has invalid currency or negative price") } if input.EffectiveFrom.IsZero() { input.EffectiveFrom = time.Now().UTC() } tx, err := s.db.Begin(ctx) if err != nil { return 0, err } defer tx.Rollback(ctx) // Lock the model row before allocating the next version. PostgreSQL does // not allow FOR UPDATE on an aggregate query. var modelExists bool if err := tx.QueryRow(ctx, `SELECT true FROM models WHERE id=$1 FOR UPDATE`, modelID).Scan(&modelExists); errors.Is(err, pgx.ErrNoRows) { return 0, ErrNotFound } else if err != nil { return 0, err } var currentEffectiveFrom time.Time err = tx.QueryRow(ctx, `SELECT effective_from FROM model_price_versions WHERE model_id=$1 AND effective_to IS NULL`, modelID).Scan(¤tEffectiveFrom) if err != nil && !errors.Is(err, pgx.ErrNoRows) { return 0, err } if err == nil && !input.EffectiveFrom.After(currentEffectiveFrom) { return 0, errors.New("new price version must become effective after the current open version") } var version int if err := tx.QueryRow(ctx, `SELECT COALESCE(max(version),0)+1 FROM model_price_versions WHERE model_id=$1`, modelID).Scan(&version); err != nil { return 0, err } if _, err := tx.Exec(ctx, `UPDATE model_price_versions SET effective_to=$2 WHERE model_id=$1 AND effective_to IS NULL AND effective_from < $2`, modelID, input.EffectiveFrom); err != nil { return 0, err } if _, err := tx.Exec(ctx, `INSERT INTO model_price_versions (model_id,version,currency,input_price_micros_per_million, output_price_micros_per_million,cache_read_price_micros_per_million,cache_write_price_micros_per_million,effective_from) VALUES ($1,$2,$3,$4,$5,$6,$7,$8)`, modelID, version, input.Currency, input.InputPriceMicrosPerMillion, input.OutputPriceMicrosPerMillion, input.CacheReadPriceMicrosPerMillion, input.CacheWritePriceMicrosPerMillion, input.EffectiveFrom); err != nil { return 0, err } 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 (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 }