diff options
Diffstat (limited to 'internal/controlplane/mutations.go')
| -rw-r--r-- | internal/controlplane/mutations.go | 143 |
1 files changed, 137 insertions, 6 deletions
diff --git a/internal/controlplane/mutations.go b/internal/controlplane/mutations.go index 3cedf70..c2cb7d8 100644 --- a/internal/controlplane/mutations.go +++ b/internal/controlplane/mutations.go @@ -11,6 +11,7 @@ import ( "net/url" "regexp" "strings" + "time" "github.com/jackc/pgx/v5" ) @@ -172,10 +173,45 @@ func (s *Store) SetProviderEnabled(ctx context.Context, id string, enabled bool) 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") } @@ -194,21 +230,63 @@ func (s *Store) CreateModel(ctx context.Context, input CreateModelInput) (Model, 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, owned_by, 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) - RETURNING id::text, public_id, owned_by, input_price_micros_per_million, output_price_micros_per_million, + 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.OwnedBy, input.InputPriceMicrosPerMillion, input.OutputPriceMicrosPerMillion, + 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.OwnedBy, &result.InputPriceMicrosPerMillion, + ).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 @@ -233,6 +311,59 @@ func (s *Store) CreateModel(ctx context.Context, input CreateModelInput) (Model, 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) } |
