mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:28:39 +08:00
Merge pull request #6109 from feeeei/main
feat(model-plaza): 模型广场增加长上下文阶梯计价显示 & 分时段计价显示
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
+193
-28
@@ -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)
|
||||
}
|
||||
@@ -908,6 +908,7 @@ var ProviderSet = wire.NewSet(
|
||||
NewChannelService,
|
||||
wire.Bind(new(ChannelCacheInvalidator), new(*ChannelService)),
|
||||
NewModelPricingResolver,
|
||||
NewModelPlazaService,
|
||||
NewContentModerationService,
|
||||
NewAffiliateService,
|
||||
ProvidePaymentConfigService,
|
||||
|
||||
@@ -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)
|
||||
})
|
||||
})
|
||||
|
||||
@@ -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',
|
||||
|
||||
@@ -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: '折扣倍率',
|
||||
|
||||
Reference in New Issue
Block a user