Merge pull request #6109 from feeeei/main

feat(model-plaza): 模型广场增加长上下文阶梯计价显示 & 分时段计价显示
This commit is contained in:
Wesley Liddick
2026-08-24 14:09:56 +08:00
committed by GitHub
23 changed files with 2787 additions and 405 deletions
+2 -1
View File
@@ -304,7 +304,8 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
handlerPaymentHandler := handler.NewPaymentHandler(paymentService, paymentConfigService)
paymentWebhookHandler := handler.NewPaymentWebhookHandler(paymentService, registry)
availableChannelHandler := handler.NewAvailableChannelHandler(channelService, apiKeyService, settingService)
modelPlazaHandler := handler.NewModelPlazaHandler(channelService, apiKeyService, settingService)
modelPlazaService := service.NewModelPlazaService(channelRepository, groupRepository, pricingService, billingService, modelPricingResolver)
modelPlazaHandler := handler.NewModelPlazaHandler(modelPlazaService, apiKeyService, settingService)
imageTaskStore := repository.NewImageTaskStore(redisClient)
imageTaskService := service.ProvideImageTaskService(imageTaskStore, imageStorageSettingService)
asyncImageHandler := handler.NewAsyncImageHandler(imageTaskService, openAIGatewayHandler)
@@ -284,13 +284,13 @@ func toUserSupportedModels(
return out
}
// toUserPricing 将 service 层定价转换为用户 DTO;入参为 nil 时返回 nil。
func toUserPricing(p *service.ChannelModelPricing) *userSupportedModelPricing {
if p == nil {
// toUserPricingIntervals 将定价区间转换为用户 DTO 白名单形态;nil 入参返回 nil(JSON omitempty 可省略)。
func toUserPricingIntervals(src []service.PricingInterval) []userPricingIntervalDTO {
if src == nil {
return nil
}
intervals := make([]userPricingIntervalDTO, 0, len(p.Intervals))
for _, iv := range p.Intervals {
intervals := make([]userPricingIntervalDTO, 0, len(src))
for _, iv := range src {
intervals = append(intervals, userPricingIntervalDTO{
MinTokens: iv.MinTokens,
MaxTokens: iv.MaxTokens,
@@ -302,6 +302,19 @@ func toUserPricing(p *service.ChannelModelPricing) *userSupportedModelPricing {
PerRequestPrice: iv.PerRequestPrice,
})
}
return intervals
}
// toUserPricing 将 service 层定价转换为用户 DTO;入参为 nil 时返回 nil。
func toUserPricing(p *service.ChannelModelPricing) *userSupportedModelPricing {
if p == nil {
return nil
}
intervals := toUserPricingIntervals(p.Intervals)
if intervals == nil {
// 用户侧定价的 intervals 固定输出数组(空配置为 []),保持既有契约。
intervals = []userPricingIntervalDTO{}
}
billingMode := string(p.BillingMode)
if billingMode == "" {
billingMode = string(service.BillingModeToken)
@@ -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,
+70 -27
View File
@@ -17,39 +17,60 @@ import (
// - 匿名:仅非专属分组(订阅型照常展示);
// - 登录:非专属分组 + user_allowed_groups 授权的专属分组(不检查订阅有效性)。
type ModelPlazaHandler struct {
channelService *service.ChannelService
plazaService *service.ModelPlazaService
apiKeyService *service.APIKeyService
settingService *service.SettingService
}
// NewModelPlazaHandler 创建模型广场 handler。
func NewModelPlazaHandler(
channelService *service.ChannelService,
plazaService *service.ModelPlazaService,
apiKeyService *service.APIKeyService,
settingService *service.SettingService,
) *ModelPlazaHandler {
return &ModelPlazaHandler{
channelService: channelService,
plazaService: plazaService,
apiKeyService: apiKeyService,
settingService: settingService,
}
}
// modelPlazaOfficialPricing LiteLLM 官方参考价(USD per token)。
// modelPlazaOfficialPricing 官方参考价(USD per token,与计费目录同源)。
type modelPlazaOfficialPricing struct {
InputPrice *float64 `json:"input_price"`
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheWrite1hPrice *float64 `json:"cache_write_1h_price,omitempty"`
CacheReadPrice *float64 `json:"cache_read_price"`
// Intervals 官方长上下文阶梯,仅多档模型给出。
Intervals []userPricingIntervalDTO `json:"intervals,omitempty"`
}
// modelPlazaModel 广场模型条目:渠道定价(白名单形态)+ 官方参考价。
// modelPlazaTimePricingPeriod 分时倍率时段(配置时区当天 [start, end))。
type modelPlazaTimePricingPeriod struct {
StartTime string `json:"start_time"`
EndTime string `json:"end_time"`
Multiplier float64 `json:"multiplier"`
}
// modelPlazaTimePricing 计费会生效的分时倍率(仅倍率 ≠ 1 的时段)。
// WeekdaysOnly 为 true 时时段仅周一至周五生效,周末整天按标准价计费。
type modelPlazaTimePricing struct {
Timezone string `json:"timezone"`
WeekdaysOnly bool `json:"weekdays_only,omitempty"`
Periods []modelPlazaTimePricingPeriod `json:"periods"`
}
// modelPlazaModel 广场模型条目:实收口径展示定价(白名单形态)+ 官方参考价。
type modelPlazaModel struct {
Name string `json:"name"`
Platform string `json:"platform"`
Pricing *userSupportedModelPricing `json:"pricing"`
OfficialPricing *modelPlazaOfficialPricing `json:"official_pricing"`
// LongContextBasis 多档时的计价基准:"whole_request"(整单按档)| "marginal"(仅超出部分)。
LongContextBasis string `json:"long_context_basis,omitempty"`
// TimePricing 分时倍率时段,落在时段内的请求整单乘倍率;无分时省略。
TimePricing *modelPlazaTimePricing `json:"time_pricing,omitempty"`
}
// modelPlazaGroup 广场分组条目(白名单字段)。
@@ -68,9 +89,11 @@ type modelPlazaGroup struct {
IsExclusive bool `json:"is_exclusive"`
// 生图独立倍率:为 true 时图片计费模型的实付倍率取 ImageRateMultiplier,
// 不取分组/用户专属倍率。
ImageRateIndependent bool `json:"image_rate_independent"`
ImageRateMultiplier float64 `json:"image_rate_multiplier"`
Models []modelPlazaModel `json:"models"`
ImageRateIndependent bool `json:"image_rate_independent"`
ImageRateMultiplier float64 `json:"image_rate_multiplier"`
// 分组是否启用长上下文阶梯计费;关闭时模型实付列只展示最低档/基础价。
LongContextPricingEnabled bool `json:"long_context_pricing_enabled"`
Models []modelPlazaModel `json:"models"`
}
// modelPlazaResponse 广场页响应。
@@ -98,7 +121,7 @@ func (h *ModelPlazaHandler) Get(c *gin.Context) {
return
}
groups, err := h.channelService.ListPlazaGroups(c.Request.Context())
groups, err := h.plazaService.ListGroups(c.Request.Context())
if err != nil {
response.ErrorFrom(c, err)
return
@@ -161,27 +184,30 @@ func toModelPlazaGroupDTO(g *service.PlazaGroup, userRates map[int64]float64) mo
for i := range g.Models {
m := &g.Models[i]
models = append(models, modelPlazaModel{
Name: m.Name,
Platform: m.Platform,
Pricing: toUserPricing(m.Pricing),
OfficialPricing: toModelPlazaOfficialPricing(m.OfficialPricing),
Name: m.Name,
Platform: m.Platform,
Pricing: toUserPricing(m.Pricing),
OfficialPricing: toModelPlazaOfficialPricing(m.OfficialPricing),
LongContextBasis: string(m.LongContextBasis),
TimePricing: toModelPlazaTimePricing(m.TimePricing),
})
}
dto := modelPlazaGroup{
ID: g.ID,
Name: g.Name,
Description: g.Description,
Platform: g.Platform,
SubscriptionType: g.SubscriptionType,
RateMultiplier: g.RateMultiplier,
PeakRateEnabled: g.PeakRateEnabled,
PeakStart: g.PeakStart,
PeakEnd: g.PeakEnd,
PeakRateMultiplier: g.PeakRateMultiplier,
IsExclusive: g.IsExclusive,
ImageRateIndependent: g.ImageRateIndependent,
ImageRateMultiplier: g.ImageRateMultiplier,
Models: models,
ID: g.ID,
Name: g.Name,
Description: g.Description,
Platform: g.Platform,
SubscriptionType: g.SubscriptionType,
RateMultiplier: g.RateMultiplier,
PeakRateEnabled: g.PeakRateEnabled,
PeakStart: g.PeakStart,
PeakEnd: g.PeakEnd,
PeakRateMultiplier: g.PeakRateMultiplier,
IsExclusive: g.IsExclusive,
ImageRateIndependent: g.ImageRateIndependent,
ImageRateMultiplier: g.ImageRateMultiplier,
LongContextPricingEnabled: g.LongContextPricingEnabled,
Models: models,
}
if rate, ok := userRates[g.ID]; ok {
dto.UserRateMultiplier = &rate
@@ -189,6 +215,22 @@ func toModelPlazaGroupDTO(g *service.PlazaGroup, userRates map[int64]float64) mo
return dto
}
// toModelPlazaTimePricing 转换分时倍率;nil 透传(JSON 省略)。
func toModelPlazaTimePricing(p *service.TimePricingSchedule) *modelPlazaTimePricing {
if p == nil || len(p.Periods) == 0 {
return nil
}
periods := make([]modelPlazaTimePricingPeriod, 0, len(p.Periods))
for _, period := range p.Periods {
periods = append(periods, modelPlazaTimePricingPeriod{
StartTime: period.StartTime,
EndTime: period.EndTime,
Multiplier: period.Multiplier,
})
}
return &modelPlazaTimePricing{Timezone: p.Timezone, WeekdaysOnly: p.WeekdaysOnly, Periods: periods}
}
// toModelPlazaOfficialPricing 转换官方参考价;nil 透传(前端显示 "-")。
func toModelPlazaOfficialPricing(p *service.PlazaOfficialPricing) *modelPlazaOfficialPricing {
if p == nil {
@@ -200,5 +242,6 @@ func toModelPlazaOfficialPricing(p *service.PlazaOfficialPricing) *modelPlazaOff
CacheWritePrice: p.CacheWritePrice,
CacheWrite1hPrice: p.CacheWrite1hPrice,
CacheReadPrice: p.CacheReadPrice,
Intervals: toUserPricingIntervals(p.Intervals),
}
}
@@ -91,7 +91,7 @@ func TestToModelPlazaGroupDTO_UserRateAndFieldWhitelist(t *testing.T) {
"id", "name", "description", "platform", "subscription_type",
"rate_multiplier", "user_rate_multiplier", "is_exclusive", "models",
"peak_rate_enabled", "peak_start", "peak_end", "peak_rate_multiplier",
"image_rate_independent", "image_rate_multiplier",
"image_rate_independent", "image_rate_multiplier", "long_context_pricing_enabled",
} {
_, exists := decoded[key]
require.Truef(t, exists, "plaza group DTO must expose %q", key)
@@ -109,6 +109,12 @@ func TestToModelPlazaGroupDTO_UserRateAndFieldWhitelist(t *testing.T) {
require.Contains(t, official, "cache_read_price")
_, has1h := official["cache_write_1h_price"]
require.False(t, has1h, "1h 缓存写价为 nil 时应 omitempty")
_, hasOfficialIntervals := official["intervals"]
require.False(t, hasOfficialIntervals, "官方无阶梯时 intervals 应 omitempty")
_, hasBasis := model["long_context_basis"]
require.False(t, hasBasis, "单档模型不输出 long_context_basis")
_, hasTimePricing := model["time_pricing"]
require.False(t, hasTimePricing, "无分时时不输出 time_pricing")
// 无专属倍率:user_rate_multiplier 整个字段省略
dtoNoRate := toModelPlazaGroupDTO(&g, nil)
@@ -124,4 +130,94 @@ func TestToModelPlazaOfficialPricing_NilPassthrough(t *testing.T) {
require.Nil(t, toModelPlazaOfficialPricing(nil))
}
func TestToModelPlazaGroupDTO_LongContextTiersAndBasis(t *testing.T) {
maxTokens := 272000
g := service.PlazaGroup{
ID: 3, Name: "ladder", Platform: "openai", SubscriptionType: "standard", RateMultiplier: 1,
LongContextPricingEnabled: true,
Models: []service.PlazaModel{{
Name: "gpt-5.4",
Platform: "openai",
Pricing: &service.ChannelModelPricing{
BillingMode: service.BillingModeToken,
InputPrice: testPtr(2.5e-6),
Intervals: []service.PricingInterval{
{MinTokens: 0, MaxTokens: &maxTokens, TierLabel: "≤272K", InputPrice: testPtr(2.5e-6)},
{MinTokens: 272000, TierLabel: ">272K", InputPrice: testPtr(5e-6)},
},
},
OfficialPricing: &service.PlazaOfficialPricing{
InputPrice: testPtr(2.5e-6),
Intervals: []service.PricingInterval{
{MinTokens: 0, MaxTokens: &maxTokens, TierLabel: "≤272K", InputPrice: testPtr(2.5e-6)},
{MinTokens: 272000, TierLabel: ">272K", InputPrice: testPtr(5e-6)},
},
},
LongContextBasis: service.ContextPricingBasisWholeRequest,
}},
}
raw, err := json.Marshal(toModelPlazaGroupDTO(&g, nil))
require.NoError(t, err)
var decoded map[string]any
require.NoError(t, json.Unmarshal(raw, &decoded))
require.Equal(t, true, decoded["long_context_pricing_enabled"])
model := decoded["models"].([]any)[0].(map[string]any)
require.Equal(t, "whole_request", model["long_context_basis"])
pricing := model["pricing"].(map[string]any)
paidTiers := pricing["intervals"].([]any)
require.Len(t, paidTiers, 2)
require.Equal(t, ">272K", paidTiers[1].(map[string]any)["tier_label"])
official := model["official_pricing"].(map[string]any)
officialTiers := official["intervals"].([]any)
require.Len(t, officialTiers, 2)
first := officialTiers[0].(map[string]any)
require.Equal(t, "≤272K", first["tier_label"])
require.InDelta(t, 272000, first["max_tokens"].(float64), 0)
require.Contains(t, first, "cache_write_price", "区间 DTO 字段齐全(nil 输出 null)")
}
func testPtr(v float64) *float64 { return &v }
func TestToModelPlazaGroupDTO_TimePricing(t *testing.T) {
g := service.PlazaGroup{
ID: 4, Name: "cn", Platform: "deepseek", SubscriptionType: "standard", RateMultiplier: 1,
Models: []service.PlazaModel{{
Name: "deepseek-chat",
Platform: "deepseek",
Pricing: &service.ChannelModelPricing{BillingMode: service.BillingModeToken, InputPrice: testPtr(0.28e-6)},
TimePricing: &service.TimePricingSchedule{Timezone: "Asia/Shanghai", Periods: []service.TimePricingPeriod{
{StartTime: "00:30", EndTime: "08:30", Multiplier: 0.5},
}},
}, {
Name: "deepseek-reasoner",
Platform: "deepseek",
Pricing: &service.ChannelModelPricing{BillingMode: service.BillingModeToken, InputPrice: testPtr(0.56e-6)},
TimePricing: &service.TimePricingSchedule{Timezone: "Asia/Shanghai", WeekdaysOnly: true, Periods: []service.TimePricingPeriod{
{StartTime: "00:30", EndTime: "08:30", Multiplier: 0.5},
}},
}},
}
raw, err := json.Marshal(toModelPlazaGroupDTO(&g, nil))
require.NoError(t, err)
var decoded map[string]any
require.NoError(t, json.Unmarshal(raw, &decoded))
model := decoded["models"].([]any)[0].(map[string]any)
tp := model["time_pricing"].(map[string]any)
require.Equal(t, "Asia/Shanghai", tp["timezone"])
_, hasWeekdaysOnly := tp["weekdays_only"]
require.False(t, hasWeekdaysOnly, "未开启仅工作日时字段省略")
periods := tp["periods"].([]any)
require.Len(t, periods, 1)
first := periods[0].(map[string]any)
require.Equal(t, "00:30", first["start_time"])
require.Equal(t, "08:30", first["end_time"])
require.InDelta(t, 0.5, first["multiplier"].(float64), 1e-12)
weekdaysModel := decoded["models"].([]any)[1].(map[string]any)
weekdaysTP := weekdaysModel["time_pricing"].(map[string]any)
require.Equal(t, true, weekdaysTP["weekdays_only"])
}
@@ -0,0 +1,470 @@
package service
import (
"context"
"errors"
"math"
"sort"
"strconv"
"time"
)
// 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
}
// TimePricingPeriod 分时倍率时段:配置时区当天 [StartTime, EndTime) 内整单费用乘 Multiplier。
type TimePricingPeriod struct {
StartTime string
EndTime string
Multiplier float64
}
// TimePricingSchedule 分组+模型生效的分时倍率(仅列出倍率 ≠ 1 的时段,按开始时间升序)。
// WeekdaysOnly 为 true 时时段仅周一至周五生效,周末整天按标准价计费。
type TimePricingSchedule struct {
Timezone string
WeekdaysOnly bool
Periods []TimePricingPeriod
}
// ContextPricingSchedule 分组+模型按上下文长度分档的有效单价表。
// 单价由真实计费函数探针得出,与扣费同源;单档表示无阶梯。
// Tiers 为标准时段单价;TimePricing 非 nil 时,落在时段内的请求整单再乘对应倍率。
type ContextPricingSchedule struct {
Basis ContextPricingBasis
Tiers []ContextPricingTier
TimePricing *TimePricingSchedule
}
// 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)
}
tiers = mergeEqualContextTiers(tiers)
applyContextTierLabels(tiers, plan)
basis := ContextPricingBasisWholeRequest
if legacy != nil {
basis = ContextPricingBasisMarginal
}
return &ContextPricingSchedule{Basis: basis, Tiers: tiers, TimePricing: resolvedTimePricingSchedule(resolved)}, nil
}
// resolvedTimePricingSchedule 列出计费会生效的分时倍率时段。
// 时段来自解析到的渠道定价配置,每个时段的倍率用计费自己的 resolvedChannelTimeMultiplier
// 在时段内取值:定价来源不是渠道(分组价卡覆盖)、配置非法等情况下计费按 1 计,
// 这里也就自然得到"无分时"。倍率为 1 的时段不列出。
func resolvedTimePricingSchedule(resolved *ResolvedPricing) *TimePricingSchedule {
if resolved == nil || resolved.channelPricing == nil || resolved.channelPricing.TimePricing == nil {
return nil
}
cfg := resolved.channelPricing.TimePricing
location, err := loadChannelTimePricingLocation(cfg.Timezone)
if err != nil {
return nil
}
type probedPeriod struct {
start int
period TimePricingPeriod
}
probed := make([]probedPeriod, 0, len(cfg.Periods))
for _, period := range cfg.Periods {
start, err := parseChannelTime(period.StartTime, false)
if err != nil {
continue
}
// 时段按每日循环,取该时段开始后 1 秒作为探针时刻。
// 锚点日必须是工作日(2026-01-05 为周一):weekdays_only 配置在周末恒为 1,
// 锚点落在周末会把时段整组剔除。
at := time.Date(2026, time.January, 5, 0, 0, start+1, 0, location)
multiplier := resolvedChannelTimeMultiplier(resolved, at)
if multiplier == 1 {
continue
}
probed = append(probed, probedPeriod{start: start, period: TimePricingPeriod{
StartTime: period.StartTime,
EndTime: period.EndTime,
Multiplier: multiplier,
}})
}
if len(probed) == 0 {
return nil
}
sort.SliceStable(probed, func(i, j int) bool { return probed[i].start < probed[j].start })
out := &TimePricingSchedule{
Timezone: cfg.Timezone,
WeekdaysOnly: cfg.WeekdaysOnly,
Periods: make([]TimePricingPeriod, 0, len(probed)),
}
for _, p := range probed {
out.Periods = append(out.Periods, p.period)
}
return out
}
// 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
}
// 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 && 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
}
// applyContextTierLabels 给多档阶梯打统一形态的标签:有上限的档为「≤上限」,
// 末档为「>下限」;档位按上下文升序,因此相邻的 ≤100K / ≤200K 即表示 (100K,200K]。
// 目录阶梯/旧规则在"达到阈值即进高档"时改用 < / ≥ 表达阈值本身。
// 渠道区间上的 tier_label 不用于 token 档位(token 模式的管理表单不暴露该字段)。
func applyContextTierLabels(tiers []ContextPricingTier, plan contextBreakpointPlan) {
if len(tiers) < 2 {
return
}
for i := range tiers {
t := &tiers[i]
switch {
case plan.thresholdInclusive && t.MaxTokens != nil && *t.MaxTokens == plan.thresholdBound:
t.Label = "<" + formatContextTokenCount(plan.threshold)
case plan.thresholdInclusive && t.MinTokens == plan.thresholdBound:
t.Label = "≥" + formatContextTokenCount(plan.threshold)
case t.MaxTokens != nil:
t.Label = "≤" + formatContextTokenCount(*t.MaxTokens)
default:
t.Label = ">" + formatContextTokenCount(t.MinTokens)
}
}
}
// 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,620 @@
//go:build unit
package service
import (
"context"
"math"
"math/rand"
"testing"
"time"
"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), "≤200K", p(2e-6), p(15e-6), p(3.75e-6), p(0.3e-6))
requireTier(t, s.Tiers[1], 200000, nil, ">200K", 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), "≤100K", p(1e-6), p(15e-6), p(3.75e-6), p(0.3e-6))
requireTier(t, s.Tiers[1], 100000, intPtr(200000), "≤200K", p(2e-6), p(15e-6), p(3.75e-6), p(0.3e-6))
requireTier(t, s.Tiers[2], 200000, nil, ">200K", 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), "≤200K", p(2e-6), p(15e-6), p(3.75e-6), p(0.3e-6))
requireTier(t, s.Tiers[1], 200000, nil, ">200K", 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), "≤1M", p(4e-6), p(15e-6), p(3.75e-6), p(0.3e-6))
requireTier(t, s.Tiers[2], 1000000, nil, ">1M", 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, "档位连续")
require.Less(t, sched.Tiers[i-1].MinTokens, sched.Tiers[i].MinTokens, "档位按上下文升序")
}
if len(sched.Tiers) > 1 {
for i, tier := range sched.Tiers {
require.NotEmpty(t, tier.Label, "多档时每档都有标签 #%d", i)
}
require.Nil(t, sched.Tiers[len(sched.Tiers)-1].MaxTokens, "末档无上限")
}
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...)
}
func sonnetChannelWithTimePricing(tp *ChannelTimePricing) []ChannelModelPricing {
return []ChannelModelPricing{{
Platform: PlatformAnthropic, Models: []string{"claude-sonnet-4"}, BillingMode: BillingModeToken,
InputPrice: testPtrFloat64(2e-6), TimePricing: tp,
}}
}
func TestResolveContextPricingSchedule_TimePricing(t *testing.T) {
valid := &ChannelTimePricing{Timezone: "Asia/Shanghai", Periods: []ChannelTimePricingPeriod{
{StartTime: "18:00", EndTime: "22:00:00", Multiplier: 1.2},
{StartTime: "00:30", EndTime: "08:30", Multiplier: 0.5},
{StartTime: "12:00", EndTime: "13:00", Multiplier: 1},
}}
t.Run("渠道分时按开始时间升序列出且跳过倍率 1 的时段", func(t *testing.T) {
bs, resolver := newTokenCostTestEnv(t, PlatformAnthropic, sonnetChannelWithTimePricing(valid), nil)
sched, err := bs.ResolveContextPricingSchedule(context.Background(), resolver, ContextPricingScheduleInput{
Model: "claude-sonnet-4", Group: enabledGroup(PlatformAnthropic), Platform: PlatformAnthropic,
})
require.NoError(t, err)
require.NotNil(t, sched.TimePricing)
require.Equal(t, "Asia/Shanghai", sched.TimePricing.Timezone)
require.False(t, sched.TimePricing.WeekdaysOnly)
require.Equal(t, []TimePricingPeriod{
{StartTime: "00:30", EndTime: "08:30", Multiplier: 0.5},
{StartTime: "18:00", EndTime: "22:00:00", Multiplier: 1.2},
}, sched.TimePricing.Periods)
// 阶梯表单价是标准时段价,不含分时倍率
requirePrice(t, testPtrFloat64(2e-6), sched.Tiers[0].Input, "input")
// 对账:时段内的真实计费 = 标准单价 × token × 倍率
group := enabledGroup(PlatformAnthropic)
gid := group.ID
resolved := resolver.Resolve(context.Background(), PricingInput{Model: "claude-sonnet-4", GroupID: &gid, Group: group})
loc, err := time.LoadLocation("Asia/Shanghai")
require.NoError(t, err)
for _, tc := range []struct {
at time.Time
want float64
}{
{time.Date(2026, 8, 23, 3, 0, 0, 0, loc), 0.5},
{time.Date(2026, 8, 23, 12, 30, 0, 0, loc), 1},
{time.Date(2026, 8, 23, 21, 59, 59, 0, loc), 1.2},
{time.Date(2026, 8, 23, 22, 0, 0, 0, loc), 1},
} {
cost, err := bs.CalculateTokenCostForRequest(TokenCostRequest{
Ctx: context.Background(), Model: "claude-sonnet-4", Group: group, Tokens: UsageTokens{InputTokens: 1000},
RateMultiplier: 1, PricingAt: tc.at, Resolver: resolver, Resolved: resolved,
})
require.NoError(t, err)
assertCostClose(t, 1000*2e-6*tc.want, cost.ActualCost, "at %s", tc.at)
}
})
t.Run("仅工作日配置透传标注且时段仍列出", func(t *testing.T) {
weekdays := &ChannelTimePricing{Timezone: "Asia/Shanghai", WeekdaysOnly: true, Periods: []ChannelTimePricingPeriod{
{StartTime: "00:30", EndTime: "08:30", Multiplier: 0.5},
}}
bs, resolver := newTokenCostTestEnv(t, PlatformAnthropic, sonnetChannelWithTimePricing(weekdays), nil)
sched, err := bs.ResolveContextPricingSchedule(context.Background(), resolver, ContextPricingScheduleInput{
Model: "claude-sonnet-4", Group: enabledGroup(PlatformAnthropic), Platform: PlatformAnthropic,
})
require.NoError(t, err)
require.NotNil(t, sched.TimePricing, "探针锚点必须落在工作日,否则仅工作日配置的时段会被整组剔除")
require.True(t, sched.TimePricing.WeekdaysOnly)
require.Equal(t, []TimePricingPeriod{
{StartTime: "00:30", EndTime: "08:30", Multiplier: 0.5},
}, sched.TimePricing.Periods)
// 对账:工作日时段内乘倍率,周末同一时段按标准价
group := enabledGroup(PlatformAnthropic)
gid := group.ID
resolved := resolver.Resolve(context.Background(), PricingInput{Model: "claude-sonnet-4", GroupID: &gid, Group: group})
loc, err := time.LoadLocation("Asia/Shanghai")
require.NoError(t, err)
for _, tc := range []struct {
at time.Time
want float64
}{
{time.Date(2026, 8, 24, 3, 0, 0, 0, loc), 0.5}, // 周一
{time.Date(2026, 8, 23, 3, 0, 0, 0, loc), 1}, // 周日
} {
cost, err := bs.CalculateTokenCostForRequest(TokenCostRequest{
Ctx: context.Background(), Model: "claude-sonnet-4", Group: group, Tokens: UsageTokens{InputTokens: 1000},
RateMultiplier: 1, PricingAt: tc.at, Resolver: resolver, Resolved: resolved,
})
require.NoError(t, err)
assertCostClose(t, 1000*2e-6*tc.want, cost.ActualCost, "at %s", tc.at)
}
})
t.Run("配置非法时计费按 1 计,阶梯表不列分时", func(t *testing.T) {
invalid := &ChannelTimePricing{Timezone: "Asia/Shanghai", Periods: []ChannelTimePricingPeriod{
{StartTime: "00:30", EndTime: "08:30", Multiplier: 0.5},
{StartTime: "08:00", EndTime: "09:00", Multiplier: 0.8}, // 与上一段重叠
}}
bs, resolver := newTokenCostTestEnv(t, PlatformAnthropic, sonnetChannelWithTimePricing(invalid), nil)
sched, err := bs.ResolveContextPricingSchedule(context.Background(), resolver, ContextPricingScheduleInput{
Model: "claude-sonnet-4", Group: enabledGroup(PlatformAnthropic), Platform: PlatformAnthropic,
})
require.NoError(t, err)
require.Nil(t, sched.TimePricing)
})
t.Run("分组价卡覆盖后渠道分时不再生效", func(t *testing.T) {
bs, resolver := newTokenCostTestEnv(t, PlatformAnthropic, sonnetChannelWithTimePricing(valid), nil)
group := &Group{ID: 100, Platform: PlatformAnthropic, LongContextPricingEnabled: true, ModelPricing: []ChannelModelPricing{{
Models: []string{"claude-sonnet-4"}, BillingMode: BillingModeToken, InputPrice: testPtrFloat64(1e-6),
}}}
sched, err := bs.ResolveContextPricingSchedule(context.Background(), resolver, ContextPricingScheduleInput{
Model: "claude-sonnet-4", Group: group, Platform: PlatformAnthropic,
})
require.NoError(t, err)
require.Nil(t, sched.TimePricing)
})
t.Run("无分时配置为 nil", func(t *testing.T) {
bs, resolver := newTokenCostTestEnv(t, PlatformAnthropic, sonnetChannelWithTimePricing(nil), nil)
sched, err := bs.ResolveContextPricingSchedule(context.Background(), resolver, ContextPricingScheduleInput{
Model: "claude-sonnet-4", Group: enabledGroup(PlatformAnthropic), Platform: PlatformAnthropic,
})
require.NoError(t, err)
require.Nil(t, sched.TimePricing)
})
}
@@ -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)
}
@@ -90,7 +90,7 @@ func (s *ChannelService) ListAvailable(ctx context.Context) ([]AvailableChannel,
ch.normalizeBillingModelSource()
supported := ch.SupportedModels()
s.fillGlobalPricingFallback(supported)
fillGlobalPricingFallback(s.pricingService, supported)
out = append(out, AvailableChannel{
ID: ch.ID,
@@ -117,16 +117,16 @@ func (s *ChannelService) ListAvailable(ctx context.Context) ([]AvailableChannel,
// 1. Pricing == nil(渠道完全没声明该模型的定价条目)
// 2. Pricing 非 nil 但所有价格字段为空(admin UI 建了条目但没填价格)
//
// 当 s.pricingService 为 nil(测试场景),跳过回落。
func (s *ChannelService) fillGlobalPricingFallback(models []SupportedModel) {
if s.pricingService == nil {
// 当 pricingService 为 nil(测试场景),跳过回落。可用渠道与模型广场共用。
func fillGlobalPricingFallback(pricingService *PricingService, models []SupportedModel) {
if pricingService == nil {
return
}
for i := range models {
if !pricingNeedsFallback(models[i].Pricing) {
continue
}
lp := s.pricingService.GetModelPricing(models[i].Name)
lp := pricingService.GetModelPricing(models[i].Name)
if lp == nil {
continue
}
@@ -255,7 +255,7 @@ func TestFillGlobalPricingFallback_NilPricing(t *testing.T) {
models := []SupportedModel{
{Name: "claude-opus-4-5", Platform: "anthropic"},
}
svc.fillGlobalPricingFallback(models)
fillGlobalPricingFallback(svc.pricingService, models)
require.NotNil(t, models[0].Pricing)
require.NotNil(t, models[0].Pricing.InputPrice)
require.InDelta(t, 5e-6, *models[0].Pricing.InputPrice, 1e-12)
@@ -281,7 +281,7 @@ func TestFillGlobalPricingFallback_EmptyPricingFillsFromLiteLLM(t *testing.T) {
},
},
}
svc.fillGlobalPricingFallback(models)
fillGlobalPricingFallback(svc.pricingService, models)
require.NotNil(t, models[0].Pricing)
require.Equal(t, BillingModeImage, models[0].Pricing.BillingMode)
require.NotNil(t, models[0].Pricing.ImageOutputPrice)
@@ -302,7 +302,7 @@ func TestFillGlobalPricingFallback_KeepsExistingPrice(t *testing.T) {
models := []SupportedModel{
{Name: "served-model", Platform: "anthropic", Pricing: existing},
}
svc.fillGlobalPricingFallback(models)
fillGlobalPricingFallback(svc.pricingService, models)
require.Same(t, existing, models[0].Pricing)
}
-256
View File
@@ -1,256 +0,0 @@
package service
import (
"context"
"fmt"
"sort"
"strings"
)
// PlazaOfficialPricing 模型广场展示用的 LiteLLM 官方参考价(USD per token)。
// 字段为 nil 表示官方数据中该项缺失(0 视为未配置)。
type PlazaOfficialPricing struct {
InputPrice *float64
OutputPrice *float64
CacheWritePrice *float64 // 5m 缓存写入(= LiteLLM cache_creation)
CacheWrite1hPrice *float64 // 1h 缓存写入(LiteLLM cache_creation_above_1hr)
CacheReadPrice *float64
}
// PlazaModel 模型广场中单个模型条目:渠道定价 + 官方参考价。
type PlazaModel struct {
Name string
Platform string
Pricing *ChannelModelPricing
OfficialPricing *PlazaOfficialPricing
}
// PlazaGroup 模型广场中以分组为顶层的条目。
//
// 与 AvailableGroupRef 相比多了 Description 与 Models;Models 来自该分组关联渠道的
// 支持模型(普通分组按分组平台隔离,Composite 分组展开关联渠道已配置的
// 具体平台),与「可用渠道」页口径一致。
type PlazaGroup struct {
ID int64
Name string
Description string
Platform string
SubscriptionType string
RateMultiplier float64
PeakRateEnabled bool
PeakStart string
PeakEnd string
PeakRateMultiplier float64
IsExclusive bool
// 图片按次实付倍率:ImageRateIndependent 为 true 时,图片计费模型的实付
// = 档位价 × ImageRateMultiplier,不乘分组/用户专属倍率(与计费口径一致)。
ImageRateIndependent bool
ImageRateMultiplier float64
Models []PlazaModel
}
// ListPlazaGroups 返回模型广场数据:每个活跃分组附带其可用模型与定价。
//
// 聚合口径与 ListAvailable 一致(Active 渠道、SupportedModels ∪ 全局定价回落、
// 平台隔离),仅把顶层从渠道换成分组:
// - 渠道按 lower(name) 排序后遍历,保证同名模型去重结果确定;
// - 同分组同名模型「先见者胜」,仅当已存条目无定价而新条目有定价时升级替换;
// - 图片计费模型的档位价按实收口径合成(分组图片价 > 渠道档位价 > 渠道默认按次价,
// 见 plazaImageDisplayPricing);
// - 每个模型附带 LiteLLM 官方参考价(查不到为 nil);
// - 只返回 Models 非空的分组;分组按 RateMultiplier 升序(同倍率按名称),
// 组内模型按名称排序。
//
// 可见性过滤(专属分组)不在此层做,由 handler 按登录态裁剪。
func (s *ChannelService) ListPlazaGroups(ctx context.Context) ([]PlazaGroup, error) {
channels, err := s.repo.ListAll(ctx)
if err != nil {
return nil, fmt.Errorf("list channels: %w", err)
}
groups, err := s.groupRepo.ListActive(ctx)
if err != nil {
return nil, fmt.Errorf("list active groups: %w", err)
}
sort.SliceStable(channels, func(i, j int) bool {
return strings.ToLower(channels[i].Name) < strings.ToLower(channels[j].Name)
})
byGroup := make(map[int64]*PlazaGroup, len(groups))
groupEnt := make(map[int64]*Group, len(groups))
order := make([]int64, 0, len(groups))
for i := range groups {
g := &groups[i]
byGroup[g.ID] = &PlazaGroup{
ID: g.ID,
Name: g.Name,
Description: g.Description,
Platform: g.Platform,
SubscriptionType: g.SubscriptionType,
RateMultiplier: g.RateMultiplier,
PeakRateEnabled: g.PeakRateEnabled,
PeakStart: g.PeakStart,
PeakEnd: g.PeakEnd,
PeakRateMultiplier: g.PeakRateMultiplier,
IsExclusive: g.IsExclusive,
ImageRateIndependent: g.ImageRateIndependent,
ImageRateMultiplier: g.ImageRateMultiplier,
}
groupEnt[g.ID] = g
order = append(order, g.ID)
}
type modelKey struct {
platform string
name string
}
// modelIdx[groupID][platform+modelName] = index into byGroup[groupID].Models
modelIdx := make(map[int64]map[modelKey]int, len(groups))
for i := range channels {
ch := &channels[i]
if ch.Status != StatusActive {
continue
}
ch.normalizeBillingModelSource()
supported := ch.SupportedModels()
s.fillGlobalPricingFallback(supported)
for _, gid := range ch.GroupIDs {
pg, ok := byGroup[gid]
if !ok {
continue
}
idx := modelIdx[gid]
if idx == nil {
idx = make(map[modelKey]int, len(supported))
modelIdx[gid] = idx
}
for j := range supported {
m := supported[j]
if pg.Platform == PlatformComposite {
if !isConcreteRequestPlatform(m.Platform) {
continue
}
} else if m.Platform != pg.Platform {
continue
}
pricing := plazaImageDisplayPricing(m.Pricing, groupEnt[gid])
key := modelKey{platform: m.Platform, name: m.Name}
if at, seen := idx[key]; seen {
// 先见者胜;仅当已存条目无定价而新条目有定价时升级。
if pg.Models[at].Pricing == nil && pricing != nil {
pg.Models[at].Pricing = pricing
}
continue
}
idx[key] = len(pg.Models)
pg.Models = append(pg.Models, PlazaModel{
Name: m.Name,
Platform: m.Platform,
Pricing: pricing,
})
}
}
}
officialMemo := make(map[string]*PlazaOfficialPricing)
out := make([]PlazaGroup, 0, len(order))
for _, gid := range order {
pg := byGroup[gid]
if len(pg.Models) == 0 {
continue
}
sort.SliceStable(pg.Models, func(i, j int) bool {
if pg.Models[i].Name != pg.Models[j].Name {
return pg.Models[i].Name < pg.Models[j].Name
}
return pg.Models[i].Platform < pg.Models[j].Platform
})
for j := range pg.Models {
pg.Models[j].OfficialPricing = s.lookupOfficialPricing(pg.Models[j].Name, officialMemo)
}
out = append(out, *pg)
}
sort.SliceStable(out, func(i, j int) bool {
if out[i].RateMultiplier != out[j].RateMultiplier {
return out[i].RateMultiplier < out[j].RateMultiplier
}
return out[i].Name < out[j].Name
})
return out, nil
}
// plazaImageDisplayPricing 为图片计费模型合成展示定价,使档位价与实收口径一致:
// 每档(1K/2K/4K)单价 = 分组图片价 > 渠道同档位价 > 渠道默认按次价,无价的档不展示。
// 分组未配任何图片价、或定价非图片模式时原样返回。返回克隆,不修改入参
// (渠道定价指针指向缓存共享数据)。
func plazaImageDisplayPricing(p *ChannelModelPricing, g *Group) *ChannelModelPricing {
if p == nil || g == nil || p.BillingMode != BillingModeImage {
return p
}
if g.ImagePrice1K == nil && g.ImagePrice2K == nil && g.ImagePrice4K == nil {
return p
}
channelTierPrice := func(label string) *float64 {
for i := range p.Intervals {
if p.Intervals[i].TierLabel == label && p.Intervals[i].PerRequestPrice != nil {
return p.Intervals[i].PerRequestPrice
}
}
return p.PerRequestPrice
}
tiers := []struct {
label string
groupPrice *float64
}{
{"1K", g.ImagePrice1K},
{"2K", g.ImagePrice2K},
{"4K", g.ImagePrice4K},
}
clone := *p
clone.Intervals = make([]PricingInterval, 0, len(tiers))
for i, t := range tiers {
price := t.groupPrice
if price == nil {
price = channelTierPrice(t.label)
}
if price == nil {
continue
}
v := *price
clone.Intervals = append(clone.Intervals, PricingInterval{
TierLabel: t.label,
PerRequestPrice: &v,
SortOrder: i,
})
}
return &clone
}
// lookupOfficialPricing 查询模型的 LiteLLM 官方参考价,带 memo 避免同名模型重复转换。
// pricingService 为 nil(测试场景)或查不到时返回 nil。
func (s *ChannelService) lookupOfficialPricing(modelName string, memo map[string]*PlazaOfficialPricing) *PlazaOfficialPricing {
if s.pricingService == nil {
return nil
}
if cached, ok := memo[modelName]; ok {
return cached
}
var result *PlazaOfficialPricing
if lp := s.pricingService.GetModelPricing(modelName); lp != nil && !lp.TokenPricingAbsent {
result = &PlazaOfficialPricing{
InputPrice: nonZeroPtr(lp.InputCostPerToken),
OutputPrice: nonZeroPtr(lp.OutputCostPerToken),
CacheWritePrice: nonZeroPtr(lp.CacheCreationInputTokenCost),
CacheWrite1hPrice: nonZeroPtr(lp.CacheCreationInputTokenCostAbove1hr),
CacheReadPrice: nonZeroPtr(lp.CacheReadInputTokenCost),
}
if result.InputPrice == nil && result.OutputPrice == nil &&
result.CacheWritePrice == nil && result.CacheWrite1hPrice == nil && result.CacheReadPrice == nil {
result = nil
}
}
memo[modelName] = result
return result
}
@@ -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,
@@ -0,0 +1,364 @@
package service
import (
"context"
"fmt"
"sort"
"strings"
)
// PlazaOfficialPricing 模型广场展示用的官方参考价(USD per token),与计费同源:
// LiteLLM → 内置兜底价卡 → 模型策略。字段为 nil 表示该项缺失(0 视为未配置)。
type PlazaOfficialPricing struct {
InputPrice *float64
OutputPrice *float64
CacheWritePrice *float64 // 5m 缓存写入(= LiteLLM cache_creation)
CacheWrite1hPrice *float64 // 1h 缓存写入,仅计费会区分 5m/1h 时给出
CacheReadPrice *float64
// Intervals 官方长上下文阶梯(多档时给出),不受分组开关影响。
Intervals []PricingInterval
}
// PlazaModel 模型广场中单个模型条目:按实收口径合成的展示定价 + 官方参考价。
type PlazaModel struct {
Name string
Platform string
Pricing *ChannelModelPricing
OfficialPricing *PlazaOfficialPricing
// LongContextBasis 多档时的计价基准(整单 / 仅超出部分),单档为空。
LongContextBasis ContextPricingBasis
// TimePricing 计费会生效的分时倍率时段;无分时为 nil。
TimePricing *TimePricingSchedule
}
// PlazaGroup 模型广场中以分组为顶层的条目。
//
// 与 AvailableGroupRef 相比多了 Description 与 Models;Models 来自该分组关联渠道的
// 支持模型(普通分组按分组平台隔离,Composite 分组展开关联渠道已配置的
// 具体平台),与「可用渠道」页口径一致。
type PlazaGroup struct {
ID int64
Name string
Description string
Platform string
SubscriptionType string
RateMultiplier float64
PeakRateEnabled bool
PeakStart string
PeakEnd string
PeakRateMultiplier float64
IsExclusive bool
// 图片按次实付倍率:ImageRateIndependent 为 true 时,图片计费模型的实付
// = 档位价 × ImageRateMultiplier,不乘分组/用户专属倍率(与计费口径一致)。
ImageRateIndependent bool
ImageRateMultiplier float64
// LongContextPricingEnabled 分组是否按上下文长度应用阶梯价;关闭时模型展示的是最低档。
LongContextPricingEnabled bool
Models []PlazaModel
}
// ModelPlazaService 聚合模型广场数据。
//
// 模型枚举来自渠道配置;token 模型的展示单价与阶梯由 BillingService 的阶梯表
// 查询给出(与扣费走同一条解析链与计费函数),图片/按次模型沿用渠道/分组档位价。
type ModelPlazaService struct {
channelRepo ChannelRepository
groupRepo GroupRepository
pricingService *PricingService
billingService *BillingService
resolver *ModelPricingResolver
}
// NewModelPlazaService 创建模型广场服务。
func NewModelPlazaService(
channelRepo ChannelRepository,
groupRepo GroupRepository,
pricingService *PricingService,
billingService *BillingService,
resolver *ModelPricingResolver,
) *ModelPlazaService {
return &ModelPlazaService{
channelRepo: channelRepo,
groupRepo: groupRepo,
pricingService: pricingService,
billingService: billingService,
resolver: resolver,
}
}
// ListGroups 返回模型广场数据:每个活跃分组附带其可用模型与定价。
//
// 模型枚举口径与 ListAvailable 一致(Active 渠道、SupportedModels ∪ 全局定价回落、
// 平台隔离),仅把顶层从渠道换成分组:
// - 渠道按 lower(name) 排序后遍历,保证同名模型去重结果确定;
// - 同分组同名模型「先见者胜」,仅当已存条目无定价而新条目有定价时升级替换;
// - token 模型的单价与阶梯按实收口径合成(见 ResolveContextPricingSchedule),
// 图片计费模型的档位价按实收口径合成(见 plazaImageDisplayPricing);
// - 每个模型附带官方参考价(查不到为 nil);
// - 只返回 Models 非空的分组;分组按 RateMultiplier 升序(同倍率按名称),
// 组内模型按名称排序。
//
// 可见性过滤(专属分组)不在此层做,由 handler 按登录态裁剪。
func (s *ModelPlazaService) ListGroups(ctx context.Context) ([]PlazaGroup, error) {
channels, err := s.channelRepo.ListAll(ctx)
if err != nil {
return nil, fmt.Errorf("list channels: %w", err)
}
groups, err := s.groupRepo.ListActive(ctx)
if err != nil {
return nil, fmt.Errorf("list active groups: %w", err)
}
sort.SliceStable(channels, func(i, j int) bool {
return strings.ToLower(channels[i].Name) < strings.ToLower(channels[j].Name)
})
byGroup := make(map[int64]*PlazaGroup, len(groups))
groupEnt := make(map[int64]*Group, len(groups))
order := make([]int64, 0, len(groups))
for i := range groups {
g := &groups[i]
byGroup[g.ID] = &PlazaGroup{
ID: g.ID,
Name: g.Name,
Description: g.Description,
Platform: g.Platform,
SubscriptionType: g.SubscriptionType,
RateMultiplier: g.RateMultiplier,
PeakRateEnabled: g.PeakRateEnabled,
PeakStart: g.PeakStart,
PeakEnd: g.PeakEnd,
PeakRateMultiplier: g.PeakRateMultiplier,
IsExclusive: g.IsExclusive,
ImageRateIndependent: g.ImageRateIndependent,
ImageRateMultiplier: g.ImageRateMultiplier,
LongContextPricingEnabled: g.LongContextPricingEnabled,
}
groupEnt[g.ID] = g
order = append(order, g.ID)
}
type modelKey struct {
platform string
name string
}
// modelIdx[groupID][platform+modelName] = index into byGroup[groupID].Models
modelIdx := make(map[int64]map[modelKey]int, len(groups))
for i := range channels {
ch := &channels[i]
if ch.Status != StatusActive {
continue
}
ch.normalizeBillingModelSource()
supported := ch.SupportedModels()
fillGlobalPricingFallback(s.pricingService, supported)
for _, gid := range ch.GroupIDs {
pg, ok := byGroup[gid]
if !ok {
continue
}
idx := modelIdx[gid]
if idx == nil {
idx = make(map[modelKey]int, len(supported))
modelIdx[gid] = idx
}
for j := range supported {
m := supported[j]
if pg.Platform == PlatformComposite {
if !isConcreteRequestPlatform(m.Platform) {
continue
}
} else if m.Platform != pg.Platform {
continue
}
key := modelKey{platform: m.Platform, name: m.Name}
if at, seen := idx[key]; seen {
// 先见者胜;仅当已存条目无定价而新条目有定价时升级。
if pg.Models[at].Pricing == nil && m.Pricing != nil {
pg.Models[at].Pricing = m.Pricing
}
continue
}
idx[key] = len(pg.Models)
pg.Models = append(pg.Models, PlazaModel{
Name: m.Name,
Platform: m.Platform,
Pricing: m.Pricing,
})
}
}
}
officialMemo := make(map[string]*PlazaOfficialPricing)
out := make([]PlazaGroup, 0, len(order))
for _, gid := range order {
pg := byGroup[gid]
if len(pg.Models) == 0 {
continue
}
sort.SliceStable(pg.Models, func(i, j int) bool {
if pg.Models[i].Name != pg.Models[j].Name {
return pg.Models[i].Name < pg.Models[j].Name
}
return pg.Models[i].Platform < pg.Models[j].Platform
})
g := groupEnt[gid]
for j := range pg.Models {
s.fillDisplayPricing(ctx, &pg.Models[j], g)
pg.Models[j].OfficialPricing = s.lookupOfficialPricing(ctx, pg.Models[j].Name, officialMemo)
}
out = append(out, *pg)
}
sort.SliceStable(out, func(i, j int) bool {
if out[i].RateMultiplier != out[j].RateMultiplier {
return out[i].RateMultiplier < out[j].RateMultiplier
}
return out[i].Name < out[j].Name
})
return out, nil
}
// fillDisplayPricing 把模型的展示定价换成实收口径:
// token 模型取计费阶梯表(单价与档位均由真实计费函数得出),
// 图片/按次模型(或阶梯表不可用时)沿用渠道定价与分组图片档位价。
func (s *ModelPlazaService) fillDisplayPricing(ctx context.Context, m *PlazaModel, g *Group) {
if s.billingService != nil && s.resolver != nil {
sched, err := s.billingService.ResolveContextPricingSchedule(ctx, s.resolver, ContextPricingScheduleInput{
Model: m.Name,
Group: g,
Platform: m.Platform,
})
if err == nil && sched != nil && len(sched.Tiers) > 0 {
m.Pricing = plazaPricingFromSchedule(m.Pricing, sched)
if len(sched.Tiers) > 1 {
m.LongContextBasis = sched.Basis
}
m.TimePricing = sched.TimePricing
return
}
}
m.Pricing = plazaImageDisplayPricing(m.Pricing, g)
}
// plazaPricingFromSchedule 把阶梯表压成展示用的 ChannelModelPricing:
// 平价取首档单价,多档时 Intervals 逐档给出绝对单价;图片/按次字段沿用原始定价。
func plazaPricingFromSchedule(raw *ChannelModelPricing, sched *ContextPricingSchedule) *ChannelModelPricing {
out := &ChannelModelPricing{BillingMode: BillingModeToken}
if raw != nil {
out.ImageInputPrice = raw.ImageInputPrice
out.ImageOutputPrice = raw.ImageOutputPrice
out.PerRequestPrice = raw.PerRequestPrice
}
first := sched.Tiers[0]
out.InputPrice = first.Input
out.OutputPrice = first.Output
out.CacheWritePrice = first.CacheWrite
out.CacheReadPrice = first.CacheRead
if len(sched.Tiers) > 1 {
out.Intervals = plazaIntervalsFromTiers(sched.Tiers)
}
return out
}
func plazaIntervalsFromTiers(tiers []ContextPricingTier) []PricingInterval {
intervals := make([]PricingInterval, 0, len(tiers))
for i, t := range tiers {
intervals = append(intervals, PricingInterval{
MinTokens: t.MinTokens,
MaxTokens: t.MaxTokens,
TierLabel: t.Label,
InputPrice: t.Input,
OutputPrice: t.Output,
CacheWritePrice: t.CacheWrite,
CacheReadPrice: t.CacheRead,
SortOrder: i,
})
}
return intervals
}
// plazaImageDisplayPricing 为图片计费模型合成展示定价,使档位价与实收口径一致:
// 每档(1K/2K/4K)单价 = 分组图片价 > 渠道同档位价 > 渠道默认按次价,无价的档不展示。
// 分组未配任何图片价、或定价非图片模式时原样返回。返回克隆,不修改入参
// (渠道定价指针指向缓存共享数据)。
func plazaImageDisplayPricing(p *ChannelModelPricing, g *Group) *ChannelModelPricing {
if p == nil || g == nil || p.BillingMode != BillingModeImage {
return p
}
if g.ImagePrice1K == nil && g.ImagePrice2K == nil && g.ImagePrice4K == nil {
return p
}
channelTierPrice := func(label string) *float64 {
for i := range p.Intervals {
if p.Intervals[i].TierLabel == label && p.Intervals[i].PerRequestPrice != nil {
return p.Intervals[i].PerRequestPrice
}
}
return p.PerRequestPrice
}
tiers := []struct {
label string
groupPrice *float64
}{
{"1K", g.ImagePrice1K},
{"2K", g.ImagePrice2K},
{"4K", g.ImagePrice4K},
}
clone := *p
clone.Intervals = make([]PricingInterval, 0, len(tiers))
for i, t := range tiers {
price := t.groupPrice
if price == nil {
price = channelTierPrice(t.label)
}
if price == nil {
continue
}
v := *price
clone.Intervals = append(clone.Intervals, PricingInterval{
TierLabel: t.label,
PerRequestPrice: &v,
SortOrder: i,
})
}
return &clone
}
// lookupOfficialPricing 查询模型的官方参考价(与计费同源:LiteLLM → 内置兜底 → 模型策略),
// 带 memo 避免同名模型重复解析。官方阶梯按无分组、无渠道的口径查阶梯表。
// billingService 为 nil(测试场景)或查不到时返回 nil。
func (s *ModelPlazaService) lookupOfficialPricing(ctx context.Context, modelName string, memo map[string]*PlazaOfficialPricing) *PlazaOfficialPricing {
if s.billingService == nil {
return nil
}
if cached, ok := memo[modelName]; ok {
return cached
}
var result *PlazaOfficialPricing
if mp, err := s.billingService.GetModelPricing(modelName); err == nil && mp != nil {
result = &PlazaOfficialPricing{
InputPrice: nonZeroPtr(mp.InputPricePerToken),
OutputPrice: nonZeroPtr(mp.OutputPricePerToken),
CacheWritePrice: nonZeroPtr(mp.CacheCreationPricePerToken),
CacheReadPrice: nonZeroPtr(mp.CacheReadPricePerToken),
}
// 计费只在支持 5m/1h 分档时使用 1h 价,其余情况 1h 价对用户无意义。
if mp.SupportsCacheBreakdown {
result.CacheWrite1hPrice = nonZeroPtr(mp.CacheCreation1hPrice)
}
if s.resolver != nil {
sched, schedErr := s.billingService.ResolveContextPricingSchedule(ctx, s.resolver, ContextPricingScheduleInput{Model: modelName})
if schedErr == nil && sched != nil && len(sched.Tiers) > 1 {
result.Intervals = plazaIntervalsFromTiers(sched.Tiers)
}
}
if result.InputPrice == nil && result.OutputPrice == nil && result.CacheWritePrice == nil &&
result.CacheWrite1hPrice == nil && result.CacheReadPrice == nil && len(result.Intervals) == 0 {
result = nil
}
}
memo[modelName] = result
return result
}
@@ -7,17 +7,16 @@ import (
"errors"
"testing"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/stretchr/testify/require"
)
// newPlazaChannelService 构造 ListPlazaGroups 测试用的 ChannelService。
func newPlazaChannelService(channels []Channel, groups []Group, pricing *PricingService) *ChannelService {
// newPlazaService 构造 ListGroups 测试用的 ModelPlazaService(不接计费服务:展示定价原样透传)。
func newPlazaService(channels []Channel, groups []Group, pricing *PricingService) *ModelPlazaService {
repo := &mockChannelRepository{
listAllFn: func(ctx context.Context) ([]Channel, error) { return channels, nil },
}
svc := NewChannelService(repo, &stubGroupRepoForAvailable{activeGroups: groups}, nil, nil)
svc.pricingService = pricing
return svc
return NewModelPlazaService(repo, &stubGroupRepoForAvailable{activeGroups: groups}, pricing, nil, nil)
}
func plazaPricedChannel(id int64, name string, groupIDs []int64, platform string, models ...string) Channel {
@@ -46,8 +45,8 @@ func TestListPlazaGroups_GroupCentricAggregation(t *testing.T) {
{ID: 10, Name: "g-main", Description: "desc", Platform: "anthropic", RateMultiplier: 1},
{ID: 20, Name: "g-empty", Platform: "anthropic", RateMultiplier: 0.5},
}
svc := newPlazaChannelService(channels, groups, nil)
out, err := svc.ListPlazaGroups(context.Background())
svc := newPlazaService(channels, groups, nil)
out, err := svc.ListGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 1, "无模型的分组不应返回")
require.Equal(t, int64(10), out[0].ID)
@@ -71,8 +70,8 @@ func TestListPlazaGroups_DedupFirstWinsWithPricingUpgrade(t *testing.T) {
groups := []Group{{ID: 10, Name: "g", Platform: "anthropic", RateMultiplier: 1}}
// alpha(无价)按名称序先于 beta(有价):先见者无价,应被有价条目升级。
svc := newPlazaChannelService([]Channel{priced, unpriced}, groups, nil)
out, err := svc.ListPlazaGroups(context.Background())
svc := newPlazaService([]Channel{priced, unpriced}, groups, nil)
out, err := svc.ListGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 1)
require.Len(t, out[0].Models, 1)
@@ -93,8 +92,8 @@ func TestListPlazaGroups_PlatformIsolation(t *testing.T) {
{ID: 10, Name: "g-claude", Platform: "anthropic", RateMultiplier: 1},
{ID: 20, Name: "g-gpt", Platform: "openai", RateMultiplier: 1},
}
svc := newPlazaChannelService([]Channel{ch}, groups, nil)
out, err := svc.ListPlazaGroups(context.Background())
svc := newPlazaService([]Channel{ch}, groups, nil)
out, err := svc.ListGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 2)
byName := map[string][]PlazaModel{}
@@ -122,7 +121,7 @@ func TestListPlazaGroups_CompositeIncludesConfiguredConcretePlatforms(t *testing
}
groups := []Group{{ID: 10, Name: "composite", Platform: PlatformComposite, RateMultiplier: 1}}
out, err := newPlazaChannelService([]Channel{ch}, groups, nil).ListPlazaGroups(context.Background())
out, err := newPlazaService([]Channel{ch}, groups, nil).ListGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 1)
@@ -146,7 +145,7 @@ func TestListPlazaGroups_CompositeAndOrdinaryGroupsDoNotLeakPlatforms(t *testing
{ID: 20, Name: "composite", Platform: PlatformComposite, RateMultiplier: 1},
}
out, err := newPlazaChannelService([]Channel{ch}, groups, nil).ListPlazaGroups(context.Background())
out, err := newPlazaService([]Channel{ch}, groups, nil).ListGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 2)
@@ -173,8 +172,8 @@ func TestListPlazaGroups_InactiveChannelSkipped(t *testing.T) {
inactive := plazaPricedChannel(1, "off", []int64{10}, "anthropic", "claude-sonnet")
inactive.Status = "inactive"
groups := []Group{{ID: 10, Name: "g", Platform: "anthropic", RateMultiplier: 1}}
svc := newPlazaChannelService([]Channel{inactive}, groups, nil)
out, err := svc.ListPlazaGroups(context.Background())
svc := newPlazaService([]Channel{inactive}, groups, nil)
out, err := svc.ListGroups(context.Background())
require.NoError(t, err)
require.Empty(t, out)
}
@@ -188,8 +187,8 @@ func TestListPlazaGroups_SortedByRateMultiplierAsc(t *testing.T) {
{ID: 20, Name: "a-standard", Platform: "anthropic", RateMultiplier: 1},
{ID: 30, Name: "cheap", Platform: "anthropic", RateMultiplier: 0.5},
}
svc := newPlazaChannelService(channels, groups, nil)
out, err := svc.ListPlazaGroups(context.Background())
svc := newPlazaService(channels, groups, nil)
out, err := svc.ListGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 3)
require.Equal(t, "cheap", out[0].Name, "倍率低者在前")
@@ -213,8 +212,11 @@ func TestListPlazaGroups_OfficialPricingFill(t *testing.T) {
plazaPricedChannel(1, "ch", []int64{10}, "anthropic", "claude-sonnet", "unknown-model", "token-absent"),
}
groups := []Group{{ID: 10, Name: "g", Platform: "anthropic", RateMultiplier: 1}}
svc := newPlazaChannelService(channels, groups, pricingSvc)
out, err := svc.ListPlazaGroups(context.Background())
svc := newPlazaService(channels, groups, pricingSvc)
// 官方价与计费同源:需要计费服务与解析器(官方参考不查渠道,解析器无需渠道服务)。
svc.billingService = NewBillingService(&config.Config{}, pricingSvc)
svc.resolver = NewModelPricingResolver(nil, svc.billingService)
out, err := svc.ListGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 1)
require.Len(t, out[0].Models, 3)
@@ -256,8 +258,8 @@ func TestListPlazaGroups_GroupImagePriceOverridesChannelPricing(t *testing.T) {
ImagePrice1K: &imgPrice, ImageRateIndependent: true, ImageRateMultiplier: 1},
{ID: 20, Name: "g-plain", Platform: "openai", RateMultiplier: 0.1},
}
svc := newPlazaChannelService(channels, groups, nil)
out, err := svc.ListPlazaGroups(context.Background())
svc := newPlazaService(channels, groups, nil)
out, err := svc.ListGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 2)
byName := map[string]PlazaGroup{}
@@ -298,8 +300,8 @@ func TestListPlazaGroups_GroupImagePriceIgnoredForNonImageModes(t *testing.T) {
imgPrice := 0.02
channels := []Channel{plazaPricedChannel(1, "ch", []int64{10}, "openai", "gpt-5")}
groups := []Group{{ID: 10, Name: "g", Platform: "openai", RateMultiplier: 1, ImagePrice1K: &imgPrice}}
svc := newPlazaChannelService(channels, groups, nil)
out, err := svc.ListPlazaGroups(context.Background())
svc := newPlazaService(channels, groups, nil)
out, err := svc.ListGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 1)
p := out[0].Models[0].Pricing
@@ -314,17 +316,180 @@ func TestListPlazaGroups_RepoErrorsPropagate(t *testing.T) {
repo := &mockChannelRepository{
listAllFn: func(ctx context.Context) ([]Channel, error) { return nil, sentinel },
}
svc := NewChannelService(repo, &stubGroupRepoForAvailable{}, nil, nil)
out, err := svc.ListPlazaGroups(context.Background())
svc := NewModelPlazaService(repo, &stubGroupRepoForAvailable{}, nil, nil, nil)
out, err := svc.ListGroups(context.Background())
require.Nil(t, out)
require.ErrorIs(t, err, sentinel)
svc2 := NewChannelService(
svc2 := NewModelPlazaService(
&mockChannelRepository{listAllFn: func(ctx context.Context) ([]Channel, error) { return nil, nil }},
&stubGroupRepoForAvailable{listActiveErr: sentinel},
nil, nil,
nil, nil, nil,
)
out2, err2 := svc2.ListPlazaGroups(context.Background())
out2, err2 := svc2.ListGroups(context.Background())
require.Nil(t, out2)
require.ErrorIs(t, err2, sentinel)
}
// newPlazaServiceWithBilling 构造接入计费服务与解析器的广场服务:解析器的渠道服务与广场共用同一份渠道数据。
func newPlazaServiceWithBilling(channels []Channel, groups []Group, groupPlatforms map[int64]string, catalog *PricingService) *ModelPlazaService {
repo := &mockChannelRepository{
listAllFn: func(ctx context.Context) ([]Channel, error) { return channels, nil },
getGroupPlatformsFn: func(ctx context.Context, _ []int64) (map[int64]string, error) {
return groupPlatforms, nil
},
}
cs := NewChannelService(repo, nil, nil, nil)
bs := NewBillingService(&config.Config{}, catalog)
return NewModelPlazaService(repo, &stubGroupRepoForAvailable{activeGroups: groups}, catalog, bs, NewModelPricingResolver(cs, bs))
}
func plazaModelsByName(models []PlazaModel) map[string]PlazaModel {
out := make(map[string]PlazaModel, len(models))
for _, m := range models {
out[m.Name] = m
}
return out
}
func TestListGroups_TokenLadderFollowsGroupToggle(t *testing.T) {
// 同一渠道挂开启/关闭阶梯的两个分组:实付档位随分组开关,官方阶梯不受影响。
channels := []Channel{{
ID: 1, Name: "ch", Status: StatusActive, GroupIDs: []int64{10, 20},
ModelPricing: []ChannelModelPricing{{Platform: PlatformOpenAI, Models: []string{"gpt-5.4"}, BillingMode: BillingModeToken}},
}}
groups := []Group{
{ID: 10, Name: "on", Platform: PlatformOpenAI, RateMultiplier: 1, LongContextPricingEnabled: true},
{ID: 20, Name: "off", Platform: PlatformOpenAI, RateMultiplier: 2, LongContextPricingEnabled: false},
}
svc := newPlazaServiceWithBilling(channels, groups, map[int64]string{10: PlatformOpenAI, 20: PlatformOpenAI}, nil)
out, err := svc.ListGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 2)
on, off := out[0], out[1]
require.True(t, on.LongContextPricingEnabled)
require.False(t, off.LongContextPricingEnabled)
onModel := on.Models[0]
require.Equal(t, ContextPricingBasisWholeRequest, onModel.LongContextBasis)
require.Len(t, onModel.Pricing.Intervals, 2)
require.Equal(t, "≤272K", onModel.Pricing.Intervals[0].TierLabel)
require.Equal(t, ">272K", onModel.Pricing.Intervals[1].TierLabel)
require.InDelta(t, 2.5e-6, *onModel.Pricing.InputPrice, 1e-15)
require.InDelta(t, 5e-6, *onModel.Pricing.Intervals[1].InputPrice, 1e-15)
require.InDelta(t, 22.5e-6, *onModel.Pricing.Intervals[1].OutputPrice, 1e-15)
require.InDelta(t, 5e-6, *onModel.Pricing.Intervals[1].CacheWritePrice, 1e-15)
require.InDelta(t, 0.5e-6, *onModel.Pricing.Intervals[1].CacheReadPrice, 1e-15)
offModel := off.Models[0]
require.Empty(t, offModel.LongContextBasis)
require.Empty(t, offModel.Pricing.Intervals)
require.InDelta(t, 2.5e-6, *offModel.Pricing.InputPrice, 1e-15)
for _, m := range []PlazaModel{onModel, offModel} {
require.NotNil(t, m.OfficialPricing)
require.Len(t, m.OfficialPricing.Intervals, 2, "官方阶梯不受分组开关影响")
require.InDelta(t, 5e-6, *m.OfficialPricing.Intervals[1].InputPrice, 1e-15)
require.InDelta(t, 2.5e-6, *m.OfficialPricing.InputPrice, 1e-15)
}
}
func TestListGroups_GeminiLegacyRuleShownAsMarginal(t *testing.T) {
channels := []Channel{{
ID: 1, Name: "ch", Status: StatusActive, GroupIDs: []int64{10},
ModelMapping: map[string]map[string]string{PlatformGemini: {"gemini-2.5-pro": "gemini-2.5-pro"}},
}}
groups := []Group{{ID: 10, Name: "g", Platform: PlatformGemini, RateMultiplier: 1, LongContextPricingEnabled: true}}
svc := newPlazaServiceWithBilling(channels, groups, map[int64]string{10: PlatformGemini}, geminiCatalogStub())
out, err := svc.ListGroups(context.Background())
require.NoError(t, err)
require.Len(t, out, 1)
m := out[0].Models[0]
require.Equal(t, ContextPricingBasisMarginal, m.LongContextBasis)
require.Len(t, m.Pricing.Intervals, 2)
require.Equal(t, "≤200K", m.Pricing.Intervals[0].TierLabel)
require.Equal(t, ">200K", m.Pricing.Intervals[1].TierLabel)
require.InDelta(t, 2.5e-6, *m.Pricing.Intervals[1].InputPrice, 1e-15)
require.InDelta(t, 10e-6, *m.Pricing.Intervals[1].OutputPrice, 1e-15)
// 官方参考不套用站内旧规则
require.NotNil(t, m.OfficialPricing)
require.Empty(t, m.OfficialPricing.Intervals)
}
func TestListGroups_GroupTokenCardOverridesChannelPricing(t *testing.T) {
channels := []Channel{plazaPricedChannel(1, "ch", []int64{10}, PlatformAnthropic, "claude-sonnet-4")}
groups := []Group{{
ID: 10, Name: "g", Platform: PlatformAnthropic, RateMultiplier: 1, LongContextPricingEnabled: true,
ModelPricing: []ChannelModelPricing{{Models: []string{"claude-sonnet-*"}, BillingMode: BillingModeToken, InputPrice: testPtrFloat64(1e-6)}},
}}
svc := newPlazaServiceWithBilling(channels, groups, map[int64]string{10: PlatformAnthropic}, nil)
out, err := svc.ListGroups(context.Background())
require.NoError(t, err)
m := out[0].Models[0]
require.InDelta(t, 1e-6, *m.Pricing.InputPrice, 1e-15, "分组价卡优先于渠道平价")
require.InDelta(t, 15e-6, *m.Pricing.OutputPrice, 1e-15, "卡未配置的项回落目录价")
require.Empty(t, m.Pricing.Intervals)
}
func TestListGroups_ImageModelKeepsTierSynthesisWithBilling(t *testing.T) {
channels := []Channel{{
ID: 1, Name: "ch", Status: StatusActive, GroupIDs: []int64{10},
ModelPricing: []ChannelModelPricing{{
Platform: PlatformOpenAI, Models: []string{"gpt-image-2"}, BillingMode: BillingModeImage,
PerRequestPrice: testPtrFloat64(0.04),
}},
}}
groups := []Group{{
ID: 10, Name: "g", Platform: PlatformOpenAI, RateMultiplier: 1, LongContextPricingEnabled: true,
ImagePrice1K: testPtrFloat64(0.02),
}}
svc := newPlazaServiceWithBilling(channels, groups, map[int64]string{10: PlatformOpenAI}, nil)
out, err := svc.ListGroups(context.Background())
require.NoError(t, err)
m := out[0].Models[0]
require.Equal(t, BillingModeImage, m.Pricing.BillingMode)
require.Empty(t, m.LongContextBasis)
require.Len(t, m.Pricing.Intervals, 3)
require.InDelta(t, 0.02, *m.Pricing.Intervals[0].PerRequestPrice, 1e-12)
require.InDelta(t, 0.04, *m.Pricing.Intervals[1].PerRequestPrice, 1e-12)
}
func TestListGroups_CatalogMissingStillShowsChannelFlatPricing(t *testing.T) {
// 目录查不到的模型:计费按渠道平价(未配置项 $0),广场单档展示渠道平价,官方价为空。
channels := []Channel{plazaPricedChannel(1, "ch", []int64{10}, PlatformAnthropic, "unknown-model-xyz")}
groups := []Group{{ID: 10, Name: "g", Platform: PlatformAnthropic, RateMultiplier: 1, LongContextPricingEnabled: true}}
svc := newPlazaServiceWithBilling(channels, groups, map[int64]string{10: PlatformAnthropic}, nil)
out, err := svc.ListGroups(context.Background())
require.NoError(t, err)
m := out[0].Models[0]
require.NotNil(t, m.Pricing)
require.InDelta(t, 3e-6, *m.Pricing.InputPrice, 1e-15)
require.Empty(t, m.Pricing.Intervals)
require.Nil(t, m.Pricing.CacheWritePrice, "目录无价且渠道未配置 → 无价")
require.Nil(t, m.OfficialPricing)
}
func TestListGroups_TimePricingPassthrough(t *testing.T) {
channels := []Channel{{
ID: 1, Name: "ch", Status: StatusActive, GroupIDs: []int64{10},
ModelPricing: []ChannelModelPricing{{
Platform: PlatformDeepseek, Models: []string{"deepseek-chat"}, BillingMode: BillingModeToken,
InputPrice: testPtrFloat64(0.28e-6), OutputPrice: testPtrFloat64(0.42e-6),
TimePricing: &ChannelTimePricing{Timezone: "Asia/Shanghai", Periods: []ChannelTimePricingPeriod{
{StartTime: "00:30", EndTime: "08:30", Multiplier: 0.5},
}},
}},
}}
groups := []Group{{ID: 10, Name: "cn", Platform: PlatformDeepseek, RateMultiplier: 1, LongContextPricingEnabled: true}}
svc := newPlazaServiceWithBilling(channels, groups, map[int64]string{10: PlatformDeepseek}, nil)
out, err := svc.ListGroups(context.Background())
require.NoError(t, err)
m := out[0].Models[0]
require.NotNil(t, m.TimePricing)
require.Equal(t, "Asia/Shanghai", m.TimePricing.Timezone)
require.Len(t, m.TimePricing.Periods, 1)
require.InDelta(t, 0.5, m.TimePricing.Periods[0].Multiplier, 1e-12)
// 展示单价为标准时段价
require.InDelta(t, 0.28e-6, *m.Pricing.InputPrice, 1e-15)
}
+1
View File
@@ -908,6 +908,7 @@ var ProviderSet = wire.NewSet(
NewChannelService,
wire.Bind(new(ChannelCacheInvalidator), new(*ChannelService)),
NewModelPricingResolver,
NewModelPlazaService,
NewContentModerationService,
NewAffiliateService,
ProvidePaymentConfigService,
+34 -2
View File
@@ -5,9 +5,9 @@
*/
import { apiClient } from './client'
import type { UserSupportedModelPricing } from './channels'
import type { UserPricingInterval, UserSupportedModelPricing } from './channels'
/** LiteLLM 官方参考价(USD per token,字段缺失 = 官方数据未覆盖)。 */
/** 官方参考价(USD per token,与计费目录同源;字段缺失 = 目录未覆盖)。 */
export interface PlazaOfficialPricing {
input_price: number | null
output_price: number | null
@@ -16,13 +16,43 @@ export interface PlazaOfficialPricing {
/** 1h 缓存写入(LiteLLM cache_creation_above_1hr),多数模型缺失。 */
cache_write_1h_price?: number | null
cache_read_price: number | null
/** 官方长上下文阶梯(多档模型才有),不受分组开关影响。 */
intervals?: UserPricingInterval[]
}
/**
* 多档时的计价基准:
* - whole_request:整单按所在档单价计价(目录阶梯、渠道区间);
* - marginal:仅超出阈值的部分按该档单价计价(平台旧规则)。
*/
export type PlazaLongContextBasis = 'whole_request' | 'marginal'
/** 分时倍率时段:配置时区当天 [start_time, end_time) 内整单实付乘 multiplier。 */
export interface PlazaTimePricingPeriod {
start_time: string
end_time: string
multiplier: number
}
/** 计费会生效的分时倍率(仅倍率 ≠ 1 的时段,已按开始时间升序)。 */
export interface PlazaTimePricing {
/** IANA 时区名,如 Asia/Shanghai。 */
timezone: string
/** true 时时段仅周一至周五生效,周末整天按标准价计费。 */
weekdays_only?: boolean
periods: PlazaTimePricingPeriod[]
}
export interface PlazaModel {
name: string
platform: string
/** 实收口径的展示定价:多档时 intervals 为各档绝对单价(已由计费服务折算);均为标准时段价。 */
pricing: UserSupportedModelPricing | null
official_pricing: PlazaOfficialPricing | null
/** 仅多档模型返回。 */
long_context_basis?: PlazaLongContextBasis
/** 仅配置了分时倍率的模型返回。 */
time_pricing?: PlazaTimePricing
}
export interface ModelPlazaGroup {
@@ -43,6 +73,8 @@ export interface ModelPlazaGroup {
/** 生图独立倍率:true 时图片计费模型的实付倍率取 image_rate_multiplier,不取分组/专属倍率。 */
image_rate_independent: boolean
image_rate_multiplier: number
/** 分组是否启用长上下文阶梯计费;false 时实付列只展示最低档,官方阶梯仅供参考。 */
long_context_pricing_enabled: boolean
models: PlazaModel[]
}
@@ -42,6 +42,13 @@
<Icon name="clock" size="xs" class="h-3 w-3" />
{{ peakNote }}
</p>
<p
v-if="longContextNote"
class="mt-1.5 flex items-center gap-1 text-xs text-gray-500 dark:text-dark-400"
>
<Icon name="infoCircle" size="xs" class="h-3 w-3" />
{{ longContextNote }}
</p>
</header>
<!-- 模型价格表:整行(含 hover 底色/分区底色)顶到卡片边缘,左右留白由表格首列/末列的 padding 提供 -->
@@ -54,6 +61,8 @@
:user-rate-multiplier="group.user_rate_multiplier ?? null"
:image-rate-independent="group.image_rate_independent"
:image-rate-multiplier="group.image_rate_multiplier"
:peak-window="peakWindow"
:peak-rate-multiplier="group.peak_rate_multiplier"
/>
<p v-else class="px-5 py-4 text-center text-sm text-gray-400 dark:text-dark-500">
{{ t('modelPlaza.detail.noModels') }}
@@ -81,15 +90,32 @@ const props = defineProps<{
const { t } = useI18n()
const appStore = useAppStore()
const peakNote = computed(() => {
/** 高峰窗口描述(含倍率与服务器时区标注);分组未启用高峰为空串。 */
const peakWindow = computed(() => {
if (!hasPeakRate(props.group)) return ''
const window = formatPeakRateWindow(
return formatPeakRateWindow(
props.group,
serverTimezoneLabel(appStore.cachedPublicSettings?.server_utc_offset)
)
})
const peakNote = computed(() => {
if (!peakWindow.value) return ''
return t('modelPlaza.detail.peakNote', {
window,
window: peakWindow.value,
multiplier: props.group.peak_rate_multiplier
})
})
/**
* 分组关闭了长上下文阶梯、但组内有模型官方带阶梯时提示:实付列只展示基础档,
* 官方阶梯仅供参考。字段缺失(旧后端)不提示。
*/
const longContextNote = computed(() => {
if (props.group.long_context_pricing_enabled !== false) return ''
const hasOfficialLadder = props.group.models.some(
(m) => (m.official_pricing?.intervals?.length ?? 0) > 1
)
return hasOfficialLadder ? t('modelPlaza.detail.longContextDisabledNote') : ''
})
</script>
@@ -1,15 +1,15 @@
<template>
<div class="plaza-pricing-table overflow-x-auto" :style="accentStyle">
<table class="w-full min-w-[860px] table-fixed border-collapse text-sm tabular-nums">
<table class="w-full min-w-[1000px] table-fixed border-collapse text-sm tabular-nums">
<colgroup>
<col class="w-[22%]" />
<col class="w-[10%]" />
<col class="w-[10%]" />
<col class="w-[25%]" />
<col class="w-[11%]" />
<col class="w-[9%]" />
<col class="w-[14%]" />
<col class="w-[10%]" />
<col class="w-[10%]" />
<col class="w-[11%]" />
<col class="w-[8%]" />
<col class="w-[14%]" />
<col class="w-[10%]" />
<col class="w-[8%]" />
</colgroup>
<thead>
<tr
@@ -58,14 +58,25 @@
</thead>
<tbody>
<tr
v-for="m in sortedModels"
:key="`${m.platform}:${m.name}`"
v-for="{ model: m, period, key } in rows"
:key="key"
class="border-b border-gray-100 transition-colors last:border-b-0 hover:bg-gray-50/70 dark:border-dark-800 dark:hover:bg-dark-800/50"
>
<!-- 模型名 + 非 token 计费模式徽章 -->
<!-- 模型名 + 非 token 计费模式徽章;分时时段行额外标注时段 -->
<td class="border-r border-gray-100 py-2.5 pl-5 pr-4 align-middle dark:border-dark-700/60">
<div class="flex flex-wrap items-center gap-1.5">
<span class="font-medium text-gray-900 dark:text-white">{{ m.name }}</span>
<!-- 时段徽章紧跟模型名,其余徽章排在后面,空间不足时先换行的是它们 -->
<span
v-if="period"
class="inline-flex items-center whitespace-nowrap rounded-md bg-gray-100 px-1 py-0.5 font-mono text-[10px] font-medium text-gray-500 dark:bg-dark-700/70 dark:text-dark-300"
:title="timePricingRowHint(m)"
>
<span v-if="m.time_pricing?.weekdays_only" class="mr-1 font-sans">{{
t('modelPlaza.table.timePricingWeekdays')
}}</span>
{{ formatTimeWindow(period) }}
</span>
<span
v-if="platform && m.platform !== platform"
:class="[
@@ -81,10 +92,17 @@
>
{{ billingModeLabel(m) }}
</span>
<span
v-if="m.long_context_basis === 'marginal'"
class="rounded-md bg-gray-100 px-1.5 py-0.5 text-[10px] font-medium text-gray-500 dark:bg-dark-700/70 dark:text-dark-300"
:title="t('modelPlaza.table.tierHintMarginal')"
>
{{ t('modelPlaza.table.marginalBadge') }}
</span>
</div>
</td>
<!-- token 计费:输入 / 输出(阶梯内联)/ 缓存(写/读) -->
<!-- token 计费:输入 / 输出 / 缓存(写/读),有阶梯时每档一行;档位标签只放输入列,其余列按行对齐 -->
<template v-if="billingMode(m) === BILLING_MODE_TOKEN">
<td class="pz-cell px-3 py-2.5 align-middle font-mono font-semibold text-gray-900 dark:text-gray-50">
<template v-if="tokenIntervals(m).length">
@@ -93,11 +111,11 @@
:key="idx"
class="whitespace-nowrap text-xs leading-5"
>
<span class="mr-1 font-sans font-normal text-gray-400 dark:text-dark-500">{{ tierLabel(iv) }}</span>
{{ paidPerMillion(iv.input_price) }}
<span class="mr-1 font-sans font-normal text-gray-400 dark:text-dark-500" :title="tierHint(m)">{{ tierLabel(iv) }}</span>
{{ paidPerMillion(iv.input_price, period) }}
</div>
</template>
<template v-else>{{ paidPerMillion(m.pricing?.input_price) }}</template>
<template v-else>{{ paidPerMillion(m.pricing?.input_price, period) }}</template>
</td>
<td class="pz-cell px-3 py-2.5 align-middle font-mono font-semibold text-gray-900 dark:text-gray-50">
<template v-if="tokenIntervals(m).length">
@@ -105,25 +123,41 @@
v-for="(iv, idx) in tokenIntervals(m)"
:key="idx"
class="whitespace-nowrap text-xs leading-5"
:title="tierHint(m)"
>
<span class="mr-1 font-sans font-normal text-gray-400 dark:text-dark-500">{{ tierLabel(iv) }}</span>
{{ paidPerMillion(iv.output_price) }}
{{ paidPerMillion(iv.output_price, period) }}
</div>
</template>
<template v-else>{{ paidPerMillion(m.pricing?.output_price) }}</template>
<template v-else>{{ paidPerMillion(m.pricing?.output_price, period) }}</template>
</td>
<td class="pz-cell px-3 py-2.5 align-middle">
<template v-if="hasTierCachePricing(tokenIntervals(m))">
<div
v-for="(iv, idx) in tokenIntervals(m)"
:key="idx"
class="whitespace-nowrap font-mono text-xs leading-5 text-gray-800 dark:text-gray-200"
:title="tierHint(m)"
>
<template v-if="iv.cache_write_price != null || iv.cache_read_price != null">
<span class="font-sans font-normal text-gray-400 dark:text-dark-500">{{ t('modelPlaza.table.cacheWriteShort') }}</span>
{{ paidPerMillion(iv.cache_write_price, period) }}
<span class="ml-1 font-sans font-normal text-gray-400 dark:text-dark-500">{{ t('modelPlaza.table.cacheReadShort') }}</span>
{{ paidPerMillion(iv.cache_read_price, period) }}
</template>
<span v-else class="text-gray-400 dark:text-dark-500">-</span>
</div>
</template>
<div
v-if="hasCachePricing(m)"
v-else-if="hasCachePricing(m)"
class="space-y-0.5 font-mono text-xs text-gray-800 dark:text-gray-200"
>
<div>
<span class="mr-1 font-sans font-normal text-gray-400 dark:text-dark-500">{{ t('modelPlaza.table.cacheWrite') }}</span>
{{ paidPerMillion(m.pricing?.cache_write_price) }}
{{ paidPerMillion(m.pricing?.cache_write_price, period) }}
</div>
<div>
<span class="mr-1 font-sans font-normal text-gray-400 dark:text-dark-500">{{ t('modelPlaza.table.cacheRead') }}</span>
{{ paidPerMillion(m.pricing?.cache_read_price) }}
{{ paidPerMillion(m.pricing?.cache_read_price, period) }}
</div>
</div>
<span v-else class="text-gray-400 dark:text-dark-500">-</span>
@@ -157,18 +191,54 @@
</td>
</template>
<!-- 官方价格(LiteLLM 参考价,不乘倍率) -->
<!-- 官方价格(参考价,不乘倍率;官方有阶梯时每档一行) -->
<td
class="border-l border-gray-100 px-3 py-2.5 align-middle font-mono text-xs text-gray-500 dark:border-dark-700/60 dark:text-dark-400"
>
{{ official(m.official_pricing?.input_price) }}
<template v-if="officialIntervals(m).length">
<div
v-for="(iv, idx) in officialIntervals(m)"
:key="idx"
class="whitespace-nowrap leading-5"
>
<span class="mr-1 font-sans text-gray-400 dark:text-dark-500" :title="t('modelPlaza.table.tierHint')">{{ tierLabel(iv) }}</span>
{{ official(iv.input_price) }}
</div>
</template>
<template v-else>{{ official(m.official_pricing?.input_price) }}</template>
</td>
<td class="px-3 py-2.5 align-middle font-mono text-xs text-gray-500 dark:text-dark-400">
{{ official(m.official_pricing?.output_price) }}
<template v-if="officialIntervals(m).length">
<div
v-for="(iv, idx) in officialIntervals(m)"
:key="idx"
class="whitespace-nowrap leading-5"
:title="t('modelPlaza.table.tierHint')"
>
{{ official(iv.output_price) }}
</div>
</template>
<template v-else>{{ official(m.official_pricing?.output_price) }}</template>
</td>
<td class="px-3 py-2.5 align-middle">
<template v-if="hasTierCachePricing(officialIntervals(m))">
<div
v-for="(iv, idx) in officialIntervals(m)"
:key="idx"
class="whitespace-nowrap font-mono text-xs leading-5 text-gray-500 dark:text-dark-400"
:title="t('modelPlaza.table.tierHint')"
>
<template v-if="iv.cache_write_price != null || iv.cache_read_price != null">
<span class="font-sans text-gray-400 dark:text-dark-500">{{ t('modelPlaza.table.cacheWriteShort') }}</span>
{{ official(iv.cache_write_price) }}
<span class="ml-1 font-sans text-gray-400 dark:text-dark-500">{{ t('modelPlaza.table.cacheReadShort') }}</span>
{{ official(iv.cache_read_price) }}
</template>
<span v-else class="text-gray-400 dark:text-dark-500">-</span>
</div>
</template>
<div
v-if="m.official_pricing && hasOfficialCache(m.official_pricing)"
v-else-if="m.official_pricing && hasOfficialCache(m.official_pricing)"
class="space-y-0.5 font-mono text-xs text-gray-500 dark:text-dark-400"
>
<div>
@@ -187,12 +257,18 @@
<span v-else class="text-gray-400 dark:text-dark-500">-</span>
</td>
<!-- 折扣倍率(生图独立倍率行展示独立倍率;专属倍率划线展示原倍率) -->
<!-- 折扣倍率(分时时段行展示 生效倍率×时段倍率;生图独立倍率行展示独立倍率;专属倍率划线展示原倍率) -->
<td
class="border-l border-gray-100 py-2.5 pl-3 pr-5 text-right align-middle font-mono text-xs dark:border-dark-700/60"
>
<span
v-if="usesIndependentImageRate(m)"
v-if="period"
class="font-bold text-primary-600 dark:text-primary-400"
:title="t('modelPlaza.table.timePricingRateHint', { rate: effectiveRate, multiplier: period.multiplier })"
>{{ periodRate(period) }}x</span
>
<span
v-else-if="usesIndependentImageRate(m)"
class="font-bold text-gray-700 dark:text-gray-300"
>{{ requestRate(m) }}x</span
>
@@ -218,7 +294,7 @@ import {
BILLING_MODE_IMAGE,
type BillingMode
} from '@/constants/channel'
import type { PlazaModel } from '@/api/modelPlaza'
import type { PlazaModel, PlazaTimePricingPeriod } from '@/api/modelPlaza'
import type { UserPricingInterval } from '@/api/channels'
const props = defineProps<{
@@ -232,6 +308,13 @@ const props = defineProps<{
/** 生图独立倍率:true 时图片计费模型的实付倍率取 imageRateMultiplier,不取分组/专属倍率。 */
imageRateIndependent?: boolean
imageRateMultiplier?: number | null
/**
* 高峰窗口描述(含倍率与服务器时区标注),空串/缺省 = 分组未启用高峰。
* 表格所有价格均为不含高峰因子的口径,该窗口仅用于分时时段行的 tooltip 披露:
* 与高峰重叠的部分实付还会再乘高峰倍率。
*/
peakWindow?: string
peakRateMultiplier?: number | null
}>()
const { t } = useI18n()
@@ -279,10 +362,35 @@ function billingModeLabel(m: PlazaModel): string {
/** 价格统一保底 2 位小数,更长的有效小数原样保留。 */
const MIN_DECIMALS = 2
/** 实付价 = 渠道单价 × 生效倍率,按 $/1M token 展示。 */
function paidPerMillion(value: number | null | undefined): string {
/** 表格行:每个模型一行标准价;配置了分时倍率的模型再按时段各加一行。 */
interface PlazaRow {
model: PlazaModel
period: PlazaTimePricingPeriod | null
key: string
}
const rows = computed<PlazaRow[]>(() =>
sortedModels.value.flatMap((m) => {
const base: PlazaRow = { model: m, period: null, key: `${m.platform}:${m.name}` }
const periodRows = timePeriods(m).map<PlazaRow>((p, idx) => ({
model: m,
period: p,
key: `${m.platform}:${m.name}:${idx}`
}))
return [base, ...periodRows]
})
)
/** 时段行的生效倍率 = 生效倍率 × 时段倍率(去掉浮点噪声)。 */
function periodRate(period: PlazaTimePricingPeriod): number {
return Math.round(effectiveRate.value * period.multiplier * 1000) / 1000
}
/** 实付价 = 渠道单价 × 生效倍率(时段行再乘时段倍率),按 $/1M token 展示。 */
function paidPerMillion(value: number | null | undefined, period: PlazaTimePricingPeriod | null = null): string {
if (value == null) return '-'
return formatScaled(value * effectiveRate.value, PER_MILLION, MIN_DECIMALS)
const rate = period ? periodRate(period) : effectiveRate.value
return formatScaled(value * rate, PER_MILLION, MIN_DECIMALS)
}
/** 图片计费模型且分组开启生图独立倍率:实付倍率取独立倍率,与计费口径一致。 */
@@ -322,23 +430,76 @@ function hasOfficialCache(o: NonNullable<PlazaModel['official_pricing']>): boole
return o.cache_write_price != null || o.cache_read_price != null || o.cache_write_1h_price != null
}
/** token 模式的阶梯定价(内联进输入/输出列)。 */
function tokenIntervals(m: PlazaModel): UserPricingInterval[] {
return m.pricing?.intervals ?? []
/** 分时倍率时段(后端只给出倍率 ≠ 1 的时段,已升序)。 */
function timePeriods(m: PlazaModel): PlazaTimePricingPeriod[] {
return m.time_pricing?.periods ?? []
}
/**
* 时段行 tooltip:仅工作日生效的配置换用带周末回落说明的文案;
* 分组启用高峰倍率时追加披露——本行价格不含高峰因子,与高峰窗口重叠的部分实付再乘高峰倍率。
*/
function timePricingRowHint(m: PlazaModel): string {
const key = m.time_pricing?.weekdays_only
? 'modelPlaza.table.timePricingRowHintWeekdays'
: 'modelPlaza.table.timePricingRowHint'
let hint = t(key, { timezone: m.time_pricing?.timezone })
if (props.peakWindow) {
hint += t('modelPlaza.table.timePricingRowHintPeak', {
window: props.peakWindow,
multiplier: props.peakRateMultiplier ?? 1
})
}
return hint
}
/** “00:30–08:30”;整分钟的 HH:mm:ss 省略秒。 */
function formatTimeWindow(p: PlazaTimePricingPeriod): string {
const clock = (v: string) => v.replace(/^(\d{2}:\d{2}):00$/, '$1')
return `${clock(p.start_time)}–${clock(p.end_time)}`
}
/** 上下文档位按下限升序展示(后端已升序,此处兜底)。 */
function sortByContext(intervals: UserPricingInterval[]): UserPricingInterval[] {
return [...intervals].sort((a, b) => a.min_tokens - b.min_tokens)
}
/** token 模式的阶梯定价(内联进输入/输出/缓存列)。 */
function tokenIntervals(m: PlazaModel): UserPricingInterval[] {
return sortByContext(m.pricing?.intervals ?? [])
}
/** 官方阶梯(后端按目录规则合成,不受分组开关影响)。 */
function officialIntervals(m: PlazaModel): UserPricingInterval[] {
return sortByContext(m.official_pricing?.intervals ?? [])
}
/** 任一档带缓存价才按档渲染缓存列;否则沿用平价的写入/读取两行。 */
function hasTierCachePricing(intervals: UserPricingInterval[]): boolean {
return intervals.some((iv) => iv.cache_write_price != null || iv.cache_read_price != null)
}
/** 档位说明:整单按档计价,或(平台旧规则)仅超出部分按档计价。 */
function tierHint(m: PlazaModel): string {
return m.long_context_basis === 'marginal'
? t('modelPlaza.table.tierHintMarginal')
: t('modelPlaza.table.tierHint')
}
/** 按次/按图模式的阶梯定价(仅保留配了按次价的档位)。 */
function requestIntervals(m: PlazaModel): UserPricingInterval[] {
return (m.pricing?.intervals ?? []).filter((iv) => iv.per_request_price != null)
}
/** 档位标签:优先管理员配置的 tier_label,否则按 token 区间生成(≤200K / >200K / 200K–1M)。 */
/**
* 档位标签:优先后端/管理员给出的 tier_label,否则按区间生成统一形态——
* 有上限为「≤上限」,末档为「>下限」;档位升序排列,相邻的 ≤100K / ≤200K 即表示 (100K,200K]。
*/
function tierLabel(iv: UserPricingInterval): string {
if (iv.tier_label) return iv.tier_label
const { min_tokens: min, max_tokens: max } = iv
if (max == null) return `>${formatTokenCount(min)}`
if (min === 0) return `≤${formatTokenCount(max)}`
return `${formatTokenCount(min)}–${formatTokenCount(max)}`
return max == null ? `>${formatTokenCount(min)}` : `≤${formatTokenCount(max)}`
}
function formatTokenCount(n: number): string {
@@ -0,0 +1,139 @@
import { describe, expect, it, vi } from 'vitest'
import { mount } from '@vue/test-utils'
import PlazaGroupSection from '../PlazaGroupSection.vue'
import PlazaModelPricingTable from '../PlazaModelPricingTable.vue'
import type { ModelPlazaGroup, PlazaModel } from '@/api/modelPlaza'
vi.mock('vue-i18n', async () => {
const actual = await vi.importActual<typeof import('vue-i18n')>('vue-i18n')
return {
...actual,
useI18n: () => ({
t: (key: string) => key
})
}
})
vi.mock('@/stores/app', () => ({
useAppStore: () => ({ cachedPublicSettings: null })
}))
function ladderModel(tiers: number): PlazaModel {
const intervals = Array.from({ length: tiers }, (_, i) => ({
min_tokens: i * 272000,
max_tokens: i === tiers - 1 ? null : (i + 1) * 272000,
tier_label: '',
input_price: 5e-6,
output_price: 3e-5,
cache_write_price: null,
cache_read_price: null,
per_request_price: null
}))
return {
name: 'gpt-5.6-sol',
platform: 'openai',
pricing: {
billing_mode: 'token',
input_price: 5e-6,
output_price: 3e-5,
cache_write_price: null,
cache_read_price: null,
image_input_price: null,
image_output_price: null,
per_request_price: null,
intervals: []
},
official_pricing: {
input_price: 5e-6,
output_price: 3e-5,
cache_write_price: null,
cache_read_price: null,
intervals
}
}
}
function group(overrides: Partial<ModelPlazaGroup> = {}): ModelPlazaGroup {
return {
id: 1,
name: 'g',
description: '',
platform: 'openai',
subscription_type: 'standard',
rate_multiplier: 1,
peak_rate_enabled: false,
peak_start: '',
peak_end: '',
peak_rate_multiplier: 1,
is_exclusive: false,
image_rate_independent: false,
image_rate_multiplier: 1,
long_context_pricing_enabled: true,
models: [ladderModel(2)],
...overrides
}
}
function mountSection(g: ModelPlazaGroup) {
return mount(PlazaGroupSection, {
props: { group: g },
global: {
stubs: {
GroupBadge: true,
Icon: true,
PlazaModelPricingTable: true
}
}
})
}
const NOTE = 'modelPlaza.detail.longContextDisabledNote'
describe('PlazaGroupSection 长上下文说明', () => {
it('分组关闭阶梯且组内有官方阶梯模型时显示说明', () => {
const wrapper = mountSection(group({ long_context_pricing_enabled: false }))
expect(wrapper.text()).toContain(NOTE)
})
it('分组开启阶梯时不显示', () => {
const wrapper = mountSection(group({ long_context_pricing_enabled: true }))
expect(wrapper.text()).not.toContain(NOTE)
})
it('分组关闭但没有官方阶梯模型时不显示', () => {
const wrapper = mountSection(
group({ long_context_pricing_enabled: false, models: [ladderModel(1)] })
)
expect(wrapper.text()).not.toContain(NOTE)
})
it('旧后端缺少开关字段时不显示', () => {
const g = group()
delete (g as Partial<ModelPlazaGroup>).long_context_pricing_enabled
const wrapper = mountSection(g)
expect(wrapper.text()).not.toContain(NOTE)
})
})
describe('PlazaGroupSection 高峰配置传递', () => {
it('分组启用高峰时把窗口描述与倍率传给价格表', () => {
const wrapper = mountSection(
group({
subscription_type: 'subscription',
peak_rate_enabled: true,
peak_start: '14:00',
peak_end: '18:00',
peak_rate_multiplier: 1.5
})
)
const table = wrapper.findComponent(PlazaModelPricingTable)
// appStore mock 无 server_utc_offset,窗口描述不带时区标注
expect(table.props('peakWindow')).toBe('14:00-18:00 ×1.5')
expect(table.props('peakRateMultiplier')).toBe(1.5)
})
it('分组未启用高峰时窗口描述为空串', () => {
const wrapper = mountSection(group())
expect(wrapper.findComponent(PlazaModelPricingTable).props('peakWindow')).toBe('')
})
})
@@ -43,7 +43,12 @@ function mountTable(
models: PlazaModel[],
rateMultiplier: number,
userRateMultiplier?: number | null,
extraProps?: { imageRateIndependent?: boolean; imageRateMultiplier?: number | null }
extraProps?: {
imageRateIndependent?: boolean
imageRateMultiplier?: number | null
peakWindow?: string
peakRateMultiplier?: number | null
}
) {
return mount(PlazaModelPricingTable, {
props: { models, rateMultiplier, userRateMultiplier: userRateMultiplier ?? null, ...extraProps }
@@ -405,3 +410,225 @@ describe('PlazaModelPricingTable', () => {
expect(wrapper.text()).toContain('OpenAI')
})
})
describe('PlazaModelPricingTable 长上下文阶梯', () => {
function ladderIntervals() {
return [
{
min_tokens: 0,
max_tokens: 272000,
tier_label: '≤272K',
input_price: 5e-6,
output_price: 3e-5,
cache_write_price: 6.25e-6,
cache_read_price: 5e-7,
per_request_price: null
},
{
min_tokens: 272000,
max_tokens: null,
tier_label: '>272K',
input_price: 1e-5,
output_price: 4.5e-5,
cache_write_price: 1.25e-5,
cache_read_price: 1e-6,
per_request_price: null
}
]
}
function ladderModel(overrides: Partial<PlazaModel> = {}): PlazaModel {
return tokenModel({
name: 'gpt-5.6-sol',
platform: 'openai',
pricing: {
billing_mode: 'token',
input_price: 5e-6,
output_price: 3e-5,
cache_write_price: 6.25e-6,
cache_read_price: 5e-7,
image_input_price: null,
image_output_price: null,
per_request_price: null,
intervals: ladderIntervals()
},
official_pricing: {
input_price: 5e-6,
output_price: 3e-5,
cache_write_price: 6.25e-6,
cache_read_price: 5e-7,
intervals: ladderIntervals()
},
long_context_basis: 'whole_request',
...overrides
})
}
it('实付缓存列按档分行并乘倍率,每档一行与输入/输出列对齐;档位标签只在输入列', () => {
const wrapper = mountTable([ladderModel()], 0.5)
const cells = wrapper.findAll('tbody td')
const cacheCell = cells[3]
const rows = cacheCell.findAll('.leading-5')
expect(rows).toHaveLength(2)
// 写 6.25 × 0.5 / 读 0.5 × 0.5;高档 12.5 × 0.5 / 1 × 0.5
expect(rows[0].text()).toContain('modelPlaza.table.cacheWriteShort')
expect(rows[0].text()).toContain('$3.125')
expect(rows[0].text()).toContain('$0.25')
expect(rows[1].text()).toContain('$6.25')
expect(rows[1].text()).toContain('$0.50')
// 输入列带标签,输出/缓存列只按行对齐不重复标签
expect(cells[1].text()).toContain('≤272K')
expect(cells[1].text()).toContain('>272K')
expect(cells[2].text()).not.toContain('272K')
expect(cacheCell.text()).not.toContain('272K')
expect(cells[1].findAll('.leading-5')).toHaveLength(2)
expect(cells[2].findAll('.leading-5')).toHaveLength(2)
})
it('官方三列按 official_pricing.intervals 分档且不乘倍率,不内联 1h', () => {
const wrapper = mountTable([ladderModel()], 0.5)
const cells = wrapper.findAll('tbody td')
expect(cells[4].text()).toContain('≤272K')
expect(cells[4].text()).toContain('$5.00')
expect(cells[4].text()).toContain('>272K')
expect(cells[4].text()).toContain('$10.00')
expect(cells[5].text()).toContain('$30.00')
expect(cells[5].text()).toContain('$45.00')
expect(cells[6].text()).toContain('$6.25')
expect(cells[6].text()).toContain('$12.50')
expect(cells[6].text()).toContain('$1.00')
expect(cells[6].text()).not.toContain('(1h')
})
it('整单计价的档位标签带 tooltip;边际计价在模型名旁加徽章并换用边际说明', () => {
const whole = mountTable([ladderModel()], 1)
const wholeLabels = whole.findAll('tbody td span[title="modelPlaza.table.tierHint"]')
expect(wholeLabels.length).toBeGreaterThan(0)
expect(whole.text()).not.toContain('modelPlaza.table.marginalBadge')
const marginal = mountTable([ladderModel({ long_context_basis: 'marginal' })], 1)
const marginalLabels = marginal.findAll('tbody td span[title="modelPlaza.table.tierHintMarginal"]')
expect(marginalLabels.length).toBeGreaterThan(0)
expect(marginal.findAll('tbody td')[0].text()).toContain('modelPlaza.table.marginalBadge')
})
it('无标签的多档按区间生成统一形态(≤上限 / >下限),并按下限升序展示', () => {
const model = ladderModel({
pricing: {
...ladderModel().pricing!,
// 故意乱序:展示必须按上下文从低到高
intervals: [
{ ...ladderIntervals()[1], min_tokens: 1000000, tier_label: '' },
{ ...ladderIntervals()[0], min_tokens: 100000, max_tokens: 200000, tier_label: '' },
{ ...ladderIntervals()[0], max_tokens: 100000, tier_label: '' },
{ ...ladderIntervals()[1], min_tokens: 200000, max_tokens: 1000000, tier_label: '' }
]
}
})
const rows = mountTable([model], 1).findAll('tbody td')[1].findAll('.leading-5')
expect(rows.map((r) => r.text().split(/\s+/)[0])).toEqual(['≤100K', '≤200K', '≤1M', '>1M'])
})
it('官方无 intervals 字段(旧响应)时官方列保持平价,实付无阶梯时缓存列保持两行', () => {
const wrapper = mountTable([tokenModel()], 1)
const cells = wrapper.findAll('tbody td')
expect(cells[3].text()).toContain('modelPlaza.table.cacheWrite')
expect(cells[3].text()).toContain('modelPlaza.table.cacheRead')
expect(cells[3].findAll('.leading-5')).toHaveLength(0)
expect(cells[6].text()).toContain('(1h')
})
})
describe('PlazaModelPricingTable 分时计价', () => {
function timePricedModel() {
return tokenModel({
name: 'deepseek-chat',
platform: 'deepseek',
time_pricing: {
timezone: 'Asia/Shanghai',
periods: [
{ start_time: '00:30', end_time: '08:30:00', multiplier: 0.5 },
{ start_time: '18:00', end_time: '22:00', multiplier: 1.2 }
]
}
})
}
it('有分时倍率的模型展开为标准行 + 每时段一行,时段行价格按倍率折算且倍率列显示生效倍率', () => {
const wrapper = mountTable([timePricedModel()], 0.8)
const trs = wrapper.findAll('tbody tr')
expect(trs).toHaveLength(3)
// 标准行:输入 3 × 0.8
const baseCells = trs[0].findAll('td')
expect(baseCells[0].text()).toBe('deepseek-chat')
expect(baseCells[1].text()).toContain('$2.40')
expect(baseCells[7].text()).toContain('0.8x')
// 夜间时段行:输入 3 × 0.8 × 0.5,倍率 0.4x,标注时段不含时区
const nightCells = trs[1].findAll('td')
expect(nightCells[0].text()).toContain('deepseek-chat')
expect(nightCells[0].text()).toContain('00:30–08:30')
expect(nightCells[0].text()).not.toContain('Asia/Shanghai')
// 时区只放在 tooltip 里(i18n mock 不做插值,这里只断言挂了说明)
expect(nightCells[0].find('[title="modelPlaza.table.timePricingRowHint"]').exists()).toBe(true)
expect(nightCells[1].text()).toContain('$1.20')
expect(nightCells[2].text()).toContain('$6.00')
expect(nightCells[3].text()).toContain('$1.50')
expect(nightCells[7].text()).toContain('0.4x')
// 晚高峰行:3 × 0.8 × 1.2 = 2.88,倍率 0.96x
const peakCells = trs[2].findAll('td')
expect(peakCells[0].text()).toContain('18:00–22:00')
expect(peakCells[1].text()).toContain('$2.88')
expect(peakCells[7].text()).toContain('0.96x')
// 官方列不受时段影响
expect(nightCells[4].text()).toContain('$3.00')
})
it('仅工作日生效时时段行带工作日前缀,tooltip 换用周末回落文案', () => {
const model = timePricedModel()
model.time_pricing!.weekdays_only = true
const wrapper = mountTable([model], 1)
const trs = wrapper.findAll('tbody tr')
expect(trs).toHaveLength(3)
const nightCells = trs[1].findAll('td')
expect(nightCells[0].text()).toContain('modelPlaza.table.timePricingWeekdays')
expect(nightCells[0].text()).toContain('00:30–08:30')
expect(nightCells[0].find('[title="modelPlaza.table.timePricingRowHintWeekdays"]').exists()).toBe(true)
expect(nightCells[0].find('[title="modelPlaza.table.timePricingRowHint"]').exists()).toBe(false)
})
it('每日生效(无 weekdays_only)不渲染工作日前缀', () => {
const wrapper = mountTable([timePricedModel()], 1)
expect(wrapper.find('tbody').text()).not.toContain('modelPlaza.table.timePricingWeekdays')
expect(wrapper.find('[title="modelPlaza.table.timePricingRowHint"]').exists()).toBe(true)
})
it('分组启用高峰倍率时时段行 tooltip 追加高峰披露,价格与倍率列保持不含高峰的口径', () => {
const wrapper = mountTable([timePricedModel()], 0.8, null, {
peakWindow: '14:00-18:00 ×1.5 (UTC+08:00)',
peakRateMultiplier: 1.5
})
const nightCells = wrapper.findAll('tbody tr')[1].findAll('td')
const title = nightCells[0].find('[title*="modelPlaza.table.timePricingRowHint"]').attributes('title')
expect(title).toContain('modelPlaza.table.timePricingRowHintPeak')
// 行内数字仍是 基础倍率 × 时段倍率(0.8 × 0.5),高峰只进披露不进价格
expect(nightCells[1].text()).toContain('$1.20')
expect(nightCells[7].text()).toContain('0.4x')
})
it('分组未启用高峰(peakWindow 缺省)时 tooltip 不含高峰披露', () => {
const wrapper = mountTable([timePricedModel()], 1)
const badge = wrapper.find('[title*="modelPlaza.table.timePricingRowHint"]')
expect(badge.attributes('title')).not.toContain('modelPlaza.table.timePricingRowHintPeak')
})
it('无分时倍率时只有一行,不渲染时段标注', () => {
const wrapper = mountTable([tokenModel()], 1)
expect(wrapper.findAll('tbody tr')).toHaveLength(1)
expect(wrapper.find('[title*="modelPlaza.table.timePricingRowHint"]').exists()).toBe(false)
})
})
+14 -1
View File
@@ -589,7 +589,8 @@ export default {
detail: {
noModels: 'No models configured for this group',
noPricing: 'Pricing not configured',
peakNote: 'Peak hours {window}: billing rate ×{multiplier}'
peakNote: 'Peak hours {window}: billing rate ×{multiplier}',
longContextDisabledNote: 'Long-context tier pricing is disabled for this group: requests above the threshold are billed at the base tier; official tiers are for reference only'
},
table: {
model: 'Model',
@@ -598,6 +599,18 @@ export default {
cache: 'Cache',
cacheWrite: 'Write',
cacheRead: 'Read',
cacheWriteShort: 'W',
cacheReadShort: 'R',
tierHint: 'The whole request is billed at the tier matching its total context (input + cache write + cache read)',
tierHintMarginal: 'Only the portion above the threshold is billed at this tier; output is unaffected',
marginalBadge: 'excess-only tiers',
timePricingRowHint: 'Requests made within this period ({timezone} time) are billed at the prices in this row',
timePricingRowHintWeekdays:
'On weekdays (Mon–Fri) only, requests made within this period ({timezone} time) are billed at the prices in this row; weekends use the standard prices',
timePricingRowHintPeak:
'; prices in this row exclude the peak-hour rate — where this period overlaps the peak hours {window}, the overlapping portion is additionally multiplied by ×{multiplier}',
timePricingWeekdays: 'Weekdays',
timePricingRateHint: 'Effective rate {rate} × period multiplier {multiplier}',
paidPrice: 'Your Price (Discounted)',
officialPrice: 'Official Price',
rate: 'Rate',
+13 -1
View File
@@ -594,7 +594,8 @@ export default {
detail: {
noModels: '该分组暂未配置模型',
noPricing: '未配置定价',
peakNote: '高峰时段 {window} 计费倍率 ×{multiplier}'
peakNote: '高峰时段 {window} 计费倍率 ×{multiplier}',
longContextDisabledNote: '该分组未启用长上下文阶梯计费,超阈值请求仍按基础档计费,官方阶梯仅供参考'
},
table: {
model: '模型',
@@ -603,6 +604,17 @@ export default {
cache: '缓存',
cacheWrite: '写入',
cacheRead: '读取',
cacheWriteShort: '写',
cacheReadShort: '读',
tierHint: '按单次请求的总上下文(输入 + 缓存写入 + 缓存读取)所在档位对整单计价',
tierHintMarginal: '仅超过阈值的部分按该档计价,输出不加价',
marginalBadge: '超出部分计价',
timePricingRowHint: '按 {timezone} 时间,在该时段内发起的请求按本行价格计费',
timePricingRowHintWeekdays:
'按 {timezone} 时间,仅工作日(周一至周五)在该时段内发起的请求按本行价格计费,周末全天按标准价',
timePricingRowHintPeak: ';本行价格未含高峰倍率,与高峰时段 {window} 重叠的部分实付再乘 ×{multiplier}',
timePricingWeekdays: '工作日',
timePricingRateHint: '生效倍率 {rate} × 时段倍率 {multiplier}',
paidPrice: '实付价格(折后)',
officialPrice: '官方价格',
rate: '折扣倍率',