summaryrefslogtreecommitdiff
path: root/internal/controlplane/mutations.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/controlplane/mutations.go')
-rw-r--r--internal/controlplane/mutations.go130
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
+}