计费:统一 token 计费路径选择并提供上下文阶梯单价表查询

- BillingService.CalculateTokenCostForRequest 承接网关的路径选择
  (分组/渠道定价 → 平台旧长上下文规则 → 内置目录),网关改为调用该入口
- Gemini /v1beta 的 200K 边际翻倍常量从 handler 移入
  BillingService.LegacyLongContextRule,入口只声明适用
- 新增 ResolveContextPricingSchedule:沿用 Resolve 解析链收集断点
  (渠道区间边界、目录阶梯阈值、旧规则阈值),每档单价由真实计费函数
  探针差商得出,倍率/策略变更无需同步;附阶梯表 vs 计费函数的对账测试
This commit is contained in:
feeeei
2026-08-24 10:50:52 +08:00
parent 3b8a148bcf
commit 6466978d2f
6 changed files with 1199 additions and 34 deletions
@@ -569,6 +569,13 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
forceCacheBilling := fs.ForceCacheBilling
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
sessionID := service.ExtractClientSessionID(c)
// 长上下文规则由计费服务统一持有(模型广场展示同源),入口只负责声明自己适用该规则。
var longContextThreshold int
var longContextMultiplier float64
if rule := h.gatewayService.LegacyLongContextRule(service.PlatformGemini); rule != nil {
longContextThreshold = rule.Threshold
longContextMultiplier = rule.Multiplier
}
h.submitUsageRecordTask(c.Request.Context(), func(ctx context.Context) {
if err := h.gatewayService.RecordUsageWithLongContext(ctx, &service.RecordUsageLongContextInput{
Result: result,
@@ -583,8 +590,8 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
UserAgent: userAgent,
IPAddress: clientIP,
RequestPayloadHash: requestPayloadHash,
LongContextThreshold: 200000, // Gemini 200K 阈值
LongContextMultiplier: 2.0, // 超出部分双倍计费
LongContextThreshold: longContextThreshold,
LongContextMultiplier: longContextMultiplier,
ForceCacheBilling: forceCacheBilling,
APIKeyService: h.apiKeyService,
SessionID: sessionID,
@@ -0,0 +1,424 @@
package service
import (
"context"
"errors"
"math"
"sort"
"strconv"
)
// ContextPricingBasis 阶梯的计价基准。
type ContextPricingBasis string
const (
// ContextPricingBasisWholeRequest 整单按所在档单价计价(目录阶梯、渠道区间)。
ContextPricingBasisWholeRequest ContextPricingBasis = "whole_request"
// ContextPricingBasisMarginal 仅超出阈值的部分按该档单价计价(平台旧规则)。
ContextPricingBasisMarginal ContextPricingBasis = "marginal"
)
// ContextPricingTier (MinTokens, MaxTokens] 区间内的有效 per-token 单价(USD)。
// nil 表示该项无价/不计费;MaxTokens 为 nil 表示无上限。
type ContextPricingTier struct {
MinTokens int
MaxTokens *int
Label string
Input *float64
Output *float64
CacheWrite *float64
CacheRead *float64
}
// ContextPricingSchedule 分组+模型按上下文长度分档的有效单价表。
// 单价由真实计费函数探针得出,与扣费同源;单档表示无阶梯。
type ContextPricingSchedule struct {
Basis ContextPricingBasis
Tiers []ContextPricingTier
}
// ContextPricingScheduleInput 阶梯表查询输入。
type ContextPricingScheduleInput struct {
Model string
// Group 为 nil 表示查官方参考价:无分组、无渠道定价,也不套用平台旧规则。
Group *Group
// Platform 为请求的具体平台(composite 分组传模型所属平台),
// 决定渠道定价查找与平台旧规则的适用。
Platform string
}
var errContextPricingResolverRequired = errors.New("context pricing schedule: resolver is required")
// 探针步长:相邻两个探针点相差 contextProbeDelta 个 token,单价 = Δcost / Δtoken。
const contextProbeDelta = 1000
// ResolveContextPricingSchedule 解析分组+模型的上下文阶梯单价表。
//
// 解析链与扣费完全一致:Resolver.Resolve(分组卡 → 渠道 → 目录 → 策略)给出定价,
// CalculateTokenCostForRequest 给出路径(分组/渠道定价 → 平台旧规则 → 内置目录)。
// 断点只取自计费自身的规则输入(渠道区间边界、目录阶梯阈值、旧规则阈值),
// 每一段的单价由真实计费函数在该段内两点探针的差商得到,因此倍率、策略等
// 规则变更无需同步到这里;相邻同价段会合并。
//
// 非 token 计费模式返回 (nil, nil);模型无任何定价来源时返回 ErrModelPricingUnavailable。
func (s *BillingService) ResolveContextPricingSchedule(ctx context.Context, resolver *ModelPricingResolver, in ContextPricingScheduleInput) (*ContextPricingSchedule, error) {
if s == nil || resolver == nil {
return nil, errContextPricingResolverRequired
}
if ctx == nil {
ctx = context.Background()
}
if in.Platform != "" {
ctx = WithResolvedTargetPlatform(ctx, in.Platform)
}
pricingInput := PricingInput{Model: in.Model, Group: in.Group}
if in.Group != nil {
gid := in.Group.ID
pricingInput.GroupID = &gid
}
resolved := resolver.Resolve(ctx, pricingInput)
if resolved == nil {
return nil, ErrModelPricingUnavailable
}
if resolved.Mode != "" && resolved.Mode != BillingModeToken {
return nil, nil
}
var legacy *LegacyLongContextRule
if in.Group != nil {
legacy = s.LegacyLongContextRule(in.Platform)
}
if !legacyLongContextApplies(resolved, in.Group, legacy) {
legacy = nil
}
req := TokenCostRequest{
Ctx: ctx,
Model: in.Model,
Group: in.Group,
RateMultiplier: 1,
Resolver: resolver,
Resolved: resolved,
LegacyLongContext: legacy,
}
probe := func(tokens UsageTokens) (*CostBreakdown, error) {
r := req
r.Tokens = tokens
return s.CalculateTokenCostForRequest(r)
}
plan := s.contextPricingBreakpoints(resolver, resolved, in.Model, legacy)
segments := buildContextSegments(plan.bounds)
tiers := make([]ContextPricingTier, 0, len(segments))
for _, seg := range segments {
tier, err := probeContextTier(seg, resolved, probe)
if err != nil {
return nil, err
}
tiers = append(tiers, tier)
}
if legacy == nil {
applyIntervalContextLabels(tiers, resolved.Intervals)
}
tiers = mergeEqualContextTiers(tiers)
applyGeneratedContextLabels(tiers, plan)
basis := ContextPricingBasisWholeRequest
if legacy != nil {
basis = ContextPricingBasisMarginal
}
return &ContextPricingSchedule{Basis: basis, Tiers: tiers}, nil
}
// contextBreakpointPlan 描述断点来源。
type contextBreakpointPlan struct {
bounds []int
// thresholdBound 为目录阶梯/旧规则的断点值((0,b] / (b,∞)),0 表示无。
thresholdBound int
// thresholdInclusive 为真表示达到阈值即进入高档(断点 = 阈值-1)。
thresholdInclusive bool
threshold int
}
// contextPricingBreakpoints 从计费自身的规则输入收集价格断点(不读取任何倍率)。
func (s *BillingService) contextPricingBreakpoints(resolver *ModelPricingResolver, resolved *ResolvedPricing, model string, legacy *LegacyLongContextRule) contextBreakpointPlan {
plan := contextBreakpointPlan{}
if legacy != nil {
plan.bounds = []int{legacy.Threshold}
plan.thresholdBound = legacy.Threshold
plan.threshold = legacy.Threshold
return plan
}
if !resolved.longContextPricingEnabled {
return plan
}
if len(resolved.Intervals) > 0 {
// 区间边界即断点;空洞段(含末档上限之外)由计费回落 base,探针会自然得到基础价。
set := make(map[int]struct{}, len(resolved.Intervals)*2)
for i := range resolved.Intervals {
iv := &resolved.Intervals[i]
if iv.MinTokens > 0 {
set[iv.MinTokens] = struct{}{}
}
if iv.MaxTokens != nil {
set[*iv.MaxTokens] = struct{}{}
}
}
for b := range set {
plan.bounds = append(plan.bounds, b)
}
sort.Ints(plan.bounds)
return plan
}
pricing := resolver.GetIntervalPricing(resolved, 1)
if pricing == nil {
return plan
}
pricing = s.applyModelSpecificPricingPolicy(model, pricing)
if pricing.LongContextInputThreshold <= 0 {
return plan
}
bound := pricing.LongContextInputThreshold
if pricing.LongContextThresholdInclusive {
bound--
}
if bound <= 0 {
return plan
}
plan.bounds = []int{bound}
plan.thresholdBound = bound
plan.threshold = pricing.LongContextInputThreshold
plan.thresholdInclusive = pricing.LongContextThresholdInclusive
return plan
}
// contextSegment 探针用的 (min, max] 段;max 为 nil 表示无上限。
type contextSegment struct {
min int
max *int
}
// buildContextSegments 把升序断点切成 (0,b1], (b1,b2], …, (bn,∞);无断点时为单个开区间。
func buildContextSegments(bounds []int) []contextSegment {
segments := make([]contextSegment, 0, len(bounds)+1)
prev := 0
for _, b := range bounds {
if b <= prev {
continue
}
upper := b
segments = append(segments, contextSegment{min: prev, max: &upper})
prev = b
}
segments = append(segments, contextSegment{min: prev})
return segments
}
// probeContextTier 在段内两点探针,单价 = ΔActualCost / Δtoken(倍率固定为 1)。
// 每次只喂一种 token,ActualCost 即该项费用;旧边际规则的加倍只体现在 ActualCost
// 而不在分项费用里,因此统一读 ActualCost。整单阶梯与边际规则在同一段内都是
// 线性函数,差商同时适用,无需区分规则类型。
func probeContextTier(seg contextSegment, resolved *ResolvedPricing, probe func(UsageTokens) (*CostBreakdown, error)) (ContextPricingTier, error) {
tier := ContextPricingTier{MinTokens: seg.min, MaxTokens: seg.max}
c := seg.min + 1
delta := contextProbeDelta
if seg.max != nil {
if width := *seg.max - seg.min; width-1 < delta {
delta = width - 1
}
}
var err error
tier.Input, err = probeComponentPrice(func(n int) UsageTokens { return UsageTokens{InputTokens: n} }, c, delta, probe)
if err != nil {
return tier, err
}
tier.CacheRead, err = probeComponentPrice(func(n int) UsageTokens { return UsageTokens{CacheReadTokens: n} }, c, delta, probe)
if err != nil {
return tier, err
}
tier.CacheWrite, err = probeComponentPrice(func(n int) UsageTokens { return UsageTokens{CacheCreationTokens: n} }, c, delta, probe)
if err != nil {
return tier, err
}
// 输出价只随上下文所在档变化:固定上下文 c,对输出 token 数做差商(固定部分相减抵消)。
tier.Output, err = probeComponentPrice(func(n int) UsageTokens { return UsageTokens{InputTokens: c, OutputTokens: n} }, 0, contextProbeDelta, probe)
if err != nil {
return tier, err
}
explicit := explicitContextPricingFields(resolved, c)
tier.Input = contextPricePtr(tier.Input, explicit.input)
tier.Output = contextPricePtr(tier.Output, explicit.output)
tier.CacheWrite = contextPricePtr(tier.CacheWrite, explicit.cacheWrite)
tier.CacheRead = contextPricePtr(tier.CacheRead, explicit.cacheRead)
return tier, nil
}
// probeComponentPrice 返回 [from, from+delta] 上 ActualCost 的差商;delta 为 0(退化段)时退回平均单价。
func probeComponentPrice(tokensAt func(int) UsageTokens, from, delta int, probe func(UsageTokens) (*CostBreakdown, error)) (*float64, error) {
if delta <= 0 {
n := from
if n <= 0 {
n = 1
}
cost, err := probe(tokensAt(n))
if err != nil {
return nil, err
}
v := roundContextPrice(cost.ActualCost / float64(n))
return &v, nil
}
lo, err := probe(tokensAt(from))
if err != nil {
return nil, err
}
hi, err := probe(tokensAt(from + delta))
if err != nil {
return nil, err
}
v := roundContextPrice((hi.ActualCost - lo.ActualCost) / float64(delta))
return &v, nil
}
// roundContextPrice 去掉差商带来的浮点噪声(保留 12 位有效数字)。
func roundContextPrice(v float64) float64 {
if v == 0 || math.IsNaN(v) || math.IsInf(v, 0) {
return 0
}
r, err := strconv.ParseFloat(strconv.FormatFloat(v, 'g', 12, 64), 64)
if err != nil {
return v
}
return r
}
type explicitContextFields struct {
input, output, cacheWrite, cacheRead bool
}
// explicitContextPricingFields 判断各项是否被分组卡/渠道定价(含命中区间)显式配置。
// 显式配置为 0 时计费按 $0 收,展示应为 $0 而非“无价”。
func explicitContextPricingFields(resolved *ResolvedPricing, contextTokens int) explicitContextFields {
var out explicitContextFields
if resolved == nil || resolved.channelPricing == nil {
return out
}
cp := resolved.channelPricing
out.input = cp.InputPrice != nil
out.output = cp.OutputPrice != nil
out.cacheWrite = cp.CacheWritePrice != nil
out.cacheRead = cp.CacheReadPrice != nil
if iv := FindMatchingInterval(resolved.Intervals, contextTokens); iv != nil {
out.input = out.input || iv.InputPrice != nil
out.output = out.output || iv.OutputPrice != nil
out.cacheWrite = out.cacheWrite || iv.CacheWritePrice != nil
out.cacheRead = out.cacheRead || iv.CacheReadPrice != nil
}
return out
}
func contextPricePtr(v *float64, explicit bool) *float64 {
if v == nil {
return nil
}
if *v == 0 && !explicit {
return nil
}
return v
}
// applyIntervalContextLabels 把管理员在渠道区间上配置的 tier_label 带到对应档位。
func applyIntervalContextLabels(tiers []ContextPricingTier, intervals []PricingInterval) {
for i := range tiers {
for j := range intervals {
iv := &intervals[j]
if iv.TierLabel == "" || iv.MinTokens != tiers[i].MinTokens {
continue
}
if (iv.MaxTokens == nil) != (tiers[i].MaxTokens == nil) {
continue
}
if iv.MaxTokens != nil && *iv.MaxTokens != *tiers[i].MaxTokens {
continue
}
tiers[i].Label = iv.TierLabel
break
}
}
}
// mergeEqualContextTiers 合并相邻、四项单价相同且都没有管理员标签的段
// (倍率 ≤1 的目录、关闭阶梯等场景塌成单档)。
func mergeEqualContextTiers(tiers []ContextPricingTier) []ContextPricingTier {
if len(tiers) < 2 {
return tiers
}
merged := make([]ContextPricingTier, 0, len(tiers))
for _, t := range tiers {
if n := len(merged); n > 0 && merged[n-1].Label == "" && t.Label == "" && sameContextPrices(merged[n-1], t) {
merged[n-1].MaxTokens = t.MaxTokens
continue
}
merged = append(merged, t)
}
return merged
}
func sameContextPrices(a, b ContextPricingTier) bool {
return samePricePtr(a.Input, b.Input) && samePricePtr(a.Output, b.Output) &&
samePricePtr(a.CacheWrite, b.CacheWrite) && samePricePtr(a.CacheRead, b.CacheRead)
}
func samePricePtr(a, b *float64) bool {
if a == nil || b == nil {
return a == nil && b == nil
}
if *a == *b {
return true
}
scale := math.Max(math.Abs(*a), math.Abs(*b))
return math.Abs(*a-*b) <= scale*1e-9
}
// applyGeneratedContextLabels 给档位打标签:渠道区间沿用管理员配置的 tier_label;
// 目录阶梯/旧规则的两档按阈值生成(达到阈值即进高档时用 < / ≥)。
func applyGeneratedContextLabels(tiers []ContextPricingTier, plan contextBreakpointPlan) {
if len(tiers) < 2 {
return
}
if plan.thresholdBound > 0 {
label := formatContextTokenCount(plan.threshold)
lowPrefix, highPrefix := "≤", ">"
if plan.thresholdInclusive {
lowPrefix, highPrefix = "<", "≥"
}
for i := range tiers {
switch {
case tiers[i].MaxTokens != nil && *tiers[i].MaxTokens == plan.thresholdBound:
tiers[i].Label = lowPrefix + label
case tiers[i].MinTokens == plan.thresholdBound:
tiers[i].Label = highPrefix + label
}
}
}
}
// formatContextTokenCount 把 token 数格式化为 272K / 1M 等短标签。
func formatContextTokenCount(n int) string {
switch {
case n >= 1_000_000 && n%1_000_000 == 0:
return strconv.Itoa(n/1_000_000) + "M"
case n >= 1_000_000:
return trimFloatLabel(float64(n)/1_000_000) + "M"
case n >= 1_000:
return trimFloatLabel(float64(n)/1_000) + "K"
}
return strconv.Itoa(n)
}
func trimFloatLabel(v float64) string {
return strconv.FormatFloat(math.Round(v*100)/100, 'f', -1, 64)
}
@@ -0,0 +1,486 @@
//go:build unit
package service
import (
"context"
"math"
"math/rand"
"testing"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/stretchr/testify/require"
)
// scheduleScenario 阶梯表测试场景:同一份场景既做断言也做与计费函数的对账。
type scheduleScenario struct {
name string
model string
platform string
groupPlatform string
group *Group
channel []ChannelModelPricing
catalog *PricingService
wantErr bool
wantNil bool
wantBasis ContextPricingBasis
check func(t *testing.T, s *ContextPricingSchedule)
}
func enabledGroup(platform string) *Group {
return &Group{ID: 100, Platform: platform, LongContextPricingEnabled: true}
}
func disabledGroup(platform string) *Group {
return &Group{ID: 100, Platform: platform, LongContextPricingEnabled: false}
}
func sonnetChannel(iv ...PricingInterval) []ChannelModelPricing {
return []ChannelModelPricing{{
Platform: PlatformAnthropic, Models: []string{"claude-sonnet-4"}, BillingMode: BillingModeToken,
InputPrice: testPtrFloat64(2e-6), Intervals: iv,
}}
}
func requireTier(t *testing.T, tier ContextPricingTier, min int, max *int, label string, input, output, cacheWrite, cacheRead *float64) {
t.Helper()
require.Equal(t, min, tier.MinTokens, "min")
if max == nil {
require.Nil(t, tier.MaxTokens, "max should be unbounded")
} else {
require.NotNil(t, tier.MaxTokens, "max")
require.Equal(t, *max, *tier.MaxTokens, "max")
}
require.Equal(t, label, tier.Label, "label")
requirePrice(t, input, tier.Input, "input")
requirePrice(t, output, tier.Output, "output")
requirePrice(t, cacheWrite, tier.CacheWrite, "cache_write")
requirePrice(t, cacheRead, tier.CacheRead, "cache_read")
}
func requirePrice(t *testing.T, want, got *float64, field string) {
t.Helper()
if want == nil {
require.Nil(t, got, field)
return
}
require.NotNil(t, got, field)
require.InDelta(t, *want, *got, 1e-15, field)
}
func scheduleScenarios() []scheduleScenario {
p := testPtrFloat64
return []scheduleScenario{
{
name: "官方阶梯 gpt-5.4 整单两档", model: "gpt-5.4", platform: PlatformOpenAI, groupPlatform: PlatformOpenAI,
group: enabledGroup(PlatformOpenAI), wantBasis: ContextPricingBasisWholeRequest,
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 2)
requireTier(t, s.Tiers[0], 0, intPtr(272000), "≤272K", p(2.5e-6), p(15e-6), p(2.5e-6), p(0.25e-6))
requireTier(t, s.Tiers[1], 272000, nil, ">272K", p(5e-6), p(22.5e-6), p(5e-6), p(0.5e-6))
},
},
{
name: "Grok 达到阈值即进高档", model: "grok-4.5", platform: PlatformGrok, groupPlatform: PlatformGrok,
group: enabledGroup(PlatformGrok), wantBasis: ContextPricingBasisWholeRequest,
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 2)
requireTier(t, s.Tiers[0], 0, intPtr(199999), "<200K", p(2e-6), p(6e-6), nil, p(0.3e-6))
requireTier(t, s.Tiers[1], 199999, nil, "≥200K", p(4e-6), p(12e-6), nil, p(0.6e-6))
},
},
{
name: "分组关闭阶梯只剩基础档", model: "gpt-5.4", platform: PlatformOpenAI, groupPlatform: PlatformOpenAI,
group: disabledGroup(PlatformOpenAI), wantBasis: ContextPricingBasisWholeRequest,
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 1)
requireTier(t, s.Tiers[0], 0, nil, "", p(2.5e-6), p(15e-6), p(2.5e-6), p(0.25e-6))
},
},
{
name: "官方参考价(无分组)带目录阶梯", model: "gpt-5.4", platform: "", groupPlatform: PlatformOpenAI,
group: nil, wantBasis: ContextPricingBasisWholeRequest,
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 2)
requireTier(t, s.Tiers[1], 272000, nil, ">272K", p(5e-6), p(22.5e-6), p(5e-6), p(0.5e-6))
},
},
{
name: "渠道倍率区间按渠道平价覆盖后的 base 折算", model: "claude-sonnet-4", platform: PlatformAnthropic, groupPlatform: PlatformAnthropic,
group: enabledGroup(PlatformAnthropic), wantBasis: ContextPricingBasisWholeRequest,
channel: sonnetChannel(
PricingInterval{MinTokens: 0, MaxTokens: intPtr(200000), InputMultiplier: p(1)},
PricingInterval{MinTokens: 200000, TierLabel: "long", InputMultiplier: p(2), OutputMultiplier: p(1.5), CacheWriteMultiplier: p(2), CacheReadMultiplier: p(2)},
),
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 2)
requireTier(t, s.Tiers[0], 0, intPtr(200000), "", p(2e-6), p(15e-6), p(3.75e-6), p(0.3e-6))
requireTier(t, s.Tiers[1], 200000, nil, "long", p(4e-6), p(22.5e-6), p(7.5e-6), p(0.6e-6))
},
},
{
name: "区间显式价优先于倍率", model: "claude-sonnet-4", platform: PlatformAnthropic, groupPlatform: PlatformAnthropic,
group: enabledGroup(PlatformAnthropic), wantBasis: ContextPricingBasisWholeRequest,
channel: sonnetChannel(
PricingInterval{MinTokens: 0, MaxTokens: intPtr(200000), InputMultiplier: p(1)},
PricingInterval{MinTokens: 200000, InputPrice: p(9e-6), InputMultiplier: p(2)},
),
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 2)
requirePrice(t, p(9e-6), s.Tiers[1].Input, "input")
},
},
{
name: "区间空洞按 base 补档且同价段合并", model: "claude-sonnet-4", platform: PlatformAnthropic, groupPlatform: PlatformAnthropic,
group: enabledGroup(PlatformAnthropic), wantBasis: ContextPricingBasisWholeRequest,
channel: sonnetChannel(
PricingInterval{MinTokens: 0, MaxTokens: intPtr(100000), InputMultiplier: p(0.5)},
PricingInterval{MinTokens: 200000, InputMultiplier: p(2)},
),
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 3)
requireTier(t, s.Tiers[0], 0, intPtr(100000), "", p(1e-6), p(15e-6), p(3.75e-6), p(0.3e-6))
requireTier(t, s.Tiers[1], 100000, intPtr(200000), "", p(2e-6), p(15e-6), p(3.75e-6), p(0.3e-6))
requireTier(t, s.Tiers[2], 200000, nil, "", p(4e-6), p(15e-6), p(3.75e-6), p(0.3e-6))
},
},
{
name: "首档同价空洞被合并", model: "claude-sonnet-4", platform: PlatformAnthropic, groupPlatform: PlatformAnthropic,
group: enabledGroup(PlatformAnthropic), wantBasis: ContextPricingBasisWholeRequest,
channel: sonnetChannel(
PricingInterval{MinTokens: 0, MaxTokens: intPtr(100000), InputMultiplier: p(1)},
PricingInterval{MinTokens: 200000, InputMultiplier: p(2)},
),
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 2)
requireTier(t, s.Tiers[0], 0, intPtr(200000), "", p(2e-6), p(15e-6), p(3.75e-6), p(0.3e-6))
requireTier(t, s.Tiers[1], 200000, nil, "", p(4e-6), p(15e-6), p(3.75e-6), p(0.3e-6))
},
},
{
name: "末档有上限时尾段回落 base", model: "claude-sonnet-4", platform: PlatformAnthropic, groupPlatform: PlatformAnthropic,
group: enabledGroup(PlatformAnthropic), wantBasis: ContextPricingBasisWholeRequest,
channel: sonnetChannel(
PricingInterval{MinTokens: 0, MaxTokens: intPtr(200000), InputMultiplier: p(1)},
PricingInterval{MinTokens: 200000, MaxTokens: intPtr(1000000), InputMultiplier: p(2)},
),
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 3)
requireTier(t, s.Tiers[1], 200000, intPtr(1000000), "", p(4e-6), p(15e-6), p(3.75e-6), p(0.3e-6))
requireTier(t, s.Tiers[2], 1000000, nil, "", p(2e-6), p(15e-6), p(3.75e-6), p(0.3e-6))
},
},
{
name: "分组关闭时渠道区间折叠到 1-token 档", model: "claude-sonnet-4", platform: PlatformAnthropic, groupPlatform: PlatformAnthropic,
group: disabledGroup(PlatformAnthropic), wantBasis: ContextPricingBasisWholeRequest,
channel: sonnetChannel(
PricingInterval{MinTokens: 0, MaxTokens: intPtr(200000), InputMultiplier: p(0.5)},
PricingInterval{MinTokens: 200000, InputMultiplier: p(2)},
),
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 1)
requireTier(t, s.Tiers[0], 0, nil, "", p(1e-6), p(15e-6), p(3.75e-6), p(0.3e-6))
},
},
{
name: "分组关闭且首档不从 0 起时平价为 base", model: "claude-sonnet-4", platform: PlatformAnthropic, groupPlatform: PlatformAnthropic,
group: disabledGroup(PlatformAnthropic), wantBasis: ContextPricingBasisWholeRequest,
channel: sonnetChannel(PricingInterval{MinTokens: 200000, InputMultiplier: p(2)}),
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 1)
requirePrice(t, p(2e-6), s.Tiers[0].Input, "input")
},
},
{
name: "分组 token 价卡整张替换渠道定价并剥区间", model: "claude-sonnet-4", platform: PlatformAnthropic, groupPlatform: PlatformAnthropic,
group: &Group{ID: 100, Platform: PlatformAnthropic, LongContextPricingEnabled: true, ModelPricing: []ChannelModelPricing{{
Models: []string{"claude-sonnet-*"}, BillingMode: BillingModeToken, InputPrice: p(1e-6),
Intervals: []PricingInterval{{MinTokens: 200000, InputMultiplier: p(5)}},
}}},
channel: sonnetChannel(PricingInterval{MinTokens: 200000, InputMultiplier: p(3)}),
wantBasis: ContextPricingBasisWholeRequest,
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 1)
requireTier(t, s.Tiers[0], 0, nil, "", p(1e-6), p(15e-6), p(3.75e-6), p(0.3e-6))
},
},
{
name: "分组价卡之上叠加官方阶梯", model: "gpt-5.4", platform: PlatformOpenAI, groupPlatform: PlatformOpenAI,
group: &Group{ID: 100, Platform: PlatformOpenAI, LongContextPricingEnabled: true, ModelPricing: []ChannelModelPricing{{
Models: []string{"gpt-5.4"}, BillingMode: BillingModeToken, InputPrice: p(1e-6),
}}},
wantBasis: ContextPricingBasisWholeRequest,
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 2)
requirePrice(t, p(1e-6), s.Tiers[0].Input, "input")
requirePrice(t, p(2e-6), s.Tiers[1].Input, "input")
requirePrice(t, p(22.5e-6), s.Tiers[1].Output, "output")
},
},
{
name: "Gemini 旧规则按超出部分计价", model: "gemini-2.5-pro", platform: PlatformGemini, groupPlatform: PlatformGemini,
group: enabledGroup(PlatformGemini), catalog: geminiCatalogStub(), wantBasis: ContextPricingBasisMarginal,
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 2)
requireTier(t, s.Tiers[0], 0, intPtr(200000), "≤200K", p(1.25e-6), p(10e-6), nil, p(0.3125e-6))
requireTier(t, s.Tiers[1], 200000, nil, ">200K", p(2.5e-6), p(10e-6), nil, p(0.625e-6))
},
},
{
name: "Gemini 分组关闭时不用旧规则", model: "gemini-2.5-pro", platform: PlatformGemini, groupPlatform: PlatformGemini,
group: disabledGroup(PlatformGemini), catalog: geminiCatalogStub(), wantBasis: ContextPricingBasisWholeRequest,
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 1)
},
},
{
name: "Gemini 有渠道定价时旧规则让位", model: "gemini-2.5-pro", platform: PlatformGemini, groupPlatform: PlatformGemini,
group: enabledGroup(PlatformGemini), catalog: geminiCatalogStub(),
channel: []ChannelModelPricing{{
Platform: PlatformGemini, Models: []string{"gemini-2.5-pro"}, BillingMode: BillingModeToken, InputPrice: p(3e-6),
}},
wantBasis: ContextPricingBasisWholeRequest,
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 1)
requirePrice(t, p(3e-6), s.Tiers[0].Input, "input")
},
},
{
name: "Gemini 官方参考不套用站内旧规则", model: "gemini-2.5-pro", platform: "", groupPlatform: PlatformGemini,
group: nil, catalog: geminiCatalogStub(), wantBasis: ContextPricingBasisWholeRequest,
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 1)
},
},
{
name: "composite 分组按模型平台取渠道定价并叠加官方阶梯", model: "gpt-5.4", platform: PlatformOpenAI, groupPlatform: PlatformComposite,
group: enabledGroup(PlatformComposite),
channel: []ChannelModelPricing{{
Platform: PlatformOpenAI, Models: []string{"gpt-5.4"}, BillingMode: BillingModeToken, InputPrice: p(1e-6),
}},
wantBasis: ContextPricingBasisWholeRequest,
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 2)
requirePrice(t, p(1e-6), s.Tiers[0].Input, "input")
requirePrice(t, p(2e-6), s.Tiers[1].Input, "input")
},
},
{
name: "渠道显式 cache_write=0 保留为 0", model: "claude-sonnet-4", platform: PlatformAnthropic, groupPlatform: PlatformAnthropic,
group: enabledGroup(PlatformAnthropic),
channel: []ChannelModelPricing{{
Platform: PlatformAnthropic, Models: []string{"claude-sonnet-4"}, BillingMode: BillingModeToken,
InputPrice: p(2e-6), CacheWritePrice: p(0),
}},
wantBasis: ContextPricingBasisWholeRequest,
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 1)
require.NotNil(t, s.Tiers[0].CacheWrite)
require.Zero(t, *s.Tiers[0].CacheWrite)
},
},
{
name: "gpt-5.6 缺 cache_write 时按策略补 1.25 倍并带阶梯", model: "gpt-5.6-sol", platform: PlatformOpenAI, groupPlatform: PlatformOpenAI,
group: enabledGroup(PlatformOpenAI),
catalog: newStubPricingServiceFromMap(map[string]*LiteLLMModelPricing{
"gpt-5.6-sol": {Mode: "chat", InputCostPerToken: 5e-6, OutputCostPerToken: 30e-6, CacheReadInputTokenCost: 0.5e-6},
}),
wantBasis: ContextPricingBasisWholeRequest,
check: func(t *testing.T, s *ContextPricingSchedule) {
require.Len(t, s.Tiers, 2)
requireTier(t, s.Tiers[0], 0, intPtr(272000), "≤272K", p(5e-6), p(30e-6), p(6.25e-6), p(0.5e-6))
requireTier(t, s.Tiers[1], 272000, nil, ">272K", p(10e-6), p(45e-6), p(12.5e-6), p(1e-6))
},
},
{
name: "图片模式返回 nil", model: "gpt-image-2", platform: PlatformOpenAI, groupPlatform: PlatformOpenAI,
group: enabledGroup(PlatformOpenAI),
channel: []ChannelModelPricing{{
Platform: PlatformOpenAI, Models: []string{"gpt-image-2"}, BillingMode: BillingModeImage, PerRequestPrice: p(0.04),
}},
wantNil: true,
},
{
name: "无任何定价来源报错", model: "unknown-model-xyz", platform: PlatformOpenAI, groupPlatform: PlatformOpenAI,
group: enabledGroup(PlatformOpenAI), wantErr: true,
},
}
}
func newScheduleTestEnv(t *testing.T, sc scheduleScenario) (*BillingService, *ModelPricingResolver) {
t.Helper()
return newTokenCostTestEnv(t, sc.groupPlatform, sc.channel, sc.catalog)
}
func TestResolveContextPricingSchedule_Scenarios(t *testing.T) {
for _, sc := range scheduleScenarios() {
t.Run(sc.name, func(t *testing.T) {
bs, resolver := newScheduleTestEnv(t, sc)
sched, err := bs.ResolveContextPricingSchedule(context.Background(), resolver, ContextPricingScheduleInput{
Model: sc.model, Group: sc.group, Platform: sc.platform,
})
if sc.wantErr {
require.Error(t, err)
require.Nil(t, sched)
return
}
require.NoError(t, err)
if sc.wantNil {
require.Nil(t, sched)
return
}
require.NotNil(t, sched)
require.Equal(t, sc.wantBasis, sched.Basis)
require.NotEmpty(t, sched.Tiers)
require.Zero(t, sched.Tiers[0].MinTokens, "首档从 0 起")
for i := 1; i < len(sched.Tiers); i++ {
require.NotNil(t, sched.Tiers[i-1].MaxTokens)
require.Equal(t, *sched.Tiers[i-1].MaxTokens, sched.Tiers[i].MinTokens, "档位连续")
}
if sc.check != nil {
sc.check(t, sched)
}
})
}
}
func TestResolveContextPricingSchedule_NilResolver(t *testing.T) {
bs := NewBillingService(&config.Config{}, nil)
sched, err := bs.ResolveContextPricingSchedule(context.Background(), nil, ContextPricingScheduleInput{Model: "gpt-5.4"})
require.Error(t, err)
require.Nil(t, sched)
}
// --- 对账:阶梯表推算的费用必须等于真实计费函数 ---
type tokenKind int
const (
kindInput tokenKind = iota
kindOutput
kindCacheRead
kindCacheWrite
)
func tierPrice(tier ContextPricingTier, kind tokenKind) float64 {
var p *float64
switch kind {
case kindInput:
p = tier.Input
case kindOutput:
p = tier.Output
case kindCacheRead:
p = tier.CacheRead
case kindCacheWrite:
p = tier.CacheWrite
}
if p == nil {
return 0
}
return *p
}
func tierAt(tiers []ContextPricingTier, contextTokens int) ContextPricingTier {
for _, tier := range tiers {
if contextTokens > tier.MinTokens && (tier.MaxTokens == nil || contextTokens <= *tier.MaxTokens) {
return tier
}
}
return tiers[len(tiers)-1]
}
// expectedCostFromSchedule 按阶梯表推算 contextTokens 个某类 token 的费用:
// 整单基准取所在档单价 × 全量;边际基准逐段累加。
func expectedCostFromSchedule(s *ContextPricingSchedule, kind tokenKind, contextTokens int) float64 {
if s.Basis == ContextPricingBasisMarginal {
total := 0.0
for _, tier := range s.Tiers {
if contextTokens <= tier.MinTokens {
break
}
upper := contextTokens
if tier.MaxTokens != nil && *tier.MaxTokens < upper {
upper = *tier.MaxTokens
}
total += float64(upper-tier.MinTokens) * tierPrice(tier, kind)
}
return total
}
return float64(contextTokens) * tierPrice(tierAt(s.Tiers, contextTokens), kind)
}
func TestResolveContextPricingSchedule_ParityWithBilling(t *testing.T) {
const outputProbe = 1000
rng := rand.New(rand.NewSource(20260823))
for _, sc := range scheduleScenarios() {
if sc.wantErr || sc.wantNil {
continue
}
t.Run(sc.name, func(t *testing.T) {
bs, resolver := newScheduleTestEnv(t, sc)
ctx := context.Background()
if sc.platform != "" {
ctx = WithResolvedTargetPlatform(ctx, sc.platform)
}
sched, err := bs.ResolveContextPricingSchedule(ctx, resolver, ContextPricingScheduleInput{
Model: sc.model, Group: sc.group, Platform: sc.platform,
})
require.NoError(t, err)
require.NotNil(t, sched)
// 消费方视角重建同一份计费请求(与网关一致)。
pricingInput := PricingInput{Model: sc.model, Group: sc.group}
if sc.group != nil {
gid := sc.group.ID
pricingInput.GroupID = &gid
}
resolved := resolver.Resolve(ctx, pricingInput)
var legacy *LegacyLongContextRule
if sc.group != nil {
legacy = bs.LegacyLongContextRule(sc.platform)
}
if !legacyLongContextApplies(resolved, sc.group, legacy) {
legacy = nil
}
cost := func(tokens UsageTokens) float64 {
bd, err := bs.CalculateTokenCostForRequest(TokenCostRequest{
Ctx: ctx, Model: sc.model, Group: sc.group, Tokens: tokens, RateMultiplier: 1,
Resolver: resolver, Resolved: resolved, LegacyLongContext: legacy,
})
require.NoError(t, err)
return bd.ActualCost
}
probes := []int{1, 2, 999, 1000, 1001, 10_000_000}
for _, tier := range sched.Tiers {
for _, b := range []*int{tier.MaxTokens} {
if b == nil {
continue
}
probes = append(probes, *b-1, *b, *b+1)
}
}
for i := 0; i < 8; i++ {
probes = append(probes, 1+rng.Intn(3_000_000))
}
for _, c := range probes {
if c <= 0 {
continue
}
assertCostClose(t, expectedCostFromSchedule(sched, kindInput, c), cost(UsageTokens{InputTokens: c}), "input@%d", c)
assertCostClose(t, expectedCostFromSchedule(sched, kindCacheRead, c), cost(UsageTokens{CacheReadTokens: c}), "cache_read@%d", c)
assertCostClose(t, expectedCostFromSchedule(sched, kindCacheWrite, c), cost(UsageTokens{CacheCreationTokens: c}), "cache_write@%d", c)
wantOutput := float64(outputProbe) * tierPrice(tierAt(sched.Tiers, c), kindOutput)
gotOutput := cost(UsageTokens{InputTokens: c, OutputTokens: outputProbe}) - cost(UsageTokens{InputTokens: c})
assertCostClose(t, wantOutput, gotOutput, "output@%d", c)
}
})
}
}
func assertCostClose(t *testing.T, want, got float64, format string, args ...any) {
t.Helper()
tolerance := 1e-12 + 1e-9*math.Abs(want)
require.InDeltaf(t, want, got, tolerance, format, args...)
}
@@ -0,0 +1,103 @@
package service
import (
"context"
"time"
)
// LegacyLongContextRule 平台级"超出阈值部分按倍率计费"的旧规则。
//
// 语义为边际计费:仅 input/cache_read 中超过 Threshold 的部分乘 Multiplier,
// output 与 cache_write 不受影响(见 CalculateCostWithLongContext)。
// 只有 Gemini 原生 /v1beta 入口在该分组对该模型没有分组/渠道定价时才使用;
// 规则常量由 BillingService 统一持有,网关与模型广场都从这里读取。
type LegacyLongContextRule struct {
Threshold int
Multiplier float64
}
const (
geminiLegacyLongContextThreshold = 200000
geminiLegacyLongContextMultiplier = 2.0
)
// LegacyLongContextRule 返回平台的旧长上下文规则;无规则的平台返回 nil。
func (s *BillingService) LegacyLongContextRule(platform string) *LegacyLongContextRule {
if platform == PlatformGemini {
return &LegacyLongContextRule{
Threshold: geminiLegacyLongContextThreshold,
Multiplier: geminiLegacyLongContextMultiplier,
}
}
return nil
}
// TokenCostRequest 通用网关 token 计费请求。
type TokenCostRequest struct {
Ctx context.Context
Model string
Group *Group
Tokens UsageTokens
RateMultiplier float64
PricingAt time.Time
ServiceTier string
Resolver *ModelPricingResolver
// Resolved 为调用方预先解析的定价(Resolver.Resolve 的结果),nil 表示未解析。
Resolved *ResolvedPricing
// LegacyLongContext 入口携带的旧长上下文规则,nil 表示该入口不使用。
LegacyLongContext *LegacyLongContextRule
}
// legacyLongContextApplies 判定请求是否走旧长上下文规则:
// 分组/渠道显式定价优先;否则在规则存在且分组长上下文开关开启时生效。
func legacyLongContextApplies(resolved *ResolvedPricing, group *Group, rule *LegacyLongContextRule) bool {
if rule == nil || rule.Threshold <= 0 {
return false
}
if resolved != nil && (resolved.Source == PricingSourceGroup || resolved.Source == PricingSourceChannel) {
return false
}
return group == nil || group.LongContextPricingEnabled
}
// CalculateTokenCostForRequest 按通用网关的路径选择计算 token 费用:
// 1. 分组/渠道显式定价 → 统一计费(区间、分组卡、目录阶梯均在其中);
// 2. 否则入口带旧长上下文规则且分组开关开启 → 旧边际计费;
// 3. 否则有解析器与分组 → 统一计费(内置目录定价);
// 4. 否则按模型目录直接计费。
//
// 模型广场的阶梯表查询与网关使用同一入口,保证展示与扣费同源。
func (s *BillingService) CalculateTokenCostForRequest(req TokenCostRequest) (*CostBreakdown, error) {
resolved := req.Resolved
if resolved != nil && (resolved.Source == PricingSourceGroup || resolved.Source == PricingSourceChannel) {
return s.CalculateCostUnified(s.tokenCostInput(req, resolved))
}
if legacyLongContextApplies(resolved, req.Group, req.LegacyLongContext) {
return s.CalculateCostWithLongContext(req.Model, req.Tokens, req.RateMultiplier,
req.LegacyLongContext.Threshold, req.LegacyLongContext.Multiplier)
}
if req.Resolver != nil && req.Group != nil {
return s.CalculateCostUnified(s.tokenCostInput(req, resolved))
}
return s.CalculateCost(req.Model, req.Tokens, req.RateMultiplier)
}
func (s *BillingService) tokenCostInput(req TokenCostRequest, resolved *ResolvedPricing) CostInput {
input := CostInput{
Ctx: req.Ctx,
Model: req.Model,
Group: req.Group,
Tokens: req.Tokens,
RequestCount: 1,
RateMultiplier: req.RateMultiplier,
PricingAt: req.PricingAt,
ServiceTier: req.ServiceTier,
Resolver: req.Resolver,
Resolved: resolved,
}
if req.Group != nil {
gid := req.Group.ID
input.GroupID = &gid
}
return input
}
@@ -0,0 +1,147 @@
//go:build unit
package service
import (
"context"
"testing"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/stretchr/testify/require"
)
// newTokenCostTestEnv 构造带渠道定价的计费环境:group 100 挂一个渠道,定价由 pricing 指定。
func newTokenCostTestEnv(t *testing.T, groupPlatform string, pricing []ChannelModelPricing, catalog *PricingService) (*BillingService, *ModelPricingResolver) {
t.Helper()
repo := &mockChannelRepository{
listAllFn: func(_ context.Context) ([]Channel, error) {
return []Channel{{
ID: 1, Name: "ch", Status: StatusActive, GroupIDs: []int64{100}, ModelPricing: pricing,
}}, nil
},
getGroupPlatformsFn: func(_ context.Context, _ []int64) (map[int64]string, error) {
return map[int64]string{100: groupPlatform}, nil
},
}
cs := NewChannelService(repo, nil, nil, nil)
bs := NewBillingService(&config.Config{}, catalog)
return bs, NewModelPricingResolver(cs, bs)
}
func geminiCatalogStub() *PricingService {
return newStubPricingServiceFromMap(map[string]*LiteLLMModelPricing{
"gemini-2.5-pro": {
Mode: "chat",
InputCostPerToken: 1.25e-6,
OutputCostPerToken: 10e-6,
CacheReadInputTokenCost: 0.3125e-6,
},
})
}
func TestLegacyLongContextRule_OnlyGemini(t *testing.T) {
bs := NewBillingService(&config.Config{}, nil)
rule := bs.LegacyLongContextRule(PlatformGemini)
require.NotNil(t, rule)
require.Equal(t, 200000, rule.Threshold)
require.InDelta(t, 2.0, rule.Multiplier, 1e-12)
for _, platform := range []string{PlatformOpenAI, PlatformAnthropic, PlatformAntigravity, PlatformComposite, ""} {
require.Nil(t, bs.LegacyLongContextRule(platform), platform)
}
}
func TestCalculateTokenCostForRequest_ChannelPricingWinsOverLegacyRule(t *testing.T) {
bs, resolver := newTokenCostTestEnv(t, PlatformGemini, []ChannelModelPricing{{
Platform: PlatformGemini, Models: []string{"gemini-2.5-pro"}, BillingMode: BillingModeToken,
InputPrice: testPtrFloat64(10e-6), OutputPrice: testPtrFloat64(40e-6),
}}, geminiCatalogStub())
group := &Group{ID: 100, Platform: PlatformGemini, LongContextPricingEnabled: true}
gid := group.ID
resolved := resolver.Resolve(context.Background(), PricingInput{Model: "gemini-2.5-pro", GroupID: &gid, Group: group})
require.Equal(t, PricingSourceChannel, resolved.Source)
tokens := UsageTokens{InputTokens: 300000, OutputTokens: 1000}
got, err := bs.CalculateTokenCostForRequest(TokenCostRequest{
Ctx: context.Background(), Model: "gemini-2.5-pro", Group: group, Tokens: tokens, RateMultiplier: 1,
Resolver: resolver, Resolved: resolved, LegacyLongContext: bs.LegacyLongContextRule(PlatformGemini),
})
require.NoError(t, err)
want, err := bs.CalculateCostUnified(CostInput{
Ctx: context.Background(), Model: "gemini-2.5-pro", GroupID: &gid, Group: group, Tokens: tokens,
RequestCount: 1, RateMultiplier: 1, Resolver: resolver, Resolved: resolved,
})
require.NoError(t, err)
require.Equal(t, want, got)
// 渠道平价 10e-6 × 300K,旧规则未叠加
require.InDelta(t, 3.0, got.InputCost, 1e-9)
}
func TestCalculateTokenCostForRequest_LegacyRuleFollowsGroupToggle(t *testing.T) {
bs, resolver := newTokenCostTestEnv(t, PlatformGemini, nil, geminiCatalogStub())
tokens := UsageTokens{InputTokens: 300000, OutputTokens: 1000}
rule := bs.LegacyLongContextRule(PlatformGemini)
for _, enabled := range []bool{true, false} {
group := &Group{ID: 100, Platform: PlatformGemini, LongContextPricingEnabled: enabled}
gid := group.ID
resolved := resolver.Resolve(context.Background(), PricingInput{Model: "gemini-2.5-pro", GroupID: &gid, Group: group})
require.Equal(t, PricingSourceLiteLLM, resolved.Source)
got, err := bs.CalculateTokenCostForRequest(TokenCostRequest{
Ctx: context.Background(), Model: "gemini-2.5-pro", Group: group, Tokens: tokens, RateMultiplier: 1,
Resolver: resolver, Resolved: resolved, LegacyLongContext: rule,
})
require.NoError(t, err)
if enabled {
want, err := bs.CalculateCostWithLongContext("gemini-2.5-pro", tokens, 1, rule.Threshold, rule.Multiplier)
require.NoError(t, err)
require.Equal(t, want, got)
// 输入 200K × 1.25e-6 + 超出 100K × 1.25e-6 × 2 = 0.5;输出 1000 × 10e-6 = 0.01。
// 旧路径的加倍只体现在 ActualCost(分项 InputCost 不含倍率),探针也据此取值。
require.InDelta(t, 0.51, got.ActualCost, 1e-9)
require.True(t, got.LongContextBillingApplied)
} else {
want, err := bs.CalculateCostUnified(CostInput{
Ctx: context.Background(), Model: "gemini-2.5-pro", GroupID: &gid, Group: group, Tokens: tokens,
RequestCount: 1, RateMultiplier: 1, Resolver: resolver, Resolved: resolved,
})
require.NoError(t, err)
require.Equal(t, want, got)
require.InDelta(t, 0.385, got.ActualCost, 1e-9)
require.False(t, got.LongContextBillingApplied)
}
}
}
func TestCalculateTokenCostForRequest_BuiltInPricingUsesUnifiedPath(t *testing.T) {
bs, resolver := newTokenCostTestEnv(t, PlatformOpenAI, nil, nil)
group := &Group{ID: 100, Platform: PlatformOpenAI, LongContextPricingEnabled: true}
gid := group.ID
tokens := UsageTokens{InputTokens: 300000, OutputTokens: 1000}
resolved := resolver.Resolve(context.Background(), PricingInput{Model: "gpt-5.4", GroupID: &gid, Group: group})
got, err := bs.CalculateTokenCostForRequest(TokenCostRequest{
Ctx: context.Background(), Model: "gpt-5.4", Group: group, Tokens: tokens, RateMultiplier: 1,
Resolver: resolver, Resolved: resolved,
})
require.NoError(t, err)
want, err := bs.CalculateCostUnified(CostInput{
Ctx: context.Background(), Model: "gpt-5.4", GroupID: &gid, Group: group, Tokens: tokens,
RequestCount: 1, RateMultiplier: 1, Resolver: resolver,
})
require.NoError(t, err)
require.Equal(t, want, got)
// 目录阶梯:超 272K 整单输入 ×2
require.InDelta(t, 300000*2.5e-6*2, got.InputCost, 1e-9)
require.True(t, got.LongContextBillingApplied)
}
func TestCalculateTokenCostForRequest_NoResolverFallsBackToCatalog(t *testing.T) {
bs := NewBillingService(&config.Config{}, nil)
tokens := UsageTokens{InputTokens: 1000, OutputTokens: 10}
got, err := bs.CalculateTokenCostForRequest(TokenCostRequest{Model: "gpt-5.4", Tokens: tokens, RateMultiplier: 1})
require.NoError(t, err)
want, err := bs.CalculateCost("gpt-5.4", tokens, 1)
require.NoError(t, err)
require.Equal(t, want, got)
}
@@ -1122,7 +1122,8 @@ func (s *GatewayService) calculateImageCost(
return s.billingService.CalculateImageCost(billingModel, sizeTier, result.ImageCount, groupConfig, multiplier)
}
// calculateTokenCost 计算 Token 计费:根据 opts 决定走普通/长上下文/渠道统一计费。
// calculateTokenCost 计算 Token 计费:路径选择(分组/渠道定价 → 旧长上下文规则 → 内置定价)
// 统一交给 BillingService.CalculateTokenCostForRequest,与模型广场的阶梯表查询同源。
func (s *GatewayService) calculateTokenCost(
ctx context.Context,
result *ForwardResult,
@@ -1142,39 +1143,28 @@ func (s *GatewayService) calculateTokenCost(
ImageOutputTokens: result.Usage.ImageOutputTokens,
}
var cost *CostBreakdown
var err error
// Explicit group/channel pricing wins. Built-in pricing also uses the unified
// resolver so the group long-context toggle can veto model-native tiers.
if resolved := s.resolveChannelPricing(ctx, billingModel, apiKey); resolved != nil {
var resolved *ResolvedPricing
if s.resolver != nil && apiKey.Group != nil {
gid := apiKey.Group.ID
cost, err = s.billingService.CalculateCostUnified(CostInput{
Ctx: ctx,
Model: billingModel,
GroupID: &gid,
Group: apiKey.Group,
Tokens: tokens,
RequestCount: 1,
RateMultiplier: multiplier,
PricingAt: pricingAt,
ServiceTier: optionalStringValue(result.ServiceTier),
Resolver: s.resolver,
Resolved: resolved,
})
} else if opts.LongContextThreshold > 0 && (apiKey.Group == nil || apiKey.Group.LongContextPricingEnabled) {
// 长上下文双倍计费(如 Gemini 200K 阈值)
cost, err = s.billingService.CalculateCostWithLongContext(billingModel, tokens, multiplier, opts.LongContextThreshold, opts.LongContextMultiplier)
} else if s.resolver != nil && apiKey.Group != nil {
gid := apiKey.Group.ID
cost, err = s.billingService.CalculateCostUnified(CostInput{
Ctx: ctx, Model: billingModel, GroupID: &gid, Group: apiKey.Group,
Tokens: tokens, RequestCount: 1, RateMultiplier: multiplier, PricingAt: pricingAt,
ServiceTier: optionalStringValue(result.ServiceTier), Resolver: s.resolver,
})
} else {
cost, err = s.billingService.CalculateCost(billingModel, tokens, multiplier)
resolved = s.resolver.Resolve(ctx, PricingInput{Model: billingModel, GroupID: &gid, Group: apiKey.Group})
}
var legacy *LegacyLongContextRule
if opts.LongContextThreshold > 0 {
legacy = &LegacyLongContextRule{Threshold: opts.LongContextThreshold, Multiplier: opts.LongContextMultiplier}
}
cost, err := s.billingService.CalculateTokenCostForRequest(TokenCostRequest{
Ctx: ctx,
Model: billingModel,
Group: apiKey.Group,
Tokens: tokens,
RateMultiplier: multiplier,
PricingAt: pricingAt,
ServiceTier: optionalStringValue(result.ServiceTier),
Resolver: s.resolver,
Resolved: resolved,
LegacyLongContext: legacy,
})
if err != nil {
logger.LegacyPrintf("service.gateway", "Calculate cost failed: %v", err)
return &CostBreakdown{ActualCost: 0}
@@ -1182,6 +1172,14 @@ func (s *GatewayService) calculateTokenCost(
return cost
}
// LegacyLongContextRule 透传 BillingService 的平台旧长上下文规则,供入口 handler 取用。
func (s *GatewayService) LegacyLongContextRule(platform string) *LegacyLongContextRule {
if s == nil || s.billingService == nil {
return nil
}
return s.billingService.LegacyLongContextRule(platform)
}
// buildRecordUsageLog 构建使用日志并设置计费模式。
func (s *GatewayService) buildRecordUsageLog(
ctx context.Context,