package store import ( "context" "encoding/json" "fmt" "strings" "github.com/jackc/pgx/v5" ) const baseModelColumns = ` id::text, provider_key, canonical_model_key, invocation_name, provider_model_name, model_type, display_name, capabilities, base_billing_config, default_rate_limit_policy, COALESCE(pricing_rule_set_id::text, ''), COALESCE(runtime_policy_set_id::text, ''), runtime_policy_override, metadata, catalog_type, COALESCE(default_snapshot, '{}'::jsonb), COALESCE(customized_at::text, ''), pricing_version, status, created_at, updated_at, COALESCE(( SELECT jsonb_agg(DISTINCT compatibility_alias.alias ORDER BY compatibility_alias.alias) FROM model_compatibility_aliases compatibility_alias WHERE compatibility_alias.base_model_id = base_model_catalog.id AND compatibility_alias.active = true AND (compatibility_alias.expires_at IS NULL OR compatibility_alias.expires_at > now()) ), '[]'::jsonb), (SELECT count(*)::int FROM platform_models platform_model WHERE platform_model.base_model_id = base_model_catalog.id)` type BaseModelInput struct { ProviderKey string `json:"providerKey"` CanonicalModelKey string `json:"canonicalModelKey"` InvocationName string `json:"invocationName"` ProviderModelName string `json:"providerModelName"` ModelType StringList `json:"modelType"` ModelAlias string `json:"modelAlias"` DisplayName string `json:"displayName"` LegacyAliases StringList `json:"legacyAliases"` Capabilities map[string]any `json:"capabilities"` BaseBillingConfig map[string]any `json:"baseBillingConfig"` DefaultRateLimitPolicy map[string]any `json:"defaultRateLimitPolicy"` PricingRuleSetID string `json:"pricingRuleSetId"` RuntimePolicySetID string `json:"runtimePolicySetId"` RuntimePolicyOverride map[string]any `json:"runtimePolicyOverride"` Metadata map[string]any `json:"metadata"` CatalogType string `json:"catalogType"` DefaultSnapshot map[string]any `json:"defaultSnapshot"` PricingVersion int `json:"pricingVersion"` Status string `json:"status"` } type baseModelScanner interface { Scan(dest ...any) error } type StringList []string func (list *StringList) UnmarshalJSON(data []byte) error { var values []string if err := json.Unmarshal(data, &values); err == nil { *list = uniqueStringList(values) return nil } var value string if err := json.Unmarshal(data, &value); err != nil { return err } *list = uniqueStringList([]string{value}) return nil } func (s *Store) ListBaseModels(ctx context.Context) ([]BaseModel, error) { rows, err := s.pool.Query(ctx, ` SELECT `+baseModelColumns+` FROM base_model_catalog ORDER BY provider_key ASC, canonical_model_key ASC`) if err != nil { return nil, err } defer rows.Close() items := make([]BaseModel, 0) for rows.Next() { item, err := scanBaseModel(rows) if err != nil { return nil, err } items = append(items, item) } return items, rows.Err() } func (s *Store) CreateBaseModel(ctx context.Context, input BaseModelInput) (BaseModel, error) { input = normalizeBaseModelInput(input) capabilities, _ := json.Marshal(emptyObjectIfNil(input.Capabilities)) billingConfig, _ := json.Marshal(emptyObjectIfNil(input.BaseBillingConfig)) rateLimitPolicy, _ := json.Marshal(emptyObjectIfNil(input.DefaultRateLimitPolicy)) runtimePolicyOverride, _ := json.Marshal(emptyObjectIfNil(input.RuntimePolicyOverride)) metadata, _ := json.Marshal(emptyObjectIfNil(input.Metadata)) defaultSnapshot, _ := json.Marshal(emptyObjectIfNil(input.DefaultSnapshot)) modelType, _ := json.Marshal(input.ModelType) tx, err := s.pool.Begin(ctx) if err != nil { return BaseModel{}, err } defer tx.Rollback(ctx) item, err := scanBaseModel(tx.QueryRow(ctx, ` INSERT INTO base_model_catalog ( provider_id, provider_key, canonical_model_key, invocation_name, provider_model_name, model_type, display_name, capabilities, base_billing_config, default_rate_limit_policy, pricing_rule_set_id, runtime_policy_set_id, runtime_policy_override, metadata, catalog_type, default_snapshot, pricing_version, status ) VALUES ( (SELECT id FROM model_catalog_providers WHERE provider_key = $1 OR provider_code = $1 LIMIT 1), $1, $2, $3, $4, $5::jsonb, $6, $7, $8, $9, COALESCE(NULLIF($10, '')::uuid, (SELECT id FROM model_pricing_rule_sets WHERE rule_set_key = 'default-multimodal-v1' LIMIT 1)), COALESCE(NULLIF($11, '')::uuid, (SELECT id FROM model_runtime_policy_sets WHERE policy_key = 'default-runtime-v1' LIMIT 1)), $12, $13, NULLIF($14, ''), NULLIF($15::jsonb, '{}'::jsonb), $16, $17 ) RETURNING `+baseModelColumns, input.ProviderKey, input.CanonicalModelKey, input.InvocationName, input.ProviderModelName, string(modelType), input.DisplayName, capabilities, billingConfig, rateLimitPolicy, input.PricingRuleSetID, input.RuntimePolicySetID, runtimePolicyOverride, metadata, input.CatalogType, string(defaultSnapshot), input.PricingVersion, input.Status, )) if err != nil { return BaseModel{}, err } if err := replaceBaseModelAliases(ctx, tx, item.ID, input); err != nil { return BaseModel{}, err } if err := tx.Commit(ctx); err != nil { return BaseModel{}, err } item.LegacyAliases = input.LegacyAliases return item, nil } func (s *Store) UpdateBaseModel(ctx context.Context, id string, input BaseModelInput) (BaseModel, error) { input = normalizeBaseModelInput(input) capabilities, _ := json.Marshal(emptyObjectIfNil(input.Capabilities)) billingConfig, _ := json.Marshal(emptyObjectIfNil(input.BaseBillingConfig)) rateLimitPolicy, _ := json.Marshal(emptyObjectIfNil(input.DefaultRateLimitPolicy)) runtimePolicyOverride, _ := json.Marshal(emptyObjectIfNil(input.RuntimePolicyOverride)) metadata, _ := json.Marshal(emptyObjectIfNil(input.Metadata)) defaultSnapshot, _ := json.Marshal(emptyObjectIfNil(input.DefaultSnapshot)) modelType, _ := json.Marshal(input.ModelType) tx, err := s.pool.Begin(ctx) if err != nil { return BaseModel{}, err } defer tx.Rollback(ctx) item, err := scanBaseModel(tx.QueryRow(ctx, ` UPDATE base_model_catalog SET provider_id = (SELECT id FROM model_catalog_providers WHERE provider_key = $2 OR provider_code = $2 LIMIT 1), provider_key = $2, canonical_model_key = $3, invocation_name = $4, provider_model_name = $5, model_type = $6::jsonb, display_name = $7, capabilities = $8, base_billing_config = $9, default_rate_limit_policy = $10, pricing_rule_set_id = COALESCE(NULLIF($11, '')::uuid, (SELECT id FROM model_pricing_rule_sets WHERE rule_set_key = 'default-multimodal-v1' LIMIT 1)), runtime_policy_set_id = COALESCE(NULLIF($12, '')::uuid, (SELECT id FROM model_runtime_policy_sets WHERE policy_key = 'default-runtime-v1' LIMIT 1)), runtime_policy_override = $13, metadata = $14, catalog_type = NULLIF($15, ''), default_snapshot = COALESCE(NULLIF($16::jsonb, '{}'::jsonb), default_snapshot), customized_at = CASE WHEN NULLIF($15, '') = 'system' THEN now() ELSE NULL END, pricing_version = $17, status = $18, updated_at = now() WHERE id = $1::uuid RETURNING `+baseModelColumns, id, input.ProviderKey, input.CanonicalModelKey, input.InvocationName, input.ProviderModelName, string(modelType), input.DisplayName, capabilities, billingConfig, rateLimitPolicy, input.PricingRuleSetID, input.RuntimePolicySetID, runtimePolicyOverride, metadata, input.CatalogType, string(defaultSnapshot), input.PricingVersion, input.Status, )) if err != nil { return BaseModel{}, err } if err := replaceBaseModelAliases(ctx, tx, item.ID, input); err != nil { return BaseModel{}, err } if err := tx.Commit(ctx); err != nil { return BaseModel{}, err } item.LegacyAliases = input.LegacyAliases return item, nil } func (s *Store) ResetBaseModelToDefault(ctx context.Context, id string) (BaseModel, error) { var catalogType string var snapshotBytes []byte if err := s.pool.QueryRow(ctx, ` SELECT catalog_type, COALESCE(default_snapshot, '{}'::jsonb) FROM base_model_catalog WHERE id = $1::uuid`, id).Scan(&catalogType, &snapshotBytes); err != nil { return BaseModel{}, err } if catalogType != "system" { return BaseModel{}, ErrProtectedDefault } snapshot := decodeObject(snapshotBytes) if len(snapshot) == 0 { return BaseModel{}, ErrProtectedDefault } return scanBaseModel(s.pool.QueryRow(ctx, ` UPDATE base_model_catalog SET provider_id = (SELECT id FROM model_catalog_providers WHERE provider_key = COALESCE($2::text, provider_key) OR provider_code = COALESCE($2::text, provider_key) LIMIT 1), provider_key = COALESCE($2::text, provider_key), canonical_model_key = COALESCE($3::text, canonical_model_key), invocation_name = COALESCE($4::text, invocation_name), provider_model_name = COALESCE($5::text, provider_model_name), model_type = COALESCE($6::jsonb, model_type), display_name = COALESCE($7::text, display_name), capabilities = COALESCE($8::jsonb, capabilities), base_billing_config = COALESCE($9::jsonb, base_billing_config), default_rate_limit_policy = COALESCE($10::jsonb, default_rate_limit_policy), pricing_rule_set_id = COALESCE(NULLIF($11::text, '')::uuid, pricing_rule_set_id), runtime_policy_set_id = COALESCE(NULLIF($12::text, '')::uuid, runtime_policy_set_id), runtime_policy_override = COALESCE($13::jsonb, runtime_policy_override), metadata = COALESCE($14::jsonb, metadata), pricing_version = COALESCE($15::integer, pricing_version), status = COALESCE($16::text, status), customized_at = NULL, updated_at = now() WHERE id = $1::uuid RETURNING `+baseModelColumns, id, stringFromSnapshot(snapshot, "providerKey"), stringFromSnapshot(snapshot, "canonicalModelKey"), stringFromSnapshot(snapshot, "invocationName", "modelAlias", "providerModelName"), stringFromSnapshot(snapshot, "providerModelName"), jsonStringListFromSnapshot(snapshot, "modelType"), stringFromSnapshot(snapshot, "displayName", "modelAlias", "providerModelName"), jsonFromSnapshot(snapshot, "capabilities"), jsonFromSnapshot(snapshot, "baseBillingConfig"), jsonFromSnapshot(snapshot, "defaultRateLimitPolicy"), stringFromSnapshot(snapshot, "pricingRuleSetId"), stringFromSnapshot(snapshot, "runtimePolicySetId"), jsonFromSnapshot(snapshot, "runtimePolicyOverride"), jsonFromSnapshot(snapshot, "metadata"), intFromSnapshot(snapshot, "pricingVersion"), stringFromSnapshot(snapshot, "status"), )) } func (s *Store) ResetAllBaseModelsToDefault(ctx context.Context) ([]BaseModel, error) { rows, err := s.pool.Query(ctx, ` UPDATE base_model_catalog SET provider_id = ( SELECT id FROM model_catalog_providers WHERE provider_key = COALESCE(NULLIF(default_snapshot->>'providerKey', ''), provider_key) OR provider_code = COALESCE(NULLIF(default_snapshot->>'providerKey', ''), provider_key) LIMIT 1 ), provider_key = COALESCE(NULLIF(default_snapshot->>'providerKey', ''), provider_key), canonical_model_key = COALESCE(NULLIF(default_snapshot->>'canonicalModelKey', ''), canonical_model_key), invocation_name = COALESCE(NULLIF(default_snapshot->>'invocationName', ''), NULLIF(default_snapshot->>'modelAlias', ''), invocation_name), provider_model_name = COALESCE(NULLIF(default_snapshot->>'providerModelName', ''), provider_model_name), model_type = COALESCE(NULLIF(CASE WHEN jsonb_typeof(default_snapshot->'modelType') = 'array' THEN default_snapshot->'modelType' WHEN COALESCE(default_snapshot->>'modelType', '') <> '' THEN jsonb_build_array(default_snapshot->>'modelType') ELSE NULL END, '[]'::jsonb), model_type), display_name = COALESCE(NULLIF(COALESCE(default_snapshot->>'displayName', default_snapshot->>'modelAlias'), ''), display_name), capabilities = COALESCE(default_snapshot->'capabilities', capabilities), base_billing_config = COALESCE(default_snapshot->'baseBillingConfig', base_billing_config), default_rate_limit_policy = COALESCE(default_snapshot->'defaultRateLimitPolicy', default_rate_limit_policy), pricing_rule_set_id = COALESCE(NULLIF(default_snapshot->>'pricingRuleSetId', '')::uuid, pricing_rule_set_id), runtime_policy_set_id = COALESCE(NULLIF(default_snapshot->>'runtimePolicySetId', '')::uuid, runtime_policy_set_id), runtime_policy_override = COALESCE(default_snapshot->'runtimePolicyOverride', runtime_policy_override), metadata = COALESCE(default_snapshot->'metadata', metadata), pricing_version = COALESCE(NULLIF(default_snapshot->>'pricingVersion', '')::integer, pricing_version), status = COALESCE(NULLIF(default_snapshot->>'status', ''), status), customized_at = NULL, updated_at = now() WHERE catalog_type = 'system' AND COALESCE(default_snapshot, '{}'::jsonb) <> '{}'::jsonb RETURNING `+baseModelColumns) if err != nil { return nil, err } defer rows.Close() return scanBaseModelRows(rows) } func (s *Store) DeleteBaseModel(ctx context.Context, id string) error { result, err := s.pool.Exec(ctx, ` DELETE FROM base_model_catalog base_model WHERE base_model.id = $1::uuid AND NOT EXISTS ( SELECT 1 FROM platform_models platform_model WHERE platform_model.base_model_id = base_model.id )`, id) if err != nil { return err } if result.RowsAffected() == 0 { var exists bool if err := s.pool.QueryRow(ctx, `SELECT EXISTS (SELECT 1 FROM base_model_catalog WHERE id = $1::uuid)`, id).Scan(&exists); err != nil { return err } if exists { return ErrBaseModelInUse } return pgx.ErrNoRows } return nil } func scanBaseModelRows(rows pgx.Rows) ([]BaseModel, error) { items := make([]BaseModel, 0) for rows.Next() { item, err := scanBaseModel(rows) if err != nil { return nil, err } items = append(items, item) } return items, rows.Err() } func scanBaseModel(scanner baseModelScanner) (BaseModel, error) { var item BaseModel var modelType []byte var legacyAliases []byte var capabilities []byte var billingConfig []byte var rateLimitPolicy []byte var runtimePolicyOverride []byte var metadata []byte var defaultSnapshot []byte if err := scanner.Scan( &item.ID, &item.ProviderKey, &item.CanonicalModelKey, &item.InvocationName, &item.ProviderModelName, &modelType, &item.DisplayName, &capabilities, &billingConfig, &rateLimitPolicy, &item.PricingRuleSetID, &item.RuntimePolicySetID, &runtimePolicyOverride, &metadata, &item.CatalogType, &defaultSnapshot, &item.CustomizedAt, &item.PricingVersion, &item.Status, &item.CreatedAt, &item.UpdatedAt, &legacyAliases, &item.ReferenceCount, ); err != nil { return BaseModel{}, err } item.Capabilities = decodeObject(capabilities) item.BaseBillingConfig = decodeObject(billingConfig) item.DefaultRateLimitPolicy = decodeObject(rateLimitPolicy) item.RuntimePolicyOverride = decodeObject(runtimePolicyOverride) item.Metadata = decodeObject(metadata) item.DefaultSnapshot = decodeObject(defaultSnapshot) item.ModelType = decodeStringArray(modelType) item.LegacyAliases = decodeStringArray(legacyAliases) item.ModelAlias = item.InvocationName return item, nil } func normalizeBaseModelInput(input BaseModelInput) BaseModelInput { input.ProviderKey = strings.TrimSpace(input.ProviderKey) input.CanonicalModelKey = strings.TrimSpace(input.CanonicalModelKey) input.InvocationName = strings.TrimSpace(input.InvocationName) input.ProviderModelName = strings.TrimSpace(input.ProviderModelName) input.ModelType = uniqueStringList(input.ModelType) input.ModelAlias = strings.TrimSpace(input.ModelAlias) input.DisplayName = strings.TrimSpace(input.DisplayName) input.LegacyAliases = normalizeCompatibilityAliases(input.LegacyAliases) input.PricingRuleSetID = strings.TrimSpace(input.PricingRuleSetID) input.RuntimePolicySetID = strings.TrimSpace(input.RuntimePolicySetID) input.CatalogType = strings.TrimSpace(input.CatalogType) input.Status = strings.TrimSpace(input.Status) if input.CanonicalModelKey == "" && input.ProviderKey != "" && input.ProviderModelName != "" { input.CanonicalModelKey = input.ProviderKey + ":" + input.ProviderModelName } if input.InvocationName == "" { input.InvocationName = input.ModelAlias } if input.InvocationName == "" { input.InvocationName = input.ProviderModelName } if input.DisplayName == "" { input.DisplayName = input.InvocationName } input.ModelAlias = input.InvocationName input.LegacyAliases = withoutStrings(input.LegacyAliases, input.InvocationName) if len(input.ModelType) == 0 { input.ModelType = StringList{"text_generate"} } if input.CatalogType == "" { input.CatalogType = "custom" } if input.PricingVersion <= 0 { input.PricingVersion = 1 } if input.Status == "" { input.Status = "active" } return input } func replaceBaseModelAliases(ctx context.Context, tx pgx.Tx, baseModelID string, input BaseModelInput) error { for _, alias := range input.LegacyAliases { for _, modelType := range input.ModelType { var conflictingCanonicalKey string if err := tx.QueryRow(ctx, ` SELECT COALESCE(( SELECT other.canonical_model_key FROM base_model_catalog other WHERE other.id <> $1::uuid AND other.invocation_name <> $4::text AND other.model_type ? $3::text AND ( other.invocation_name = $2::text OR EXISTS ( SELECT 1 FROM model_compatibility_aliases other_alias WHERE other_alias.base_model_id = other.id AND other_alias.alias = $2::text AND other_alias.model_type = $3::text AND other_alias.active = true AND (other_alias.expires_at IS NULL OR other_alias.expires_at > now()) ) ) LIMIT 1 ), '')`, baseModelID, alias, modelType, input.InvocationName).Scan(&conflictingCanonicalKey); err != nil { return err } if conflictingCanonicalKey != "" { return fmt.Errorf("%w: %q (%s) conflicts with %s", ErrModelAliasConflict, alias, modelType, conflictingCanonicalKey) } } } if _, err := tx.Exec(ctx, `DELETE FROM model_compatibility_aliases WHERE base_model_id = $1::uuid`, baseModelID); err != nil { return err } for _, alias := range input.LegacyAliases { for _, modelType := range input.ModelType { if _, err := tx.Exec(ctx, ` INSERT INTO model_compatibility_aliases (base_model_id, alias, model_type, expires_at) VALUES ($1::uuid, $2, $3, now() + interval '14 days')`, baseModelID, alias, modelType); err != nil { return err } } } return nil } func normalizeCompatibilityAliases(values []string) StringList { out := make(StringList, 0, len(values)) seen := map[string]bool{} for _, value := range values { value = strings.TrimSpace(value) if value == "" || seen[value] { continue } seen[value] = true out = append(out, value) } return out } func withoutStrings(values []string, excluded string) StringList { out := make(StringList, 0, len(values)) for _, value := range values { if value != excluded { out = append(out, value) } } return out } func stringFromSnapshot(snapshot map[string]any, keys ...string) any { for _, key := range keys { value, ok := snapshot[key] if !ok { continue } switch typed := value.(type) { case string: if strings.TrimSpace(typed) != "" { return typed } case []any: for _, item := range typed { if text, ok := item.(string); ok && strings.TrimSpace(text) != "" { return strings.TrimSpace(text) } } case []string: if primary := primaryString(typed, ""); primary != "" { return primary } } } return nil } func intFromSnapshot(snapshot map[string]any, key string) any { switch value := snapshot[key].(type) { case float64: return int(value) case int: return value default: return nil } } func jsonFromSnapshot(snapshot map[string]any, key string) any { value, ok := snapshot[key] if !ok || value == nil { return nil } raw, err := json.Marshal(value) if err != nil { return nil } return string(raw) } func jsonStringListFromSnapshot(snapshot map[string]any, key string) any { values := stringListFromAny(snapshot[key]) if len(values) == 0 { if value, ok := snapshot[key].(string); ok { values = []string{value} } } normalized := uniqueStringList(values) if len(normalized) == 0 { return nil } raw, err := json.Marshal(normalized) if err != nil { return nil } return string(raw) } func stringListFromAny(value any) []string { switch typed := value.(type) { case []string: return typed case []any: values := make([]string, 0, len(typed)) for _, item := range typed { if text, ok := item.(string); ok { values = append(values, text) } } return values default: return nil } } func uniqueStringList(values []string) StringList { out := make([]string, 0, len(values)) seen := map[string]bool{} for _, value := range values { for _, normalized := range modelTypeAliases(value) { normalized = strings.TrimSpace(normalized) if normalized == "" || seen[normalized] { continue } seen[normalized] = true out = append(out, normalized) } } return out } func normalizeModelTypeList(values []string) StringList { return uniqueStringList(values) } func modelTypeAliases(value string) []string { switch strings.TrimSpace(value) { case "chat", "text", "responses": return []string{"text_generate"} case "image": return []string{"image_generate", "image_edit"} case "images.generations": return []string{"image_generate"} case "images.edits": return []string{"image_edit"} case "images.vectorize", "vectorize": return []string{"image_vectorize"} case "video", "videos.generations": return []string{"video_generate"} case "videos.upscales", "video_upscale": return []string{"video_enhance"} case "omni_video": return []string{"video_generate", "image_to_video", "omni_video"} case "song", "music", "song.generations", "music.generations", "music_generate": return []string{"audio_generate"} case "speech", "speech.generations", "tts": return []string{"text_to_speech"} default: return []string{value} } } func primaryString(values []string, fallback string) string { for _, value := range values { if value = strings.TrimSpace(value); value != "" { return value } } return fallback }