diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index 1300989dc6..a51e13bb9d 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -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) diff --git a/backend/internal/handler/available_channel_handler.go b/backend/internal/handler/available_channel_handler.go index 300eb1b1a7..a6b7a36a3c 100644 --- a/backend/internal/handler/available_channel_handler.go +++ b/backend/internal/handler/available_channel_handler.go @@ -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) diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go index 2a270d3a0a..07840b7ea4 100644 --- a/backend/internal/handler/gemini_v1beta_handler.go +++ b/backend/internal/handler/gemini_v1beta_handler.go @@ -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, diff --git a/backend/internal/handler/model_plaza_handler.go b/backend/internal/handler/model_plaza_handler.go index 0bf294d443..829ef94af4 100644 --- a/backend/internal/handler/model_plaza_handler.go +++ b/backend/internal/handler/model_plaza_handler.go @@ -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), } } diff --git a/backend/internal/handler/model_plaza_handler_test.go b/backend/internal/handler/model_plaza_handler_test.go index a7fc291ab8..ab37a6db24 100644 --- a/backend/internal/handler/model_plaza_handler_test.go +++ b/backend/internal/handler/model_plaza_handler_test.go @@ -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"]) +} diff --git a/backend/internal/service/billing_context_schedule.go b/backend/internal/service/billing_context_schedule.go new file mode 100644 index 0000000000..05f2a181b1 --- /dev/null +++ b/backend/internal/service/billing_context_schedule.go @@ -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) +} diff --git a/backend/internal/service/billing_context_schedule_test.go b/backend/internal/service/billing_context_schedule_test.go new file mode 100644 index 0000000000..b13606da9f --- /dev/null +++ b/backend/internal/service/billing_context_schedule_test.go @@ -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) + }) +} diff --git a/backend/internal/service/billing_token_cost_request.go b/backend/internal/service/billing_token_cost_request.go new file mode 100644 index 0000000000..efd8aa6606 --- /dev/null +++ b/backend/internal/service/billing_token_cost_request.go @@ -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 +} diff --git a/backend/internal/service/billing_token_cost_request_test.go b/backend/internal/service/billing_token_cost_request_test.go new file mode 100644 index 0000000000..8392b8cf49 --- /dev/null +++ b/backend/internal/service/billing_token_cost_request_test.go @@ -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) +} diff --git a/backend/internal/service/channel_available.go b/backend/internal/service/channel_available.go index eeaf7dac2e..6face6e1db 100644 --- a/backend/internal/service/channel_available.go +++ b/backend/internal/service/channel_available.go @@ -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 } diff --git a/backend/internal/service/channel_available_test.go b/backend/internal/service/channel_available_test.go index d59e587ecd..2b7b23f7e5 100644 --- a/backend/internal/service/channel_available_test.go +++ b/backend/internal/service/channel_available_test.go @@ -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) } diff --git a/backend/internal/service/channel_plaza.go b/backend/internal/service/channel_plaza.go deleted file mode 100644 index 1e4dee2717..0000000000 --- a/backend/internal/service/channel_plaza.go +++ /dev/null @@ -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 -} diff --git a/backend/internal/service/gateway_usage_billing.go b/backend/internal/service/gateway_usage_billing.go index 0a770bc7a3..1832dd6d32 100644 --- a/backend/internal/service/gateway_usage_billing.go +++ b/backend/internal/service/gateway_usage_billing.go @@ -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, diff --git a/backend/internal/service/model_plaza_service.go b/backend/internal/service/model_plaza_service.go new file mode 100644 index 0000000000..fe2252c594 --- /dev/null +++ b/backend/internal/service/model_plaza_service.go @@ -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 +} diff --git a/backend/internal/service/channel_plaza_test.go b/backend/internal/service/model_plaza_service_test.go similarity index 53% rename from backend/internal/service/channel_plaza_test.go rename to backend/internal/service/model_plaza_service_test.go index 554654149c..d54aca70ad 100644 --- a/backend/internal/service/channel_plaza_test.go +++ b/backend/internal/service/model_plaza_service_test.go @@ -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) +} diff --git a/backend/internal/service/wire.go b/backend/internal/service/wire.go index 3e469a95c5..17a68799c4 100644 --- a/backend/internal/service/wire.go +++ b/backend/internal/service/wire.go @@ -908,6 +908,7 @@ var ProviderSet = wire.NewSet( NewChannelService, wire.Bind(new(ChannelCacheInvalidator), new(*ChannelService)), NewModelPricingResolver, + NewModelPlazaService, NewContentModerationService, NewAffiliateService, ProvidePaymentConfigService, diff --git a/frontend/src/api/modelPlaza.ts b/frontend/src/api/modelPlaza.ts index c17b3aa1a7..e96a621934 100644 --- a/frontend/src/api/modelPlaza.ts +++ b/frontend/src/api/modelPlaza.ts @@ -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[] } diff --git a/frontend/src/components/modelPlaza/PlazaGroupSection.vue b/frontend/src/components/modelPlaza/PlazaGroupSection.vue index 3a46c93aaf..53a9c77cf9 100644 --- a/frontend/src/components/modelPlaza/PlazaGroupSection.vue +++ b/frontend/src/components/modelPlaza/PlazaGroupSection.vue @@ -42,6 +42,13 @@ {{ peakNote }}

+

+ + {{ longContextNote }} +

@@ -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" />

{{ 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') : '' +}) diff --git a/frontend/src/components/modelPlaza/PlazaModelPricingTable.vue b/frontend/src/components/modelPlaza/PlazaModelPricingTable.vue index edc72365ff..d4aac37a07 100644 --- a/frontend/src/components/modelPlaza/PlazaModelPricingTable.vue +++ b/frontend/src/components/modelPlaza/PlazaModelPricingTable.vue @@ -1,15 +1,15 @@