fix(billing): 修正模型计费配置继承优先级

This commit is contained in:
2026-07-22 00:24:36 +08:00
parent 56d4a3a6b7
commit 6b675c406e
8 changed files with 263 additions and 68 deletions
+34 -23
View File
@@ -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 {
@@ -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)
+13 -12
View File
@@ -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 {
+28
View File
@@ -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)
}
@@ -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},
}
}
+8 -4
View File
@@ -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,
@@ -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")
}
+39 -29
View File
@@ -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)