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.go82
1 files changed, 68 insertions, 14 deletions
diff --git a/internal/controlplane/mutations.go b/internal/controlplane/mutations.go
index c2cb7d8..9f81d6e 100644
--- a/internal/controlplane/mutations.go
+++ b/internal/controlplane/mutations.go
@@ -17,8 +17,9 @@ import (
)
var (
- ErrNotFound = errors.New("control-plane resource not found")
- slugPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{1,62}[a-z0-9]$`)
+ 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) {
@@ -84,11 +85,28 @@ 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 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)
@@ -104,15 +122,32 @@ 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)
- 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)
+ 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 {
@@ -129,10 +164,29 @@ func (s *Store) RevokeAPIKey(ctx context.Context, id string) (int64, error) {
}
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), "/")
- 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")
+ 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") {
@@ -149,11 +203,11 @@ func (s *Store) CreateProvider(ctx context.Context, input CreateProviderInput) (
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)
+ 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)
}