mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 17:08:33 +08:00
计费:统一 token 计费路径选择并提供上下文阶梯单价表查询
- BillingService.CalculateTokenCostForRequest 承接网关的路径选择 (分组/渠道定价 → 平台旧长上下文规则 → 内置目录),网关改为调用该入口 - Gemini /v1beta 的 200K 边际翻倍常量从 handler 移入 BillingService.LegacyLongContextRule,入口只声明适用 - 新增 ResolveContextPricingSchedule:沿用 Resolve 解析链收集断点 (渠道区间边界、目录阶梯阈值、旧规则阈值),每档单价由真实计费函数 探针差商得出,倍率/策略变更无需同步;附阶梯表 vs 计费函数的对账测试
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user