diff options
Diffstat (limited to 'internal/controlplane/mutations.go')
| -rw-r--r-- | internal/controlplane/mutations.go | 130 |
1 files changed, 113 insertions, 17 deletions
diff --git a/internal/controlplane/mutations.go b/internal/controlplane/mutations.go index 9f81d6e..aba4d36 100644 --- a/internal/controlplane/mutations.go +++ b/internal/controlplane/mutations.go @@ -85,8 +85,8 @@ func (s *Store) CreateAPIKey(ctx context.Context, input CreateAPIKeyInput) (Crea 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 len(input.Name) > 120 || input.MonthlySpendMicros < 0 || input.DailySpendMicros < 0 || input.RequestsPerMinute < 0 || input.TokensPerMinute < 0 { + return CreatedAPIKey{}, 0, errors.New("API key name, spend limits, or rate limits are invalid") } if input.ExpiresAt != nil && !input.ExpiresAt.After(time.Now()) { return CreatedAPIKey{}, 0, errors.New("API key expiry must be in the future") @@ -107,13 +107,10 @@ func (s *Store) CreateAPIKey(ctx context.Context, input CreateAPIKeyInput) (Crea } 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, prefix, suffix, hash, err := generateAPIKeySecret() + if err != nil { + return CreatedAPIKey{}, 0, 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 { @@ -122,14 +119,17 @@ func (s *Store) CreateAPIKey(ctx context.Context, input CreateAPIKeyInput) (Crea 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, + INSERT INTO api_keys (tenant_id, project_id, name, key_prefix, key_suffix, key_hash, scopes, tags, monthly_spend_micros, + daily_spend_micros, requests_per_minute, tokens_per_minute, expires_at) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13) + RETURNING id::text, tenant_id::text, project_id::text, name, key_prefix, key_suffix, scopes, tags, + monthly_spend_micros, daily_spend_micros, requests_per_minute, tokens_per_minute, + status, expires_at, last_used_at, created_at`, + input.TenantID, input.ProjectID, input.Name, prefix, suffix, hash[:], scopesJSON, tagsJSON, + input.MonthlySpendMicros, input.DailySpendMicros, input.RequestsPerMinute, input.TokensPerMinute, input.ExpiresAt, + ).Scan(&result.ID, &result.TenantID, &result.ProjectID, &result.Name, &result.KeyPrefix, &result.KeySuffix, + &scopesJSON, &tagsJSON, &result.MonthlySpendMicros, &result.DailySpendMicros, + &result.RequestsPerMinute, &result.TokensPerMinute, &result.Status, &result.ExpiresAt, &result.LastUsedAt, &result.CreatedAt) if err != nil { return CreatedAPIKey{}, 0, fmt.Errorf("create API key: %w", err) @@ -163,6 +163,90 @@ 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) DisableAPIKey(ctx context.Context, id string) (int64, error) { + return s.toggle(ctx, `UPDATE api_keys SET status='disabled', disabled_at=now() WHERE id=$1 AND status='active'`, id) +} + +func (s *Store) EnableAPIKey(ctx context.Context, id string) (int64, error) { + return s.toggle(ctx, `UPDATE api_keys SET status='active', disabled_at=NULL WHERE id=$1 AND status='disabled' AND (expires_at IS NULL OR expires_at > now())`, id) +} + +func (s *Store) RotateAPIKey(ctx context.Context, id string) (CreatedAPIKey, int64, error) { + tx, err := s.db.Begin(ctx) + if err != nil { + return CreatedAPIKey{}, 0, err + } + defer tx.Rollback(ctx) + + var tenantID, projectID, name, status string + var scopesJSON, tagsJSON []byte + var monthlySpendMicros, dailySpendMicros, requestsPerMinute, tokensPerMinute int64 + var expiresAt *time.Time + err = tx.QueryRow(ctx, `SELECT tenant_id::text,project_id::text,name,scopes,tags,monthly_spend_micros, + daily_spend_micros,requests_per_minute,tokens_per_minute,expires_at,status + FROM api_keys WHERE id=$1 FOR UPDATE`, id).Scan(&tenantID, &projectID, &name, &scopesJSON, &tagsJSON, + &monthlySpendMicros, &dailySpendMicros, &requestsPerMinute, &tokensPerMinute, &expiresAt, &status) + if errors.Is(err, pgx.ErrNoRows) { + return CreatedAPIKey{}, 0, ErrNotFound + } + if err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("lock API key for rotation: %w", err) + } + if status == "revoked" { + return CreatedAPIKey{}, 0, errors.New("revoked API key cannot be rotated") + } + if expiresAt != nil && !expiresAt.After(time.Now()) { + return CreatedAPIKey{}, 0, errors.New("expired API key cannot be rotated") + } + rawKey, prefix, suffix, hash, err := generateAPIKeySecret() + if err != nil { + return CreatedAPIKey{}, 0, err + } + var result CreatedAPIKey + err = tx.QueryRow(ctx, `INSERT INTO api_keys (tenant_id,project_id,name,key_prefix,key_suffix,key_hash,scopes,tags, + monthly_spend_micros,daily_spend_micros,requests_per_minute,tokens_per_minute,expires_at,rotated_from_id) + VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14) + RETURNING id::text,tenant_id::text,project_id::text,name,key_prefix,key_suffix,scopes,tags,monthly_spend_micros, + daily_spend_micros,requests_per_minute,tokens_per_minute,status,expires_at,last_used_at,created_at`, + tenantID, projectID, name, prefix, suffix, hash[:], scopesJSON, tagsJSON, monthlySpendMicros, dailySpendMicros, + requestsPerMinute, tokensPerMinute, expiresAt, id).Scan(&result.ID, &result.TenantID, &result.ProjectID, + &result.Name, &result.KeyPrefix, &result.KeySuffix, &scopesJSON, &tagsJSON, &result.MonthlySpendMicros, &result.DailySpendMicros, + &result.RequestsPerMinute, &result.TokensPerMinute, &result.Status, &result.ExpiresAt, &result.LastUsedAt, &result.CreatedAt) + if err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("create rotated API key: %w", err) + } + if _, err := tx.Exec(ctx, `INSERT INTO api_key_model_restrictions (api_key_id,model_id) + SELECT $1,model_id FROM api_key_model_restrictions WHERE api_key_id=$2`, result.ID, id); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("copy rotated API key model restrictions: %w", err) + } + var allowedModelsJSON []byte + if err := tx.QueryRow(ctx, `SELECT COALESCE(jsonb_agg(m.public_id ORDER BY m.public_id),'[]'::jsonb) + FROM api_key_model_restrictions r JOIN models m ON m.id=r.model_id WHERE r.api_key_id=$1`, result.ID).Scan(&allowedModelsJSON); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("read rotated API key model restrictions: %w", err) + } + if _, err := tx.Exec(ctx, `UPDATE api_keys SET status='revoked',revoked_at=now() WHERE id=$1`, id); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("revoke rotated API key: %w", err) + } + if err := json.Unmarshal(scopesJSON, &result.Scopes); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("decode rotated API key scopes: %w", err) + } + if err := json.Unmarshal(tagsJSON, &result.Tags); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("decode rotated API key tags: %w", err) + } + if err := json.Unmarshal(allowedModelsJSON, &result.AllowedModels); err != nil { + return CreatedAPIKey{}, 0, fmt.Errorf("decode rotated API key model restrictions: %w", err) + } + 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) CreateProvider(ctx context.Context, input CreateProviderInput) (Provider, int64, error) { input.Slug = strings.ToLower(strings.TrimSpace(input.Slug)) input.Name = strings.TrimSpace(input.Name) @@ -184,7 +268,7 @@ func (s *Store) CreateProvider(ctx context.Context, input CreateProviderInput) ( 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") || + if (input.Protocol == "openai" && input.WireAPI != "chat_completions" && input.WireAPI != "responses" && input.WireAPI != "embeddings") || (input.Protocol == "anthropic" && input.WireAPI != "messages") { return Provider{}, 0, errors.New("provider wire_api is incompatible with protocol") } @@ -474,3 +558,15 @@ func uniqueStrings(values []string) []string { } return result } + +func generateAPIKeySecret() (string, string, string, [sha256.Size]byte, error) { + random := make([]byte, 32) + if _, err := rand.Read(random); err != nil { + return "", "", "", [sha256.Size]byte{}, 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))] + "..." + suffix := rawKey[max(0, len(rawKey)-6):] + return rawKey, prefix, suffix, hash, nil +} |
