diff --git a/backend/internal/service/account_stats_pricing.go b/backend/internal/service/account_stats_pricing.go index df9e8e05aa..8d5bc144fc 100644 --- a/backend/internal/service/account_stats_pricing.go +++ b/backend/internal/service/account_stats_pricing.go @@ -68,7 +68,7 @@ func tryModelFilePricing(billingService *BillingService, model string, tokens Us return nil } normalizedTier := normalizeBillingServiceTier(serviceTier) - if normalizedTier == "priority" || normalizedTier == "flex" || + if normalizedTier == "priority" || normalizedTier == "fast" || normalizedTier == "flex" || billingService.shouldApplySessionLongContextPricing(tokens, pricing) { breakdown, err := billingService.CalculateCostWithServiceTier(model, tokens, 1, normalizedTier) if err != nil || breakdown == nil || breakdown.TotalCost <= 0 { diff --git a/backend/internal/service/account_stats_pricing_test.go b/backend/internal/service/account_stats_pricing_test.go index 48336a5834..1bd28896fc 100644 --- a/backend/internal/service/account_stats_pricing_test.go +++ b/backend/internal/service/account_stats_pricing_test.go @@ -768,6 +768,26 @@ func TestResolveAccountStatsCost_FallsBackToLiteLLM(t *testing.T) { require.InDelta(t, 0.2, *result, 1e-12) } +func TestResolveAccountStatsCost_FallbackHonorsAnthropicFast(t *testing.T) { + channel := &Channel{ID: 1, Status: StatusActive} + cs := newTestChannelServiceForStats(t, channel, 10, "anthropic") + bs := newTestBillingServiceWithPrices(map[string]*ModelPricing{ + "claude-opus-5": { + InputPricePerToken: 5e-6, + OutputPricePerToken: 25e-6, + }, + }) + + result := resolveAccountStatsCost( + context.Background(), cs, bs, + 1, 10, "claude-opus-5", + UsageTokens{InputTokens: 1_000_000, OutputTokens: 1_000_000}, + 1, 0, "fast", + ) + require.NotNil(t, result) + require.InDelta(t, 60, *result, 1e-12) +} + func TestResolveAccountStatsCost_Gemini36FlashTierUsesFallbackPricing(t *testing.T) { channel := &Channel{ ID: 1, diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index 6ead8da021..af8c694495 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -91,25 +91,27 @@ type BillingCache interface { // ModelPricing 模型价格配置(per-token价格,与LiteLLM格式一致) type ModelPricing struct { - InputPricePerToken float64 // 每token输入价格 (USD) - InputPricePerTokenPriority float64 // priority service tier 下每token输入价格 (USD) - ImageInputPricePerToken float64 // 图片输入 token 价格 (USD),用于多模态 embedding 等图文不同价场景;为 0 时回退到 InputPricePerToken - OutputPricePerToken float64 // 每token输出价格 (USD) - OutputPricePerTokenPriority float64 // priority service tier 下每token输出价格 (USD) - CacheCreationPricePerToken float64 // 缓存创建每token价格 (USD) - CacheCreationPricePerTokenPriority float64 // priority service tier 下缓存创建每token价格 (USD) - CacheCreationPriceExplicit bool // 是否由渠道/区间定价显式设定(为 true 时即使 == 0 也不回退) - CacheReadPricePerToken float64 // 缓存读取每token价格 (USD) - CacheReadPricePerTokenPriority float64 // priority service tier 下缓存读取每token价格 (USD) - CacheCreation5mPrice float64 // 5分钟缓存创建每token价格 (USD) - CacheCreation1hPrice float64 // 1小时缓存创建每token价格 (USD) - SupportsCacheBreakdown bool // 是否支持详细的缓存分类 - LongContextInputThreshold int // 超过阈值后按整次会话提升输入价格 - LongContextThresholdInclusive bool // 达到阈值即应用(xAI);默认保持严格大于以兼容既有模型 - LongContextInputMultiplier float64 // 长上下文整次会话输入倍率 - LongContextOutputMultiplier float64 // 长上下文整次会话输出倍率 - ImageOutputPricePerToken float64 // 图片输出 token 价格 (USD) - ImageOutputPriceExplicit bool // 是否由渠道定价显式设定(为 true 时即使 == 0 也不回退) + InputPricePerToken float64 // 每token输入价格 (USD) + InputPricePerTokenPriority float64 // priority service tier 下每token输入价格 (USD) + ImageInputPricePerToken float64 // 图片输入 token 价格 (USD),用于多模态 embedding 等图文不同价场景;为 0 时回退到 InputPricePerToken + OutputPricePerToken float64 // 每token输出价格 (USD) + OutputPricePerTokenPriority float64 // priority service tier 下每token输出价格 (USD) + CacheCreationPricePerToken float64 // 缓存创建每token价格 (USD) + CacheCreationPricePerTokenPriority float64 // priority service tier 下缓存创建每token价格 (USD) + CacheCreationPriceExplicit bool // 是否由渠道/区间定价显式设定(为 true 时即使 == 0 也不回退) + CacheReadPricePerToken float64 // 缓存读取每token价格 (USD) + CacheReadPricePerTokenPriority float64 // priority service tier 下缓存读取每token价格 (USD) + FastMultiplier *float64 // 渠道显式 Fast/priority 倍率;nil 时沿用模型目录行为 + FlexMultiplier *float64 // 渠道显式 Flex 倍率;nil 时沿用默认行为 + CacheCreation5mPrice float64 // 5分钟缓存创建每token价格 (USD) + CacheCreation1hPrice float64 // 1小时缓存创建每token价格 (USD) + SupportsCacheBreakdown bool // 是否支持详细的缓存分类 + LongContextInputThreshold int // 超过阈值后按整次会话提升输入价格 + LongContextThresholdInclusive bool // 达到阈值即应用(xAI);默认保持严格大于以兼容既有模型 + LongContextInputMultiplier float64 // 长上下文整次会话输入倍率 + LongContextOutputMultiplier float64 // 长上下文整次会话输出倍率 + ImageOutputPricePerToken float64 // 图片输出 token 价格 (USD) + ImageOutputPriceExplicit bool // 是否由渠道定价显式设定(为 true 时即使 == 0 也不回退) } const ( @@ -123,7 +125,14 @@ func normalizeBillingServiceTier(serviceTier string) string { } func usePriorityServiceTierPricing(serviceTier string, pricing *ModelPricing) bool { - if pricing == nil || normalizeBillingServiceTier(serviceTier) != "priority" { + if pricing == nil { + return false + } + tier := normalizeBillingServiceTier(serviceTier) + if tier != "priority" && tier != "fast" { + return false + } + if pricing.FastMultiplier != nil { return false } return pricing.InputPricePerTokenPriority > 0 || pricing.OutputPricePerTokenPriority > 0 || @@ -132,7 +141,7 @@ func usePriorityServiceTierPricing(serviceTier string, pricing *ModelPricing) bo func serviceTierCostMultiplier(serviceTier string) float64 { switch normalizeBillingServiceTier(serviceTier) { - case "priority": + case "priority", "fast": return 2.0 case "flex": return 0.5 @@ -141,6 +150,34 @@ func serviceTierCostMultiplier(serviceTier string) float64 { } } +func configuredServiceTierMultiplier(serviceTier string, pricing *ModelPricing) float64 { + if pricing != nil { + switch normalizeBillingServiceTier(serviceTier) { + case "priority", "fast": + if pricing.FastMultiplier != nil { + return *pricing.FastMultiplier + } + case "flex": + if pricing.FlexMultiplier != nil { + return *pricing.FlexMultiplier + } + } + } + return serviceTierCostMultiplier(serviceTier) +} + +func pricingWithPriorityMultiplier(base *ModelPricing, multiplier float64) *ModelPricing { + if base == nil { + return nil + } + cloned := *base + cloned.InputPricePerTokenPriority = cloned.InputPricePerToken * multiplier + cloned.OutputPricePerTokenPriority = cloned.OutputPricePerToken * multiplier + cloned.CacheCreationPricePerTokenPriority = cloned.CacheCreationPricePerToken * multiplier + cloned.CacheReadPricePerTokenPriority = cloned.CacheReadPricePerToken * multiplier + return &cloned +} + // UsageTokens 使用的token数量 type UsageTokens struct { InputTokens int @@ -281,10 +318,10 @@ func (s *BillingService) initFallbackPricing() { // Claude 4.7 Opus (暂与4.6同价,待官方定价更新) s.fallbackPrices["claude-opus-4.7"] = s.fallbackPrices["claude-opus-4.6"] - // Claude 4.8 Opus / Claude Opus 5(官方同价:$5 输入 / $25 输出 per MTok)。 + // Claude 4.8 Opus / Claude Opus 5(标准 $5/$25,Fast $10/$50 per MTok)。 // 缺少这两条时 getFallbackPricing 会掉到 claude-3-opus($15/$75),造成 3 倍超收。 - s.fallbackPrices["claude-opus-4.8"] = s.fallbackPrices["claude-opus-4.7"] - s.fallbackPrices["claude-opus-5"] = s.fallbackPrices["claude-opus-4.8"] + s.fallbackPrices["claude-opus-4.8"] = pricingWithPriorityMultiplier(s.fallbackPrices["claude-opus-4.7"], 2) + s.fallbackPrices["claude-opus-5"] = pricingWithPriorityMultiplier(s.fallbackPrices["claude-opus-4.8"], 2) // Gemini 3.1 Pro s.fallbackPrices["gemini-3.1-pro"] = &ModelPricing{ @@ -320,9 +357,31 @@ func (s *BillingService) initFallbackPricing() { LongContextInputMultiplier: openAIGPT54LongContextInputMultiplier, LongContextOutputMultiplier: openAIGPT54LongContextOutputMultiplier, } - // GPT-5.5 / GPT-5.5 Pro 暂无独立定价,回退到 GPT-5.4。 - s.fallbackPrices["gpt-5.5"] = s.fallbackPrices["gpt-5.4"] - s.fallbackPrices["gpt-5.5-pro"] = s.fallbackPrices["gpt-5.4"] + // OpenAI GPT-5.5 官方价格;Fast 为标准价 2.5 倍。 + // Source: https://platform.openai.com/docs/pricing + s.fallbackPrices["gpt-5.5"] = pricingWithPriorityMultiplier(&ModelPricing{ + InputPricePerToken: 5e-6, + OutputPricePerToken: 30e-6, + // 官方未列独立 cache-write 价;内部出现 cache creation token 时按输入价兜底。 + CacheCreationPricePerToken: 5e-6, + CacheReadPricePerToken: 0.5e-6, + SupportsCacheBreakdown: false, + LongContextInputThreshold: openAIGPT54LongContextInputThreshold, + LongContextInputMultiplier: openAIGPT54LongContextInputMultiplier, + LongContextOutputMultiplier: openAIGPT54LongContextOutputMultiplier, + }, 2.5) + // GPT-5.5 Pro 当前不提供 Fast;保留标准、Flex 和长上下文 fallback 价格。 + s.fallbackPrices["gpt-5.5-pro"] = &ModelPricing{ + InputPricePerToken: 30e-6, + OutputPricePerToken: 180e-6, + // 官方未列独立 cached-input/cache-write 价;内部出现对应 token 时按输入价兜底。 + CacheCreationPricePerToken: 30e-6, + CacheReadPricePerToken: 30e-6, + SupportsCacheBreakdown: false, + LongContextInputThreshold: openAIGPT54LongContextInputThreshold, + LongContextInputMultiplier: openAIGPT54LongContextInputMultiplier, + LongContextOutputMultiplier: openAIGPT54LongContextOutputMultiplier, + } // OpenAI GPT-5.6 官方价格(USD/token)。缓存写入为输入价的 1.25 倍。 s.fallbackPrices["gpt-5.6-sol"] = &ModelPricing{ @@ -1009,25 +1068,9 @@ func (s *BillingService) GetModelPricingWithChannel(model string, channelPricing // 防止修改 fallbackPrices 中的共享指针 cloned := *pricing pricing = &cloned - if channelPricing.InputPrice != nil { - pricing.InputPricePerToken = *channelPricing.InputPrice - pricing.InputPricePerTokenPriority = *channelPricing.InputPrice - } - if channelPricing.OutputPrice != nil { - pricing.OutputPricePerToken = *channelPricing.OutputPrice - pricing.OutputPricePerTokenPriority = *channelPricing.OutputPrice - } - if channelPricing.CacheWritePrice != nil { - pricing.CacheCreationPricePerToken = *channelPricing.CacheWritePrice - pricing.CacheCreationPricePerTokenPriority = *channelPricing.CacheWritePrice - pricing.CacheCreationPriceExplicit = true - pricing.CacheCreation5mPrice = *channelPricing.CacheWritePrice - pricing.CacheCreation1hPrice = *channelPricing.CacheWritePrice - } - if channelPricing.CacheReadPrice != nil { - pricing.CacheReadPricePerToken = *channelPricing.CacheReadPrice - pricing.CacheReadPricePerTokenPriority = *channelPricing.CacheReadPrice - } + applyChannelTokenPriceOverrides(pricing, channelPricing) + pricing.FastMultiplier = channelPricing.FastMultiplier + pricing.FlexMultiplier = channelPricing.FlexMultiplier if channelPricing.ImageOutputPrice != nil { pricing.ImageOutputPricePerToken = *channelPricing.ImageOutputPrice } else { @@ -1038,6 +1081,45 @@ func (s *BillingService) GetModelPricingWithChannel(model string, channelPricing return pricing, nil } +// channelTierOverridePrice applies a Standard-tier override while preserving +// an explicit model-catalog Fast/Priority ratio. If the catalog has no tier +// price, generic service-tier defaults remain responsible for the fallback. +func channelTierOverridePrice(baseStandard, baseTier, channelStandard float64) float64 { + if baseStandard > 0 && baseTier > 0 { + return channelStandard * (baseTier / baseStandard) + } + return 0 +} + +func applyChannelTokenPriceOverrides(pricing *ModelPricing, channelPricing *ChannelModelPricing) { + if pricing == nil || channelPricing == nil { + return + } + if channelPricing.InputPrice != nil { + priority := channelTierOverridePrice(pricing.InputPricePerToken, pricing.InputPricePerTokenPriority, *channelPricing.InputPrice) + pricing.InputPricePerToken = *channelPricing.InputPrice + pricing.InputPricePerTokenPriority = priority + } + if channelPricing.OutputPrice != nil { + priority := channelTierOverridePrice(pricing.OutputPricePerToken, pricing.OutputPricePerTokenPriority, *channelPricing.OutputPrice) + pricing.OutputPricePerToken = *channelPricing.OutputPrice + pricing.OutputPricePerTokenPriority = priority + } + if channelPricing.CacheWritePrice != nil { + priority := channelTierOverridePrice(pricing.CacheCreationPricePerToken, pricing.CacheCreationPricePerTokenPriority, *channelPricing.CacheWritePrice) + pricing.CacheCreationPricePerToken = *channelPricing.CacheWritePrice + pricing.CacheCreationPricePerTokenPriority = priority + pricing.CacheCreationPriceExplicit = true + pricing.CacheCreation5mPrice = *channelPricing.CacheWritePrice + pricing.CacheCreation1hPrice = *channelPricing.CacheWritePrice + } + if channelPricing.CacheReadPrice != nil { + priority := channelTierOverridePrice(pricing.CacheReadPricePerToken, pricing.CacheReadPricePerTokenPriority, *channelPricing.CacheReadPrice) + pricing.CacheReadPricePerToken = *channelPricing.CacheReadPrice + pricing.CacheReadPricePerTokenPriority = priority + } +} + // --- 统一计费入口 --- // CostInput 统一计费输入 @@ -1113,18 +1195,27 @@ func (s *BillingService) CalculateCostUnified(input CostInput) (*CostBreakdown, func (s *BillingService) calculateTokenCost(resolved *ResolvedPricing, input CostInput) (*CostBreakdown, error) { totalContext := input.Tokens.InputTokens + input.Tokens.CacheCreationTokens + input.Tokens.CacheReadTokens - pricing := input.Resolver.GetIntervalPricing(resolved, totalContext) + // 分组开关是统一入口;账号 API 开关保留为额外开启能力,但 false 不否决分组配置。 + contextTierPricingEnabled := resolved.longContextPricingEnabled + if input.LongContextBillingEnabled != nil && *input.LongContextBillingEnabled { + contextTierPricingEnabled = true + } + + pricingContext := totalContext + if !contextTierPricingEnabled { + // 渠道可能显式配置了第一档,也可能只配置高上下文档。用 1 token + // 选择最低档;未命中时自然回退到渠道基础价。 + pricingContext = 1 + } + pricing := input.Resolver.GetIntervalPricing(resolved, pricingContext) if pricing == nil { return nil, fmt.Errorf("no pricing available for model: %s: %w", input.Model, ErrModelPricingUnavailable) } pricing = s.applyModelSpecificPricingPolicy(input.Model, pricing) - // 长上下文定价仅在无区间定价时应用(区间定价已包含上下文分层) - applyLongCtx := len(resolved.Intervals) == 0 && resolved.longContextPricingEnabled - if input.LongContextBillingEnabled != nil { - applyLongCtx = applyLongCtx && *input.LongContextBillingEnabled - } + // 官方长上下文阶梯仅在无区间定价时应用(区间定价已包含上下文分层)。 + applyLongCtx := len(resolved.Intervals) == 0 && contextTierPricingEnabled breakdown := s.computeTokenBreakdown(pricing, input.Tokens, input.RateMultiplier, input.ServiceTier, applyLongCtx) applyCostBreakdownMultiplier(breakdown, resolvedChannelTimeMultiplier(resolved, input.PricingAt)) @@ -1164,7 +1255,7 @@ func (s *BillingService) computeTokenBreakdown( cacheCreationPrice = pricing.CacheCreationPricePerTokenPriority } } else { - tierMultiplier = serviceTierCostMultiplier(serviceTier) + tierMultiplier = configuredServiceTierMultiplier(serviceTier, pricing) } longContextPricingEligible := applyLongCtx && s.shouldApplySessionLongContextPricing(tokens, pricing) diff --git a/backend/internal/service/billing_service_test.go b/backend/internal/service/billing_service_test.go index 80d07af44b..72d3627118 100644 --- a/backend/internal/service/billing_service_test.go +++ b/backend/internal/service/billing_service_test.go @@ -293,14 +293,40 @@ func TestCalculateCost_OpenAIGPT55ProUsesGPT55PricingPolicy(t *testing.T) { cost, err := svc.CalculateCost("gpt-5.5-pro", tokens, 1.0) require.NoError(t, err) - expectedInput := float64(tokens.InputTokens) * 2.5e-6 * 2.0 - expectedOutput := float64(tokens.OutputTokens) * 15e-6 * 1.5 + expectedInput := float64(tokens.InputTokens) * 30e-6 * 2.0 + expectedOutput := float64(tokens.OutputTokens) * 180e-6 * 1.5 require.InDelta(t, expectedInput, cost.InputCost, 1e-10) require.InDelta(t, expectedOutput, cost.OutputCost, 1e-10) require.InDelta(t, expectedInput+expectedOutput, cost.TotalCost, 1e-10) require.InDelta(t, expectedInput+expectedOutput, cost.ActualCost, 1e-10) } +func TestFallbackPricing_OpenAIGPT55UsesOfficialPrices(t *testing.T) { + svc := newTestBillingService() + + pricing, err := svc.GetModelPricing("gpt-5.5") + require.NoError(t, err) + require.InDelta(t, 5e-6, pricing.InputPricePerToken, 1e-12) + require.InDelta(t, 30e-6, pricing.OutputPricePerToken, 1e-12) + require.InDelta(t, 0.5e-6, pricing.CacheReadPricePerToken, 1e-12) + require.InDelta(t, 5e-6, pricing.CacheCreationPricePerToken, 1e-12) + require.InDelta(t, 12.5e-6, pricing.InputPricePerTokenPriority, 1e-12) + require.InDelta(t, 75e-6, pricing.OutputPricePerTokenPriority, 1e-12) +} + +func TestFallbackPricing_OpenAIGPT55ProUsesOfficialPrices(t *testing.T) { + svc := newTestBillingService() + + pricing, err := svc.GetModelPricing("gpt-5.5-pro") + require.NoError(t, err) + require.InDelta(t, 30e-6, pricing.InputPricePerToken, 1e-12) + require.InDelta(t, 180e-6, pricing.OutputPricePerToken, 1e-12) + require.InDelta(t, 30e-6, pricing.CacheReadPricePerToken, 1e-12) + require.InDelta(t, 30e-6, pricing.CacheCreationPricePerToken, 1e-12) + require.Zero(t, pricing.InputPricePerTokenPriority) + require.Zero(t, pricing.OutputPricePerTokenPriority) +} + // 回归测试 #2293:长上下文计费触发时,cache_read_tokens 也应应用 LongContextInputMultiplier。 // 修复前:CacheReadCost = tokens * 0.25e-6 (漏乘倍率,少计费用)。 // 修复后:CacheReadCost = tokens * 0.25e-6 * LongContextInputMultiplier(=2.0)。 @@ -1594,9 +1620,10 @@ func TestGetModelPricingWithChannel_OverrideInputPriceOnly(t *testing.T) { pricing, err := svc.GetModelPricingWithChannel("claude-sonnet-4", chPricing) require.NoError(t, err) - // InputPrice overridden (both normal and priority) + // InputPrice overridden. claude-sonnet-4 has no catalog priority price, so + // the priority slot is zeroed and serviceTierCostMultiplier owns the surcharge. require.InDelta(t, 99e-6, pricing.InputPricePerToken, 1e-12) - require.InDelta(t, 99e-6, pricing.InputPricePerTokenPriority, 1e-12) + require.Zero(t, pricing.InputPricePerTokenPriority) // OutputPrice unchanged (claude-sonnet-4 fallback = 15e-6) require.InDelta(t, 15e-6, pricing.OutputPricePerToken, 1e-12) @@ -1611,9 +1638,9 @@ func TestGetModelPricingWithChannel_OverrideOutputPriceOnly(t *testing.T) { pricing, err := svc.GetModelPricingWithChannel("claude-sonnet-4", chPricing) require.NoError(t, err) - // OutputPrice overridden + // OutputPrice overridden; no catalog priority price to scale, so the slot is zeroed. require.InDelta(t, 88e-6, pricing.OutputPricePerToken, 1e-12) - require.InDelta(t, 88e-6, pricing.OutputPricePerTokenPriority, 1e-12) + require.Zero(t, pricing.OutputPricePerTokenPriority) // InputPrice unchanged (claude-sonnet-4 fallback = 3e-6) require.InDelta(t, 3e-6, pricing.InputPricePerToken, 1e-12) @@ -1633,15 +1660,18 @@ func TestGetModelPricingWithChannel_OverrideAllFields(t *testing.T) { require.NoError(t, err) require.InDelta(t, 10e-6, pricing.InputPricePerToken, 1e-12) - require.InDelta(t, 10e-6, pricing.InputPricePerTokenPriority, 1e-12) require.InDelta(t, 20e-6, pricing.OutputPricePerToken, 1e-12) - require.InDelta(t, 20e-6, pricing.OutputPricePerTokenPriority, 1e-12) require.InDelta(t, 5e-6, pricing.CacheCreationPricePerToken, 1e-12) require.InDelta(t, 5e-6, pricing.CacheCreation5mPrice, 1e-12) require.InDelta(t, 5e-6, pricing.CacheCreation1hPrice, 1e-12) require.InDelta(t, 1e-6, pricing.CacheReadPricePerToken, 1e-12) - require.InDelta(t, 1e-6, pricing.CacheReadPricePerTokenPriority, 1e-12) require.InDelta(t, 50e-6, pricing.ImageOutputPricePerToken, 1e-12) + + // claude-sonnet-4 carries no catalog Fast/Priority tier, so every priority + // slot stays zero and computeTokenBreakdown falls back to the 2x default. + require.Zero(t, pricing.InputPricePerTokenPriority) + require.Zero(t, pricing.OutputPricePerTokenPriority) + require.Zero(t, pricing.CacheReadPricePerTokenPriority) } func TestGetModelPricingWithChannel_CacheWritePriceAffects5mAnd1h(t *testing.T) { @@ -1668,9 +1698,27 @@ func TestGetModelPricingWithChannel_CacheReadPriceAffectsPriority(t *testing.T) pricing, err := svc.GetModelPricingWithChannel("claude-sonnet-4", chPricing) require.NoError(t, err) - // CacheReadPrice should set both normal and priority + // CacheReadPrice sets the standard slot; the priority slot is zeroed because + // claude-sonnet-4 has no catalog tier ratio to preserve. require.InDelta(t, 2e-6, pricing.CacheReadPricePerToken, 1e-12) - require.InDelta(t, 2e-6, pricing.CacheReadPricePerTokenPriority, 1e-12) + require.Zero(t, pricing.CacheReadPricePerTokenPriority) +} + +// 目录带 tier 价时,渠道覆盖必须按目录比例换算 priority 价,而不是归零。 +func TestGetModelPricingWithChannel_PreservesCatalogPriorityRatio(t *testing.T) { + svc := newTestBillingService() + + // gpt-5.4 目录价:input 2.5/5(2x),output 15/30(2x)。 + pricing, err := svc.GetModelPricingWithChannel("gpt-5.4", &ChannelModelPricing{ + InputPrice: testPtrFloat64(4e-6), + OutputPrice: testPtrFloat64(30e-6), + }) + require.NoError(t, err) + + require.InDelta(t, 4e-6, pricing.InputPricePerToken, 1e-12) + require.InDelta(t, 8e-6, pricing.InputPricePerTokenPriority, 1e-12) + require.InDelta(t, 30e-6, pricing.OutputPricePerToken, 1e-12) + require.InDelta(t, 60e-6, pricing.OutputPricePerTokenPriority, 1e-12) } func TestGetModelPricingWithChannel_UnknownModelReturnsError(t *testing.T) { diff --git a/backend/internal/service/model_pricing_resolver.go b/backend/internal/service/model_pricing_resolver.go index 52c6b930c5..62283e81b7 100644 --- a/backend/internal/service/model_pricing_resolver.go +++ b/backend/internal/service/model_pricing_resolver.go @@ -120,14 +120,8 @@ func (r *ModelPricingResolver) Resolve(ctx context.Context, input PricingInput) resolved.Source = PricingSourceChannel resolved.channelPricing = chPricing r.applyTokenOverrides(chPricing, resolved) - if !longContextPricingEnabled { - r.applyFirstTokenTier(resolved, chPricing) - } } else if input.GroupID != nil && r.channelService != nil { r.applyChannelOverrides(ctx, *input.GroupID, input.Model, resolved) - if resolved.Source == PricingSourceChannel && !longContextPricingEnabled { - r.applyFirstTokenTier(resolved, resolved.channelPricing) - } } return resolved @@ -172,20 +166,6 @@ func matchGroupModelPricing(group *Group, model string) *ChannelModelPricing { return wildcard } -func (r *ModelPricingResolver) applyFirstTokenTier(resolved *ResolvedPricing, config *ChannelModelPricing) { - if resolved == nil || len(resolved.Intervals) == 0 { - return - } - first := resolved.Intervals[0] - for _, interval := range resolved.Intervals[1:] { - if interval.MinTokens < first.MinTokens { - first = interval - } - } - resolved.BasePricing = intervalToModelPricing(&first, resolved.SupportsCacheBreakdown, config) - resolved.Intervals = nil -} - // resolveBasePricing 从 LiteLLM 或 Fallback 获取基础定价 func (r *ModelPricingResolver) resolveBasePricing(model string) (*ModelPricing, string) { pricing, err := r.billingService.GetModelPricing(model) @@ -221,31 +201,6 @@ func (r *ModelPricingResolver) applyChannelOverrides(ctx context.Context, groupI // applyTokenOverrides 应用 token 模式的渠道覆盖 func (r *ModelPricingResolver) applyTokenOverrides(chPricing *ChannelModelPricing, resolved *ResolvedPricing) { - // 过滤掉所有价格字段都为空的无效 interval - validIntervals := filterValidIntervals(chPricing.Intervals) - - // 如果有有效的区间定价,使用区间 - if len(validIntervals) > 0 { - resolved.Intervals = validIntervals - // 区间不匹配时回退到 BasePricing,也需要覆盖图片价格 - if resolved.BasePricing == nil { - resolved.BasePricing = &ModelPricing{} - } else { - // 防止修改 fallbackPrices 中的共享指针 - cloned := *resolved.BasePricing - resolved.BasePricing = &cloned - } - if chPricing.ImageOutputPrice != nil { - resolved.BasePricing.ImageOutputPricePerToken = *chPricing.ImageOutputPrice - } else { - resolved.BasePricing.ImageOutputPricePerToken = 0 - } - resolved.BasePricing.ImageOutputPriceExplicit = true - applyChannelImageInputPrice(chPricing, resolved.BasePricing) - return - } - - // 否则用 flat 字段覆盖 BasePricing if resolved.BasePricing == nil { resolved.BasePricing = &ModelPricing{} } else { @@ -254,25 +209,9 @@ func (r *ModelPricingResolver) applyTokenOverrides(chPricing *ChannelModelPricin resolved.BasePricing = &cloned } - if chPricing.InputPrice != nil { - resolved.BasePricing.InputPricePerToken = *chPricing.InputPrice - resolved.BasePricing.InputPricePerTokenPriority = *chPricing.InputPrice - } - if chPricing.OutputPrice != nil { - resolved.BasePricing.OutputPricePerToken = *chPricing.OutputPrice - resolved.BasePricing.OutputPricePerTokenPriority = *chPricing.OutputPrice - } - if chPricing.CacheWritePrice != nil { - resolved.BasePricing.CacheCreationPricePerToken = *chPricing.CacheWritePrice - resolved.BasePricing.CacheCreationPricePerTokenPriority = *chPricing.CacheWritePrice - resolved.BasePricing.CacheCreationPriceExplicit = true - resolved.BasePricing.CacheCreation5mPrice = *chPricing.CacheWritePrice - resolved.BasePricing.CacheCreation1hPrice = *chPricing.CacheWritePrice - } - if chPricing.CacheReadPrice != nil { - resolved.BasePricing.CacheReadPricePerToken = *chPricing.CacheReadPrice - resolved.BasePricing.CacheReadPricePerTokenPriority = *chPricing.CacheReadPrice - } + applyChannelTokenPriceOverrides(resolved.BasePricing, chPricing) + resolved.BasePricing.FastMultiplier = chPricing.FastMultiplier + resolved.BasePricing.FlexMultiplier = chPricing.FlexMultiplier // 渠道定价覆盖一切:显式配置则用配置值,未配置则归零(不回退到 LiteLLM) if chPricing.ImageOutputPrice != nil { resolved.BasePricing.ImageOutputPricePerToken = *chPricing.ImageOutputPrice @@ -281,6 +220,9 @@ func (r *ModelPricingResolver) applyTokenOverrides(chPricing *ChannelModelPricin } resolved.BasePricing.ImageOutputPriceExplicit = true applyChannelImageInputPrice(chPricing, resolved.BasePricing) + + // 区间未命中时回退到上面已经应用渠道覆盖的基础价。 + resolved.Intervals = filterValidIntervals(chPricing.Intervals) } // applyChannelImageInputPrice 应用渠道图片输入价:显式配置则用配置值; @@ -311,7 +253,9 @@ func filterValidIntervals(intervals []PricingInterval) []PricingInterval { for _, iv := range intervals { if iv.InputPrice != nil || iv.OutputPrice != nil || iv.CacheWritePrice != nil || iv.CacheReadPrice != nil || - iv.PerRequestPrice != nil { + iv.PerRequestPrice != nil || iv.InputMultiplier != nil || + iv.OutputMultiplier != nil || iv.CacheWriteMultiplier != nil || + iv.CacheReadMultiplier != nil { valid = append(valid, iv) } } @@ -330,32 +274,57 @@ func (r *ModelPricingResolver) GetIntervalPricing(resolved *ResolvedPricing, tot return resolved.BasePricing } - return intervalToModelPricing(iv, resolved.SupportsCacheBreakdown, resolved.channelPricing) + pricing := intervalToModelPricing(iv, resolved.BasePricing, resolved.channelPricing) + // BasePricing 为 nil(仅配置区间)时拷贝不到该标志,从 resolved 回填, + // 保证 computeCacheCreationCost 的 5m/1h 分档判断不被区间路径吞掉。 + pricing.SupportsCacheBreakdown = resolved.SupportsCacheBreakdown + return pricing } // intervalToModelPricing 将区间定价转换为 ModelPricing -func intervalToModelPricing(iv *PricingInterval, supportsCacheBreakdown bool, chPricing *ChannelModelPricing) *ModelPricing { - pricing := &ModelPricing{ - SupportsCacheBreakdown: supportsCacheBreakdown, +func intervalToModelPricing(iv *PricingInterval, base *ModelPricing, chPricing *ChannelModelPricing) *ModelPricing { + pricing := &ModelPricing{} + if base != nil { + *pricing = *base + } + applyMultiplier := func(value float64, multiplier *float64) float64 { + if multiplier == nil { + return value + } + return value * *multiplier } if iv.InputPrice != nil { + pricing.InputPricePerTokenPriority = channelTierOverridePrice(pricing.InputPricePerToken, pricing.InputPricePerTokenPriority, *iv.InputPrice) pricing.InputPricePerToken = *iv.InputPrice - pricing.InputPricePerTokenPriority = *iv.InputPrice + } else if iv.InputMultiplier != nil { + pricing.InputPricePerToken = applyMultiplier(pricing.InputPricePerToken, iv.InputMultiplier) + pricing.InputPricePerTokenPriority = applyMultiplier(pricing.InputPricePerTokenPriority, iv.InputMultiplier) } if iv.OutputPrice != nil { + pricing.OutputPricePerTokenPriority = channelTierOverridePrice(pricing.OutputPricePerToken, pricing.OutputPricePerTokenPriority, *iv.OutputPrice) pricing.OutputPricePerToken = *iv.OutputPrice - pricing.OutputPricePerTokenPriority = *iv.OutputPrice + } else if iv.OutputMultiplier != nil { + pricing.OutputPricePerToken = applyMultiplier(pricing.OutputPricePerToken, iv.OutputMultiplier) + pricing.OutputPricePerTokenPriority = applyMultiplier(pricing.OutputPricePerTokenPriority, iv.OutputMultiplier) } if iv.CacheWritePrice != nil { + pricing.CacheCreationPricePerTokenPriority = channelTierOverridePrice(pricing.CacheCreationPricePerToken, pricing.CacheCreationPricePerTokenPriority, *iv.CacheWritePrice) pricing.CacheCreationPricePerToken = *iv.CacheWritePrice - pricing.CacheCreationPricePerTokenPriority = *iv.CacheWritePrice pricing.CacheCreationPriceExplicit = true pricing.CacheCreation5mPrice = *iv.CacheWritePrice pricing.CacheCreation1hPrice = *iv.CacheWritePrice + } else if iv.CacheWriteMultiplier != nil { + pricing.CacheCreationPricePerToken = applyMultiplier(pricing.CacheCreationPricePerToken, iv.CacheWriteMultiplier) + pricing.CacheCreationPricePerTokenPriority = applyMultiplier(pricing.CacheCreationPricePerTokenPriority, iv.CacheWriteMultiplier) + pricing.CacheCreation5mPrice = applyMultiplier(pricing.CacheCreation5mPrice, iv.CacheWriteMultiplier) + pricing.CacheCreation1hPrice = applyMultiplier(pricing.CacheCreation1hPrice, iv.CacheWriteMultiplier) } if iv.CacheReadPrice != nil { + pricing.CacheReadPricePerTokenPriority = channelTierOverridePrice(pricing.CacheReadPricePerToken, pricing.CacheReadPricePerTokenPriority, *iv.CacheReadPrice) pricing.CacheReadPricePerToken = *iv.CacheReadPrice - pricing.CacheReadPricePerTokenPriority = *iv.CacheReadPrice + } else if iv.CacheReadMultiplier != nil { + pricing.CacheReadPricePerToken = applyMultiplier(pricing.CacheReadPricePerToken, iv.CacheReadMultiplier) + pricing.CacheReadPricePerTokenPriority = applyMultiplier(pricing.CacheReadPricePerTokenPriority, iv.CacheReadMultiplier) } // 渠道定价存在时,ImageOutputPrice 显式覆盖;图片输入价用渠道级配置 // (区间不携带图片输入价,与 image_output 一致)。 diff --git a/backend/internal/service/model_pricing_resolver_test.go b/backend/internal/service/model_pricing_resolver_test.go index 11613fa9eb..3fb3e1c63f 100644 --- a/backend/internal/service/model_pricing_resolver_test.go +++ b/backend/internal/service/model_pricing_resolver_test.go @@ -142,7 +142,7 @@ func TestGPT56ExplicitZeroCacheWritePriceIsPreserved(t *testing.T) { }) t.Run("interval price", func(t *testing.T) { - pricing := intervalToModelPricing(&PricingInterval{CacheWritePrice: &zero}, false, nil) + pricing := intervalToModelPricing(&PricingInterval{CacheWritePrice: &zero}, &ModelPricing{}, nil) require.True(t, pricing.CacheCreationPriceExplicit) cost, err := bs.CalculateCostUnified(CostInput{ @@ -261,9 +261,9 @@ func TestResolve_WithChannelOverride_TokenFlat(t *testing.T) { require.Equal(t, "channel", resolved.Source) require.NotNil(t, resolved.BasePricing) require.InDelta(t, 10e-6, resolved.BasePricing.InputPricePerToken, 1e-12) - require.InDelta(t, 10e-6, resolved.BasePricing.InputPricePerTokenPriority, 1e-12) + require.Zero(t, resolved.BasePricing.InputPricePerTokenPriority) require.InDelta(t, 50e-6, resolved.BasePricing.OutputPricePerToken, 1e-12) - require.InDelta(t, 50e-6, resolved.BasePricing.OutputPricePerTokenPriority, 1e-12) + require.Zero(t, resolved.BasePricing.OutputPricePerTokenPriority) } func TestResolve_WithChannelOverride_TokenPartialOverride(t *testing.T) { @@ -304,10 +304,12 @@ func TestResolve_WithChannelOverride_TokenWithIntervals(t *testing.T) { resolved := r.Resolve(context.Background(), PricingInput{ Model: "claude-sonnet-4", GroupID: groupIDPtr(), + Group: &Group{LongContextPricingEnabled: false}, }) require.NotNil(t, resolved) require.Equal(t, "channel", resolved.Source) + require.False(t, resolved.longContextPricingEnabled) require.Len(t, resolved.Intervals, 2) // GetIntervalPricing should use channel intervals @@ -532,6 +534,7 @@ func TestGetIntervalPricing_ChannelIntervalsNoMatch(t *testing.T) { Platform: "anthropic", Models: []string{"claude-sonnet-4"}, BillingMode: BillingModeToken, + InputPrice: testPtrFloat64(4e-6), Intervals: []PricingInterval{ // Only covers tokens > 50000 {MinTokens: 50000, MaxTokens: testPtrInt(200000), InputPrice: testPtrFloat64(9e-6)}, @@ -545,10 +548,10 @@ func TestGetIntervalPricing_ChannelIntervalsNoMatch(t *testing.T) { // Token count 1000 doesn't match any interval (1000 <= 50000 minTokens) pricing := r.GetIntervalPricing(resolved, 1000) - // Should fall back to BasePricing (from the billing service fallback) + // Should fall back to BasePricing after applying the channel default. require.NotNil(t, pricing) require.Equal(t, resolved.BasePricing, pricing) - require.InDelta(t, 3e-6, pricing.InputPricePerToken, 1e-12) // original base price + require.InDelta(t, 4e-6, pricing.InputPricePerToken, 1e-12) } // =========================================================================== @@ -689,6 +692,13 @@ func TestFilterValidIntervals(t *testing.T) { }, wantLen: 1, }, + { + name: "interval with only multiplier kept", + intervals: []PricingInterval{ + {MinTokens: 272000, InputMultiplier: testPtrFloat64(2)}, + }, + wantLen: 1, + }, { name: "mixed valid and invalid", intervals: []PricingInterval{ diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go index 21e5cd4096..674e733f9c 100644 --- a/backend/internal/service/openai_gateway_record_usage_test.go +++ b/backend/internal/service/openai_gateway_record_usage_test.go @@ -1164,7 +1164,7 @@ func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingDisabledByDefaul Model: "gpt-5.4-2026-03-05", Duration: time.Second, }, - APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1014, true), + APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1014, false), User: &User{ID: 2014}, Account: &Account{ID: 3014, Platform: PlatformOpenAI}, }) @@ -1219,7 +1219,7 @@ func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingEnabledPerAccoun require.True(t, usageRepo.lastLog.LongContextBillingApplied) } -func TestOpenAIGatewayServiceRecordUsage_GroupAndAccountLongContextMustBothAllow(t *testing.T) { +func TestOpenAIGatewayServiceRecordUsage_GroupOrAccountLongContextAllows(t *testing.T) { tokens := OpenAIUsage{InputTokens: 300000, OutputTokens: 2000} baseInput := 300000 * 2.5e-6 baseOutput := 2000 * 15e-6 @@ -1234,9 +1234,9 @@ func TestOpenAIGatewayServiceRecordUsage_GroupAndAccountLongContextMustBothAllow Account: &Account{ID: 3020, Platform: PlatformOpenAI}, }) require.NoError(t, err) - require.False(t, usageRepo.lastLog.LongContextBillingApplied) - require.InDelta(t, baseInput, usageRepo.lastLog.InputCost, 1e-10) - require.InDelta(t, baseOutput, usageRepo.lastLog.OutputCost, 1e-10) + require.True(t, usageRepo.lastLog.LongContextBillingApplied) + require.InDelta(t, baseInput*2, usageRepo.lastLog.InputCost, 1e-10) + require.InDelta(t, baseOutput*1.5, usageRepo.lastLog.OutputCost, 1e-10) }) t.Run("group off account on", func(t *testing.T) { @@ -1252,8 +1252,9 @@ func TestOpenAIGatewayServiceRecordUsage_GroupAndAccountLongContextMustBothAllow }, }) require.NoError(t, err) - require.False(t, usageRepo.lastLog.LongContextBillingApplied) - require.InDelta(t, baseInput, usageRepo.lastLog.InputCost, 1e-10) + require.True(t, usageRepo.lastLog.LongContextBillingApplied) + require.InDelta(t, baseInput*2, usageRepo.lastLog.InputCost, 1e-10) + require.InDelta(t, baseOutput*1.5, usageRepo.lastLog.OutputCost, 1e-10) }) t.Run("group on account on", func(t *testing.T) { @@ -1361,7 +1362,7 @@ func TestOpenAIGatewayServiceRecordUsage_SparkShadowUsesCurrentParentBillingSett Model: "gpt-5.4-2026-03-05", Duration: time.Second, }, - APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1016, true), + APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1016, false), User: &User{ID: 2016}, Account: &Account{ ID: 3016,