summaryrefslogtreecommitdiff
path: root/internal/controlplane/mutations.go
diff options
context:
space:
mode:
authorChia <Chia@93.nz>2026-08-05 22:01:29 +1200
committerChia <Chia@93.nz>2026-08-05 22:07:50 +1200
commiteadb2ffe85c43cf6fc741c9823cd28eedb4a844c (patch)
tree1aba2536d57360da403aa35c9ced58b615c7064e /internal/controlplane/mutations.go
parentcd0dd91ab93653631904f2ea0e574ccde6d60339 (diff)
feat: harden prepaid billing and commercial operations
Diffstat (limited to 'internal/controlplane/mutations.go')
-rw-r--r--internal/controlplane/mutations.go143
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(&currentEffectiveFrom)
+ 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)
}