From 6b675c406e2ef77573a3f07e720a0927cbf0a504 Mon Sep 17 00:00:00 2001 From: wangbo Date: Tue, 21 Jul 2026 23:52:53 +0800 Subject: [PATCH] =?UTF-8?q?fix(billing):=20=E4=BF=AE=E6=AD=A3=E6=A8=A1?= =?UTF-8?q?=E5=9E=8B=E8=AE=A1=E8=B4=B9=E9=85=8D=E7=BD=AE=E7=BB=A7=E6=89=BF?= =?UTF-8?q?=E4=BC=98=E5=85=88=E7=BA=A7?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- apps/api/internal/httpapi/model_response.go | 57 ++++++++------- .../internal/httpapi/model_response_test.go | 35 ++++++++++ apps/api/internal/runner/pricing.go | 25 +++---- apps/api/internal/store/billing_config.go | 28 ++++++++ .../api/internal/store/billing_config_test.go | 69 +++++++++++++++++++ apps/api/internal/store/platform_models.go | 12 ++-- .../store/platform_models_integration_test.go | 37 ++++++++++ apps/api/internal/store/postgres.go | 68 ++++++++++-------- 8 files changed, 263 insertions(+), 68 deletions(-) create mode 100644 apps/api/internal/store/billing_config.go create mode 100644 apps/api/internal/store/billing_config_test.go create mode 100644 apps/api/internal/store/platform_models_integration_test.go diff --git a/apps/api/internal/httpapi/model_response.go b/apps/api/internal/httpapi/model_response.go index ded775b..0aa90dc 100644 --- a/apps/api/internal/httpapi/model_response.go +++ b/apps/api/internal/httpapi/model_response.go @@ -7,46 +7,57 @@ import ( ) func (s *Server) platformModelResponse(ctx context.Context, model store.PlatformModel) store.PlatformModel { + return s.platformModelResponseWithRuleSets(model, s.responsePricingRuleSetConfigs(ctx, []store.PlatformModel{model})) +} + +func (s *Server) platformModelResponseWithRuleSets(model store.PlatformModel, ruleSetConfigs map[string]map[string]any) store.PlatformModel { model.Capabilities = store.EffectivePlatformModelCapabilities(model.BaseCapabilities, model.Capabilities) model.Capabilities = enrichResponseCapabilities(model) - model = s.withEffectiveResponseBillingConfig(ctx, model) + model = withEffectiveResponseBillingConfig(model, ruleSetConfigs) return store.FilterPlatformModelBillingConfig(model) } func (s *Server) platformModelResponses(ctx context.Context, models []store.PlatformModel) []store.PlatformModel { + ruleSetConfigs := s.responsePricingRuleSetConfigs(ctx, models) items := make([]store.PlatformModel, len(models)) for i, model := range models { - items[i] = s.platformModelResponse(ctx, model) + items[i] = s.platformModelResponseWithRuleSets(model, ruleSetConfigs) } return items } -func (s *Server) withEffectiveResponseBillingConfig(ctx context.Context, model store.PlatformModel) store.PlatformModel { - config := model.BillingConfig - if model.PricingRuleSetID != "" { - if ruleSetConfig, err := s.store.PricingRuleSetBillingConfig(ctx, model.PricingRuleSetID); err == nil && len(ruleSetConfig) > 0 { - config = ruleSetConfig +func (s *Server) responsePricingRuleSetConfigs(ctx context.Context, models []store.PlatformModel) map[string]map[string]any { + configs := map[string]map[string]any{} + if s.store == nil { + return configs + } + ids := map[string]bool{} + for _, model := range models { + for _, id := range []string{firstNonEmpty(model.BasePricingRuleSetID, model.PlatformPricingRuleSetID), model.PricingRuleSetID} { + if id != "" { + ids[id] = true + } } } - if len(model.BillingConfigOverride) > 0 { - config = mergeResponseBillingConfig(config, model.BillingConfigOverride) + for id := range ids { + if config, err := s.store.PricingRuleSetBillingConfig(ctx, id); err == nil && len(config) > 0 { + configs[id] = config + } } - model.BillingConfig = config - return model + return configs } -func mergeResponseBillingConfig(base map[string]any, override map[string]any) map[string]any { - if len(base) == 0 && len(override) == 0 { - return nil - } - out := make(map[string]any, len(base)+len(override)) - for key, value := range base { - out[key] = value - } - for key, value := range override { - out[key] = value - } - return out +func withEffectiveResponseBillingConfig(model store.PlatformModel, ruleSetConfigs map[string]map[string]any) store.PlatformModel { + inheritedRuleSetConfig := ruleSetConfigs[firstNonEmpty(model.BasePricingRuleSetID, model.PlatformPricingRuleSetID)] + modelRuleSetConfig := ruleSetConfigs[model.PricingRuleSetID] + model.BillingConfig = store.ResolveEffectiveBillingConfig(store.EffectiveBillingConfigInput{ + BaseConfig: model.BaseBillingConfig, + LegacyPlatformModelConfig: model.BillingConfig, + InheritedRuleSetConfig: inheritedRuleSetConfig, + ModelRuleSetConfig: modelRuleSetConfig, + Override: model.BillingConfigOverride, + }) + return model } func enrichResponseCapabilities(model store.PlatformModel) map[string]any { diff --git a/apps/api/internal/httpapi/model_response_test.go b/apps/api/internal/httpapi/model_response_test.go index e3fd661..e4c5fae 100644 --- a/apps/api/internal/httpapi/model_response_test.go +++ b/apps/api/internal/httpapi/model_response_test.go @@ -173,6 +173,41 @@ func TestPlatformModelResponsePreservesTextGenerateFieldsOverFallbacks(t *testin assertStringListValue(t, textGenerate["thinkingEffortLevels"], []string{"minimal", "low", "medium"}) } +func TestPlatformModelResponseUsesBaseBillingConfigWithoutMaterializedSnapshot(t *testing.T) { + model := store.PlatformModel{ + ModelName: "base-priced-model", + ModelType: store.StringList{"video_generate"}, + BaseBillingConfig: map[string]any{ + "video": map[string]any{"basePrice": float64(416)}, + }, + } + + response := (&Server{}).platformModelResponse(context.Background(), model) + video, ok := response.BillingConfig["video"].(map[string]any) + if !ok || video["basePrice"] != float64(416) { + t.Fatalf("expected base billing price 416, got %#v", response.BillingConfig) + } +} + +func TestEffectiveResponseBillingConfigPrefersBaseRuleOverLegacySnapshot(t *testing.T) { + model := store.PlatformModel{ + BasePricingRuleSetID: "seedance-pricing", + BillingConfig: map[string]any{ + "video": map[string]any{"basePrice": float64(100)}, + }, + } + response := withEffectiveResponseBillingConfig(model, map[string]map[string]any{ + "seedance-pricing": { + "video": map[string]any{"basePrice": float64(416)}, + }, + }) + + video, ok := response.BillingConfig["video"].(map[string]any) + if !ok || video["basePrice"] != float64(416) { + t.Fatalf("expected base rule price 416, got %#v", response.BillingConfig) + } +} + func textGenerateCapabilities(t *testing.T, model store.PlatformModel) map[string]any { t.Helper() capabilities, ok := model.Capabilities["text_generate"].(map[string]any) diff --git a/apps/api/internal/runner/pricing.go b/apps/api/internal/runner/pricing.go index f149c21..7026ac9 100644 --- a/apps/api/internal/runner/pricing.go +++ b/apps/api/internal/runner/pricing.go @@ -194,24 +194,25 @@ func (s *Service) billings(ctx context.Context, user *auth.User, kind string, bo } func (s *Service) effectiveBillingConfig(ctx context.Context, candidate store.RuntimeModelCandidate) map[string]any { - base := candidate.BaseBillingConfig - if ruleSetID := firstNonEmptyString(candidate.BasePricingRuleSetID, candidate.PlatformPricingRuleSetID); ruleSetID != "" { + var inheritedRuleSetConfig map[string]any + if ruleSetID := firstNonEmptyString(candidate.BasePricingRuleSetID, candidate.PlatformPricingRuleSetID); ruleSetID != "" && s.store != nil { if ruleSetConfig, err := s.store.PricingRuleSetBillingConfig(ctx, ruleSetID); err == nil && len(ruleSetConfig) > 0 { - base = ruleSetConfig + inheritedRuleSetConfig = ruleSetConfig } } - if len(candidate.BillingConfig) > 0 { - base = candidate.BillingConfig - } - if candidate.ModelPricingRuleSetID != "" { + var modelRuleSetConfig map[string]any + if candidate.ModelPricingRuleSetID != "" && s.store != nil { if ruleSetConfig, err := s.store.PricingRuleSetBillingConfig(ctx, candidate.ModelPricingRuleSetID); err == nil && len(ruleSetConfig) > 0 { - base = ruleSetConfig + modelRuleSetConfig = ruleSetConfig } } - if len(candidate.BillingConfigOverride) > 0 { - base = mergeMap(base, candidate.BillingConfigOverride) - } - return base + return store.ResolveEffectiveBillingConfig(store.EffectiveBillingConfigInput{ + BaseConfig: candidate.BaseBillingConfig, + LegacyPlatformModelConfig: candidate.BillingConfig, + InheritedRuleSetConfig: inheritedRuleSetConfig, + ModelRuleSetConfig: modelRuleSetConfig, + Override: candidate.BillingConfigOverride, + }) } func effectiveDiscount(ctx context.Context, db *store.Store, user *auth.User, candidate store.RuntimeModelCandidate) float64 { diff --git a/apps/api/internal/store/billing_config.go b/apps/api/internal/store/billing_config.go new file mode 100644 index 0000000..9cb1114 --- /dev/null +++ b/apps/api/internal/store/billing_config.go @@ -0,0 +1,28 @@ +package store + +// EffectiveBillingConfigInput describes the billing layers used by runtime and +// catalog responses. LegacyPlatformModelConfig is retained only as a fallback +// for models that do not have an effective pricing rule set. +type EffectiveBillingConfigInput struct { + BaseConfig map[string]any + LegacyPlatformModelConfig map[string]any + InheritedRuleSetConfig map[string]any + ModelRuleSetConfig map[string]any + Override map[string]any +} + +// ResolveEffectiveBillingConfig keeps inherited pricing rules authoritative over +// the legacy materialized snapshot. Explicit model rules and overrides retain +// their higher-priority exception semantics. +func ResolveEffectiveBillingConfig(input EffectiveBillingConfigInput) map[string]any { + config := mergeObjects(input.BaseConfig, nil) + if len(input.InheritedRuleSetConfig) > 0 { + config = mergeObjects(input.InheritedRuleSetConfig, nil) + } else if len(input.LegacyPlatformModelConfig) > 0 { + config = mergeObjects(input.LegacyPlatformModelConfig, nil) + } + if len(input.ModelRuleSetConfig) > 0 { + config = mergeObjects(input.ModelRuleSetConfig, nil) + } + return mergeObjects(config, input.Override) +} diff --git a/apps/api/internal/store/billing_config_test.go b/apps/api/internal/store/billing_config_test.go new file mode 100644 index 0000000..67be617 --- /dev/null +++ b/apps/api/internal/store/billing_config_test.go @@ -0,0 +1,69 @@ +package store + +import "testing" + +func TestResolveEffectiveBillingConfigKeepsPricingRulesAuthoritative(t *testing.T) { + tests := []struct { + name string + input EffectiveBillingConfigInput + want float64 + }{ + { + name: "inherited rule replaces stale platform snapshot", + input: EffectiveBillingConfigInput{ + BaseConfig: videoBillingConfig(100), + LegacyPlatformModelConfig: videoBillingConfig(100), + InheritedRuleSetConfig: videoBillingConfig(416), + }, + want: 416, + }, + { + name: "legacy snapshot remains a fallback without a rule", + input: EffectiveBillingConfigInput{ + BaseConfig: videoBillingConfig(100), + LegacyPlatformModelConfig: videoBillingConfig(125), + }, + want: 125, + }, + { + name: "model rule remains an explicit pricing exception", + input: EffectiveBillingConfigInput{ + BaseConfig: videoBillingConfig(100), + LegacyPlatformModelConfig: videoBillingConfig(125), + InheritedRuleSetConfig: videoBillingConfig(416), + ModelRuleSetConfig: videoBillingConfig(500), + }, + want: 500, + }, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + config := ResolveEffectiveBillingConfig(test.input) + video, ok := config["video"].(map[string]any) + if !ok { + t.Fatalf("expected video billing config, got %#v", config) + } + if got := video["basePrice"]; got != test.want { + t.Fatalf("video base price = %#v, want %v", got, test.want) + } + }) + } +} + +func TestResolveEffectiveBillingConfigAppliesOverrideLast(t *testing.T) { + config := ResolveEffectiveBillingConfig(EffectiveBillingConfigInput{ + InheritedRuleSetConfig: videoBillingConfig(416), + Override: videoBillingConfig(600), + }) + video, ok := config["video"].(map[string]any) + if !ok || video["basePrice"] != float64(600) { + t.Fatalf("expected override price 600, got %#v", config) + } +} + +func videoBillingConfig(basePrice float64) map[string]any { + return map[string]any{ + "video": map[string]any{"basePrice": basePrice}, + } +} diff --git a/apps/api/internal/store/platform_models.go b/apps/api/internal/store/platform_models.go index 6c5f7d4..9eb70ca 100644 --- a/apps/api/internal/store/platform_models.go +++ b/apps/api/internal/store/platform_models.go @@ -23,6 +23,7 @@ type modelCatalogSnapshot struct { DisplayName string Capabilities map[string]any BaseBillingConfig map[string]any + PricingRuleSetID string DefaultRateLimitPolicy map[string]any RuntimePolicySetID string RuntimePolicyOverride map[string]any @@ -121,10 +122,10 @@ func (s *Store) createPlatformModel(ctx context.Context, q platformModelQuerier, if err := validateEnabledVolcesTextModelCapabilities(ctx, q, input, capabilities); err != nil { return PlatformModel{}, err } + // billing_config is a legacy, explicitly supplied compatibility field. Do + // not materialize base-model pricing into it: copied prices become stale as + // soon as the base pricing rule changes and can mask the authoritative rule. billingConfig := input.BillingConfig - if len(billingConfig) == 0 { - billingConfig = mergeObjects(base.BaseBillingConfig, input.BillingConfigOverride) - } explicitRuntimePolicySetID := strings.TrimSpace(input.RuntimePolicySetID) rateLimitPolicy := input.RateLimitPolicy if len(rateLimitPolicy) == 0 && explicitRuntimePolicySetID == "" { @@ -260,6 +261,8 @@ RETURNING id::text, platform_id::text, COALESCE(base_model_id::text, ''), model_ model.ModelType = decodeStringArray(modelTypeBytes) model.BillingConfigOverride = decodeObject(billingOverrideBytes) model.BillingConfig = decodeObject(billingBytes) + model.BaseBillingConfig = base.BaseBillingConfig + model.BasePricingRuleSetID = base.PricingRuleSetID model.PermissionConfig = decodeObject(permissionBytes) model.RetryPolicy = decodeObject(retryPolicyBytes) model.RateLimitPolicy = decodeObject(rateLimitPolicyBytes) @@ -368,7 +371,7 @@ func (s *Store) lookupBaseModel(ctx context.Context, q platformModelQuerier, id var modelTypeBytes []byte err := q.QueryRow(ctx, ` SELECT id::text, provider_key, canonical_model_key, provider_model_name, model_type, display_name, - capabilities, base_billing_config, default_rate_limit_policy, + capabilities, base_billing_config, COALESCE(pricing_rule_set_id::text, ''), default_rate_limit_policy, COALESCE(runtime_policy_set_id::text, ''), runtime_policy_override FROM base_model_catalog WHERE ($1 <> '' AND id = NULLIF($1, '')::uuid) @@ -384,6 +387,7 @@ LIMIT 1`, strings.TrimSpace(id), strings.TrimSpace(canonicalKey), strings.TrimSp &item.DisplayName, &capabilities, &billingConfig, + &item.PricingRuleSetID, &rateLimitPolicy, &item.RuntimePolicySetID, &runtimePolicyOverride, diff --git a/apps/api/internal/store/platform_models_integration_test.go b/apps/api/internal/store/platform_models_integration_test.go new file mode 100644 index 0000000..a59f3d9 --- /dev/null +++ b/apps/api/internal/store/platform_models_integration_test.go @@ -0,0 +1,37 @@ +package store + +import ( + "context" + "os" + "strings" + "testing" +) + +func TestListModelsLoadsEffectiveBillingSources(t *testing.T) { + databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL")) + if databaseURL == "" { + t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run the platform-model billing source integration test") + } + + ctx := context.Background() + db, err := Connect(ctx, databaseURL) + if err != nil { + t.Fatalf("connect store: %v", err) + } + defer db.Close() + + models, err := db.ListModels(ctx) + if err != nil { + t.Fatalf("list models with effective billing sources: %v", err) + } + for _, model := range models { + if model.BaseModelID == "" { + continue + } + if model.BaseBillingConfig == nil { + t.Fatalf("platform model %s did not load base billing config", model.ID) + } + return + } + t.Skip("database has no base-model-backed platform model") +} diff --git a/apps/api/internal/store/postgres.go b/apps/api/internal/store/postgres.go index 036793f..4840668 100644 --- a/apps/api/internal/store/postgres.go +++ b/apps/api/internal/store/postgres.go @@ -217,33 +217,36 @@ type CreatedAPIKey struct { } type PlatformModel struct { - ID string `json:"id"` - PlatformID string `json:"platformId"` - BaseModelID string `json:"baseModelId,omitempty"` - Provider string `json:"provider,omitempty"` - PlatformName string `json:"platformName,omitempty"` - ModelName string `json:"modelName"` - ProviderModelName string `json:"providerModelName,omitempty"` - ModelAlias string `json:"modelAlias,omitempty"` - ModelType StringList `json:"modelType"` - DisplayName string `json:"displayName"` - CapabilityOverride map[string]any `json:"capabilityOverride,omitempty"` - Capabilities map[string]any `json:"capabilities,omitempty"` - BaseCapabilities map[string]any `json:"-"` - PricingMode string `json:"pricingMode"` - DiscountFactor float64 `json:"discountFactor,omitempty"` - PricingRuleSetID string `json:"pricingRuleSetId,omitempty"` - BillingConfigOverride map[string]any `json:"billingConfigOverride,omitempty"` - BillingConfig map[string]any `json:"billingConfig,omitempty"` - PermissionConfig map[string]any `json:"permissionConfig,omitempty"` - RetryPolicy map[string]any `json:"retryPolicy,omitempty"` - RateLimitPolicy map[string]any `json:"rateLimitPolicy,omitempty"` - RuntimePolicySetID string `json:"runtimePolicySetId,omitempty"` - RuntimePolicyOverride map[string]any `json:"runtimePolicyOverride,omitempty"` - CooldownUntil string `json:"cooldownUntil,omitempty"` - Enabled bool `json:"enabled"` - CreatedAt time.Time `json:"createdAt"` - UpdatedAt time.Time `json:"updatedAt"` + ID string `json:"id"` + PlatformID string `json:"platformId"` + BaseModelID string `json:"baseModelId,omitempty"` + Provider string `json:"provider,omitempty"` + PlatformName string `json:"platformName,omitempty"` + ModelName string `json:"modelName"` + ProviderModelName string `json:"providerModelName,omitempty"` + ModelAlias string `json:"modelAlias,omitempty"` + ModelType StringList `json:"modelType"` + DisplayName string `json:"displayName"` + CapabilityOverride map[string]any `json:"capabilityOverride,omitempty"` + Capabilities map[string]any `json:"capabilities,omitempty"` + BaseCapabilities map[string]any `json:"-"` + BaseBillingConfig map[string]any `json:"-"` + BasePricingRuleSetID string `json:"-"` + PlatformPricingRuleSetID string `json:"-"` + PricingMode string `json:"pricingMode"` + DiscountFactor float64 `json:"discountFactor,omitempty"` + PricingRuleSetID string `json:"pricingRuleSetId,omitempty"` + BillingConfigOverride map[string]any `json:"billingConfigOverride,omitempty"` + BillingConfig map[string]any `json:"billingConfig,omitempty"` + PermissionConfig map[string]any `json:"permissionConfig,omitempty"` + RetryPolicy map[string]any `json:"retryPolicy,omitempty"` + RateLimitPolicy map[string]any `json:"rateLimitPolicy,omitempty"` + RuntimePolicySetID string `json:"runtimePolicySetId,omitempty"` + RuntimePolicyOverride map[string]any `json:"runtimePolicyOverride,omitempty"` + CooldownUntil string `json:"cooldownUntil,omitempty"` + Enabled bool `json:"enabled"` + CreatedAt time.Time `json:"createdAt"` + UpdatedAt time.Time `json:"updatedAt"` } type AccessRule struct { @@ -927,7 +930,9 @@ func (s *Store) listModels(ctx context.Context, platformID string) ([]PlatformMo rows, err := s.pool.Query(ctx, ` SELECT m.id::text, m.platform_id::text, COALESCE(m.base_model_id::text, ''), p.provider, p.name, m.model_name, COALESCE(NULLIF(m.provider_model_name, ''), m.model_name), COALESCE(m.model_alias, ''), m.model_type, m.display_name, - m.capability_override, m.capabilities, COALESCE(b.capabilities, '{}'::jsonb), m.pricing_mode, COALESCE(m.discount_factor, 0)::float8, + m.capability_override, m.capabilities, COALESCE(b.capabilities, '{}'::jsonb), + COALESCE(b.base_billing_config, '{}'::jsonb), COALESCE(b.pricing_rule_set_id::text, ''), + COALESCE(p.pricing_rule_set_id::text, ''), m.pricing_mode, COALESCE(m.discount_factor, 0)::float8, COALESCE(m.pricing_rule_set_id::text, ''), m.billing_config_override, m.billing_config, m.permission_config, m.retry_policy, m.rate_limit_policy, COALESCE(m.runtime_policy_set_id::text, ''), m.runtime_policy_override, COALESCE(to_char(m.cooldown_until AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.MS"Z"'), ''), @@ -935,7 +940,7 @@ SELECT m.id::text, m.platform_id::text, COALESCE(m.base_model_id::text, ''), p.p FROM platform_models m JOIN integration_platforms p ON p.id = m.platform_id LEFT JOIN LATERAL ( - SELECT catalog.capabilities + SELECT catalog.capabilities, catalog.base_billing_config, catalog.pricing_rule_set_id FROM base_model_catalog catalog WHERE (m.base_model_id IS NOT NULL AND catalog.id = m.base_model_id) OR ( @@ -962,6 +967,7 @@ ORDER BY m.model_type ASC, m.model_name ASC`, args...) var capabilityOverride []byte var capabilities []byte var baseCapabilities []byte + var baseBillingConfig []byte var billingConfigOverride []byte var billingConfig []byte var permissionConfig []byte @@ -983,6 +989,9 @@ ORDER BY m.model_type ASC, m.model_name ASC`, args...) &capabilityOverride, &capabilities, &baseCapabilities, + &baseBillingConfig, + &model.BasePricingRuleSetID, + &model.PlatformPricingRuleSetID, &model.PricingMode, &model.DiscountFactor, &model.PricingRuleSetID, @@ -1003,6 +1012,7 @@ ORDER BY m.model_type ASC, m.model_name ASC`, args...) model.CapabilityOverride = decodeObject(capabilityOverride) model.Capabilities = decodeObject(capabilities) model.BaseCapabilities = decodeObject(baseCapabilities) + model.BaseBillingConfig = decodeObject(baseBillingConfig) model.ModelType = decodeStringArray(modelTypeBytes) model.BillingConfigOverride = decodeObject(billingConfigOverride) model.BillingConfig = decodeObject(billingConfig)