mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 11:43:23 +08:00
Merge pull request #5851 from IanShaw027/feat/channel-pricing-tier-multipliers
fix(billing): restore Fast tier pricing under channel overrides, add configurable multipliers
This commit is contained in:
@@ -64,6 +64,8 @@ type channelModelPricingRequest struct {
|
||||
OutputPrice *float64 `json:"output_price" binding:"omitempty,min=0"`
|
||||
CacheWritePrice *float64 `json:"cache_write_price" binding:"omitempty,min=0"`
|
||||
CacheReadPrice *float64 `json:"cache_read_price" binding:"omitempty,min=0"`
|
||||
FastMultiplier *float64 `json:"fast_multiplier" binding:"omitempty,gt=0"`
|
||||
FlexMultiplier *float64 `json:"flex_multiplier" binding:"omitempty,gt=0"`
|
||||
ImageInputPrice *float64 `json:"image_input_price" binding:"omitempty,min=0"`
|
||||
ImageOutputPrice *float64 `json:"image_output_price" binding:"omitempty,min=0"`
|
||||
PerRequestPrice *float64 `json:"per_request_price" binding:"omitempty,min=0"`
|
||||
@@ -83,15 +85,19 @@ type channelTimePricingPeriodRequest struct {
|
||||
}
|
||||
|
||||
type pricingIntervalRequest struct {
|
||||
MinTokens int `json:"min_tokens"`
|
||||
MaxTokens *int `json:"max_tokens"`
|
||||
TierLabel string `json:"tier_label"`
|
||||
InputPrice *float64 `json:"input_price"`
|
||||
OutputPrice *float64 `json:"output_price"`
|
||||
CacheWritePrice *float64 `json:"cache_write_price"`
|
||||
CacheReadPrice *float64 `json:"cache_read_price"`
|
||||
PerRequestPrice *float64 `json:"per_request_price"`
|
||||
SortOrder int `json:"sort_order"`
|
||||
MinTokens int `json:"min_tokens"`
|
||||
MaxTokens *int `json:"max_tokens"`
|
||||
TierLabel string `json:"tier_label"`
|
||||
InputPrice *float64 `json:"input_price"`
|
||||
OutputPrice *float64 `json:"output_price"`
|
||||
CacheWritePrice *float64 `json:"cache_write_price"`
|
||||
CacheReadPrice *float64 `json:"cache_read_price"`
|
||||
InputMultiplier *float64 `json:"input_multiplier" binding:"omitempty,gt=0"`
|
||||
OutputMultiplier *float64 `json:"output_multiplier" binding:"omitempty,gt=0"`
|
||||
CacheWriteMultiplier *float64 `json:"cache_write_multiplier" binding:"omitempty,gt=0"`
|
||||
CacheReadMultiplier *float64 `json:"cache_read_multiplier" binding:"omitempty,gt=0"`
|
||||
PerRequestPrice *float64 `json:"per_request_price"`
|
||||
SortOrder int `json:"sort_order"`
|
||||
}
|
||||
|
||||
type accountStatsPricingRuleRequest struct {
|
||||
@@ -128,6 +134,8 @@ type channelModelPricingResponse struct {
|
||||
OutputPrice *float64 `json:"output_price"`
|
||||
CacheWritePrice *float64 `json:"cache_write_price"`
|
||||
CacheReadPrice *float64 `json:"cache_read_price"`
|
||||
FastMultiplier *float64 `json:"fast_multiplier"`
|
||||
FlexMultiplier *float64 `json:"flex_multiplier"`
|
||||
ImageInputPrice *float64 `json:"image_input_price"`
|
||||
ImageOutputPrice *float64 `json:"image_output_price"`
|
||||
PerRequestPrice *float64 `json:"per_request_price"`
|
||||
@@ -147,16 +155,20 @@ type channelTimePricingPeriodResponse struct {
|
||||
}
|
||||
|
||||
type pricingIntervalResponse struct {
|
||||
ID int64 `json:"id"`
|
||||
MinTokens int `json:"min_tokens"`
|
||||
MaxTokens *int `json:"max_tokens"`
|
||||
TierLabel string `json:"tier_label,omitempty"`
|
||||
InputPrice *float64 `json:"input_price"`
|
||||
OutputPrice *float64 `json:"output_price"`
|
||||
CacheWritePrice *float64 `json:"cache_write_price"`
|
||||
CacheReadPrice *float64 `json:"cache_read_price"`
|
||||
PerRequestPrice *float64 `json:"per_request_price"`
|
||||
SortOrder int `json:"sort_order"`
|
||||
ID int64 `json:"id"`
|
||||
MinTokens int `json:"min_tokens"`
|
||||
MaxTokens *int `json:"max_tokens"`
|
||||
TierLabel string `json:"tier_label,omitempty"`
|
||||
InputPrice *float64 `json:"input_price"`
|
||||
OutputPrice *float64 `json:"output_price"`
|
||||
CacheWritePrice *float64 `json:"cache_write_price"`
|
||||
CacheReadPrice *float64 `json:"cache_read_price"`
|
||||
InputMultiplier *float64 `json:"input_multiplier"`
|
||||
OutputMultiplier *float64 `json:"output_multiplier"`
|
||||
CacheWriteMultiplier *float64 `json:"cache_write_multiplier"`
|
||||
CacheReadMultiplier *float64 `json:"cache_read_multiplier"`
|
||||
PerRequestPrice *float64 `json:"per_request_price"`
|
||||
SortOrder int `json:"sort_order"`
|
||||
}
|
||||
|
||||
type accountStatsPricingRuleResponse struct {
|
||||
@@ -248,6 +260,8 @@ func pricingToResponse(p *service.ChannelModelPricing) channelModelPricingRespon
|
||||
OutputPrice: p.OutputPrice,
|
||||
CacheWritePrice: p.CacheWritePrice,
|
||||
CacheReadPrice: p.CacheReadPrice,
|
||||
FastMultiplier: p.FastMultiplier,
|
||||
FlexMultiplier: p.FlexMultiplier,
|
||||
ImageInputPrice: p.ImageInputPrice,
|
||||
ImageOutputPrice: p.ImageOutputPrice,
|
||||
PerRequestPrice: p.PerRequestPrice,
|
||||
@@ -273,20 +287,24 @@ func timePricingToResponse(value *service.ChannelTimePricing) *channelTimePricin
|
||||
|
||||
func intervalToResponse(iv service.PricingInterval) pricingIntervalResponse {
|
||||
return pricingIntervalResponse{
|
||||
ID: iv.ID,
|
||||
MinTokens: iv.MinTokens,
|
||||
MaxTokens: iv.MaxTokens,
|
||||
TierLabel: iv.TierLabel,
|
||||
InputPrice: iv.InputPrice,
|
||||
OutputPrice: iv.OutputPrice,
|
||||
CacheWritePrice: iv.CacheWritePrice,
|
||||
CacheReadPrice: iv.CacheReadPrice,
|
||||
PerRequestPrice: iv.PerRequestPrice,
|
||||
SortOrder: iv.SortOrder,
|
||||
ID: iv.ID,
|
||||
MinTokens: iv.MinTokens,
|
||||
MaxTokens: iv.MaxTokens,
|
||||
TierLabel: iv.TierLabel,
|
||||
InputPrice: iv.InputPrice,
|
||||
OutputPrice: iv.OutputPrice,
|
||||
CacheWritePrice: iv.CacheWritePrice,
|
||||
CacheReadPrice: iv.CacheReadPrice,
|
||||
InputMultiplier: iv.InputMultiplier,
|
||||
OutputMultiplier: iv.OutputMultiplier,
|
||||
CacheWriteMultiplier: iv.CacheWriteMultiplier,
|
||||
CacheReadMultiplier: iv.CacheReadMultiplier,
|
||||
PerRequestPrice: iv.PerRequestPrice,
|
||||
SortOrder: iv.SortOrder,
|
||||
}
|
||||
}
|
||||
|
||||
func pricingRequestToService(reqs []channelModelPricingRequest) []service.ChannelModelPricing {
|
||||
func pricingRequestToService(reqs []channelModelPricingRequest, allowChannelMultipliers bool) []service.ChannelModelPricing {
|
||||
result := make([]service.ChannelModelPricing, 0, len(reqs))
|
||||
for _, r := range reqs {
|
||||
billingMode := service.BillingMode(r.BillingMode)
|
||||
@@ -296,18 +314,34 @@ func pricingRequestToService(reqs []channelModelPricingRequest) []service.Channe
|
||||
platform := r.Platform
|
||||
intervals := make([]service.PricingInterval, 0, len(r.Intervals))
|
||||
for _, iv := range r.Intervals {
|
||||
var inputMultiplier, outputMultiplier, cacheWriteMultiplier, cacheReadMultiplier *float64
|
||||
if allowChannelMultipliers {
|
||||
inputMultiplier = iv.InputMultiplier
|
||||
outputMultiplier = iv.OutputMultiplier
|
||||
cacheWriteMultiplier = iv.CacheWriteMultiplier
|
||||
cacheReadMultiplier = iv.CacheReadMultiplier
|
||||
}
|
||||
intervals = append(intervals, service.PricingInterval{
|
||||
MinTokens: iv.MinTokens,
|
||||
MaxTokens: iv.MaxTokens,
|
||||
TierLabel: iv.TierLabel,
|
||||
InputPrice: iv.InputPrice,
|
||||
OutputPrice: iv.OutputPrice,
|
||||
CacheWritePrice: iv.CacheWritePrice,
|
||||
CacheReadPrice: iv.CacheReadPrice,
|
||||
PerRequestPrice: iv.PerRequestPrice,
|
||||
SortOrder: iv.SortOrder,
|
||||
MinTokens: iv.MinTokens,
|
||||
MaxTokens: iv.MaxTokens,
|
||||
TierLabel: iv.TierLabel,
|
||||
InputPrice: iv.InputPrice,
|
||||
OutputPrice: iv.OutputPrice,
|
||||
CacheWritePrice: iv.CacheWritePrice,
|
||||
CacheReadPrice: iv.CacheReadPrice,
|
||||
InputMultiplier: inputMultiplier,
|
||||
OutputMultiplier: outputMultiplier,
|
||||
CacheWriteMultiplier: cacheWriteMultiplier,
|
||||
CacheReadMultiplier: cacheReadMultiplier,
|
||||
PerRequestPrice: iv.PerRequestPrice,
|
||||
SortOrder: iv.SortOrder,
|
||||
})
|
||||
}
|
||||
var fastMultiplier, flexMultiplier *float64
|
||||
if allowChannelMultipliers {
|
||||
fastMultiplier = r.FastMultiplier
|
||||
flexMultiplier = r.FlexMultiplier
|
||||
}
|
||||
result = append(result, service.ChannelModelPricing{
|
||||
Platform: platform,
|
||||
Models: r.Models,
|
||||
@@ -316,6 +350,8 @@ func pricingRequestToService(reqs []channelModelPricingRequest) []service.Channe
|
||||
OutputPrice: r.OutputPrice,
|
||||
CacheWritePrice: r.CacheWritePrice,
|
||||
CacheReadPrice: r.CacheReadPrice,
|
||||
FastMultiplier: fastMultiplier,
|
||||
FlexMultiplier: flexMultiplier,
|
||||
ImageInputPrice: r.ImageInputPrice,
|
||||
ImageOutputPrice: r.ImageOutputPrice,
|
||||
PerRequestPrice: r.PerRequestPrice,
|
||||
@@ -346,7 +382,7 @@ func accountStatsPricingRuleRequestToService(r accountStatsPricingRuleRequest) s
|
||||
Name: r.Name,
|
||||
GroupIDs: r.GroupIDs,
|
||||
AccountIDs: r.AccountIDs,
|
||||
Pricing: pricingRequestToService(r.Pricing),
|
||||
Pricing: pricingRequestToService(r.Pricing, false),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -407,7 +443,7 @@ func (h *ChannelHandler) Create(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
pricing := pricingRequestToService(req.ModelPricing)
|
||||
pricing := pricingRequestToService(req.ModelPricing, true)
|
||||
// Main model_pricing requires a platform; default to anthropic for backward compatibility.
|
||||
for i := range pricing {
|
||||
if pricing[i].Platform == "" {
|
||||
@@ -481,7 +517,7 @@ func (h *ChannelHandler) Update(c *gin.Context) {
|
||||
ApplyPricingToAccountStats: req.ApplyPricingToAccountStats,
|
||||
}
|
||||
if req.ModelPricing != nil {
|
||||
pricing := pricingRequestToService(*req.ModelPricing)
|
||||
pricing := pricingRequestToService(*req.ModelPricing, true)
|
||||
for i := range pricing {
|
||||
if pricing[i].Platform == "" {
|
||||
pricing[i].Platform = service.PlatformAnthropic
|
||||
|
||||
@@ -305,7 +305,7 @@ func TestPricingRequestToService_Defaults(t *testing.T) {
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := pricingRequestToService([]channelModelPricingRequest{tt.req})
|
||||
result := pricingRequestToService([]channelModelPricingRequest{tt.req}, true)
|
||||
require.Len(t, result, 1)
|
||||
switch tt.wantField {
|
||||
case "BillingMode":
|
||||
@@ -332,7 +332,7 @@ func TestPricingRequestToService_WithAllFields(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
result := pricingRequestToService(reqs)
|
||||
result := pricingRequestToService(reqs, true)
|
||||
require.Len(t, result, 1)
|
||||
r := result[0]
|
||||
require.Equal(t, "openai", r.Platform)
|
||||
@@ -373,7 +373,7 @@ func TestPricingRequestToService_WithIntervals(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
result := pricingRequestToService(reqs)
|
||||
result := pricingRequestToService(reqs, true)
|
||||
require.Len(t, result, 1)
|
||||
require.Len(t, result[0].Intervals, 2)
|
||||
|
||||
@@ -396,7 +396,7 @@ func TestPricingRequestToService_WithIntervals(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestPricingRequestToService_EmptySlice(t *testing.T) {
|
||||
result := pricingRequestToService([]channelModelPricingRequest{})
|
||||
result := pricingRequestToService([]channelModelPricingRequest{}, true)
|
||||
require.NotNil(t, result)
|
||||
require.Empty(t, result)
|
||||
}
|
||||
@@ -410,7 +410,7 @@ func TestPricingRequestToService_NilPriceFields(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
result := pricingRequestToService(reqs)
|
||||
result := pricingRequestToService(reqs, true)
|
||||
require.Len(t, result, 1)
|
||||
r := result[0]
|
||||
require.Nil(t, r.InputPrice)
|
||||
@@ -433,16 +433,52 @@ func TestPricingRequestToService_TimePricing(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
got := pricingRequestToService([]channelModelPricingRequest{req})
|
||||
got := pricingRequestToService([]channelModelPricingRequest{req}, true)
|
||||
require.Equal(t, "Asia/Shanghai", got[0].TimePricing.Timezone)
|
||||
require.Equal(t, 2.0, got[0].TimePricing.Periods[0].Multiplier)
|
||||
}
|
||||
|
||||
func TestPricingRequestToService_TimePricingNil(t *testing.T) {
|
||||
got := pricingRequestToService([]channelModelPricingRequest{{Models: []string{"gpt-5"}}})
|
||||
got := pricingRequestToService([]channelModelPricingRequest{{Models: []string{"gpt-5"}}}, true)
|
||||
require.Nil(t, got[0].TimePricing)
|
||||
}
|
||||
|
||||
// 账号成本统计规则不支持倍率:allowChannelMultipliers=false 时必须丢弃,
|
||||
// 避免渠道倍率意外污染账号成本口径。
|
||||
func TestPricingRequestToService_MultipliersGatedByFlag(t *testing.T) {
|
||||
req := channelModelPricingRequest{
|
||||
Models: []string{"gpt-5"},
|
||||
BillingMode: "token",
|
||||
FastMultiplier: float64Ptr(2.5),
|
||||
FlexMultiplier: float64Ptr(0.5),
|
||||
Intervals: []pricingIntervalRequest{{
|
||||
MinTokens: 272000,
|
||||
InputMultiplier: float64Ptr(2),
|
||||
OutputMultiplier: float64Ptr(1.5),
|
||||
CacheWriteMultiplier: float64Ptr(2),
|
||||
CacheReadMultiplier: float64Ptr(2),
|
||||
}},
|
||||
}
|
||||
|
||||
allowed := pricingRequestToService([]channelModelPricingRequest{req}, true)
|
||||
require.Equal(t, float64Ptr(2.5), allowed[0].FastMultiplier)
|
||||
require.Equal(t, float64Ptr(0.5), allowed[0].FlexMultiplier)
|
||||
require.Equal(t, float64Ptr(2), allowed[0].Intervals[0].InputMultiplier)
|
||||
require.Equal(t, float64Ptr(1.5), allowed[0].Intervals[0].OutputMultiplier)
|
||||
require.Equal(t, float64Ptr(2), allowed[0].Intervals[0].CacheWriteMultiplier)
|
||||
require.Equal(t, float64Ptr(2), allowed[0].Intervals[0].CacheReadMultiplier)
|
||||
|
||||
dropped := pricingRequestToService([]channelModelPricingRequest{req}, false)
|
||||
require.Nil(t, dropped[0].FastMultiplier)
|
||||
require.Nil(t, dropped[0].FlexMultiplier)
|
||||
require.Nil(t, dropped[0].Intervals[0].InputMultiplier)
|
||||
require.Nil(t, dropped[0].Intervals[0].OutputMultiplier)
|
||||
require.Nil(t, dropped[0].Intervals[0].CacheWriteMultiplier)
|
||||
require.Nil(t, dropped[0].Intervals[0].CacheReadMultiplier)
|
||||
// 非倍率字段不受开关影响
|
||||
require.Equal(t, 272000, dropped[0].Intervals[0].MinTokens)
|
||||
}
|
||||
|
||||
func TestPricingToResponse_TimePricing(t *testing.T) {
|
||||
got := pricingToResponse(&service.ChannelModelPricing{
|
||||
BillingMode: service.BillingModeToken,
|
||||
|
||||
@@ -16,7 +16,7 @@ import (
|
||||
|
||||
func (r *channelRepository) ListModelPricing(ctx context.Context, channelID int64) ([]service.ChannelModelPricing, error) {
|
||||
rows, err := r.db.QueryContext(ctx,
|
||||
`SELECT id, channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, image_input_price, image_output_price, per_request_price, time_pricing, created_at, updated_at
|
||||
`SELECT id, channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, fast_multiplier, flex_multiplier, image_input_price, image_output_price, per_request_price, time_pricing, created_at, updated_at
|
||||
FROM channel_model_pricing WHERE channel_id = $1 ORDER BY id`, channelID,
|
||||
)
|
||||
if err != nil {
|
||||
@@ -61,10 +61,11 @@ func (r *channelRepository) UpdateModelPricing(ctx context.Context, pricing *ser
|
||||
}
|
||||
result, err := r.db.ExecContext(ctx,
|
||||
`UPDATE channel_model_pricing
|
||||
SET models = $1, billing_mode = $2, input_price = $3, output_price = $4, cache_write_price = $5, cache_read_price = $6, image_input_price = $7, image_output_price = $8, per_request_price = $9, time_pricing = $10, platform = $11, updated_at = NOW()
|
||||
WHERE id = $12`,
|
||||
SET models = $1, billing_mode = $2, input_price = $3, output_price = $4, cache_write_price = $5, cache_read_price = $6, fast_multiplier = $7, flex_multiplier = $8, image_input_price = $9, image_output_price = $10, per_request_price = $11, time_pricing = $12, platform = $13, updated_at = NOW()
|
||||
WHERE id = $14`,
|
||||
modelsJSON, billingMode, pricing.InputPrice, pricing.OutputPrice, pricing.CacheWritePrice, pricing.CacheReadPrice,
|
||||
pricing.ImageInputPrice, pricing.ImageOutputPrice, pricing.PerRequestPrice, timePricingJSON, pricing.Platform, pricing.ID,
|
||||
pricing.FastMultiplier, pricing.FlexMultiplier, pricing.ImageInputPrice, pricing.ImageOutputPrice, pricing.PerRequestPrice,
|
||||
timePricingJSON, pricing.Platform, pricing.ID,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("update model pricing: %w", err)
|
||||
@@ -95,7 +96,7 @@ func (r *channelRepository) ReplaceModelPricing(ctx context.Context, channelID i
|
||||
// batchLoadModelPricing 批量加载多个渠道的模型定价(含区间)
|
||||
func (r *channelRepository) batchLoadModelPricing(ctx context.Context, channelIDs []int64) (map[int64][]service.ChannelModelPricing, error) {
|
||||
rows, err := r.db.QueryContext(ctx,
|
||||
`SELECT id, channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, image_input_price, image_output_price, per_request_price, time_pricing, created_at, updated_at
|
||||
`SELECT id, channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, fast_multiplier, flex_multiplier, image_input_price, image_output_price, per_request_price, time_pricing, created_at, updated_at
|
||||
FROM channel_model_pricing WHERE channel_id = ANY($1) ORDER BY channel_id, id`,
|
||||
pq.Array(channelIDs),
|
||||
)
|
||||
@@ -136,6 +137,7 @@ func (r *channelRepository) batchLoadIntervals(ctx context.Context, pricingIDs [
|
||||
rows, err := r.db.QueryContext(ctx,
|
||||
`SELECT id, pricing_id, min_tokens, max_tokens, tier_label,
|
||||
input_price, output_price, cache_write_price, cache_read_price,
|
||||
input_multiplier, output_multiplier, cache_write_multiplier, cache_read_multiplier,
|
||||
per_request_price, sort_order, created_at, updated_at
|
||||
FROM channel_pricing_intervals
|
||||
WHERE pricing_id = ANY($1) ORDER BY pricing_id, sort_order, id`,
|
||||
@@ -152,6 +154,7 @@ func (r *channelRepository) batchLoadIntervals(ctx context.Context, pricingIDs [
|
||||
if err := rows.Scan(
|
||||
&iv.ID, &iv.PricingID, &iv.MinTokens, &iv.MaxTokens, &iv.TierLabel,
|
||||
&iv.InputPrice, &iv.OutputPrice, &iv.CacheWritePrice, &iv.CacheReadPrice,
|
||||
&iv.InputMultiplier, &iv.OutputMultiplier, &iv.CacheWriteMultiplier, &iv.CacheReadMultiplier,
|
||||
&iv.PerRequestPrice, &iv.SortOrder, &iv.CreatedAt, &iv.UpdatedAt,
|
||||
); err != nil {
|
||||
return nil, fmt.Errorf("scan interval: %w", err)
|
||||
@@ -177,6 +180,7 @@ func scanModelPricingRows(rows *sql.Rows) ([]service.ChannelModelPricing, []int6
|
||||
if err := rows.Scan(
|
||||
&p.ID, &p.ChannelID, &p.Platform, &modelsJSON, &p.BillingMode,
|
||||
&p.InputPrice, &p.OutputPrice, &p.CacheWritePrice, &p.CacheReadPrice,
|
||||
&p.FastMultiplier, &p.FlexMultiplier,
|
||||
&p.ImageInputPrice, &p.ImageOutputPrice, &p.PerRequestPrice, &timePricingJSON, &p.CreatedAt, &p.UpdatedAt,
|
||||
); err != nil {
|
||||
return nil, nil, fmt.Errorf("scan model pricing: %w", err)
|
||||
@@ -243,11 +247,12 @@ func createModelPricingExec(ctx context.Context, exec dbExec, pricing *service.C
|
||||
platform = "anthropic"
|
||||
}
|
||||
err = exec.QueryRowContext(ctx,
|
||||
`INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, image_input_price, image_output_price, per_request_price, time_pricing)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12) RETURNING id, created_at, updated_at`,
|
||||
`INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, fast_multiplier, flex_multiplier, image_input_price, image_output_price, per_request_price, time_pricing)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14) RETURNING id, created_at, updated_at`,
|
||||
pricing.ChannelID, platform, modelsJSON, billingMode,
|
||||
pricing.InputPrice, pricing.OutputPrice, pricing.CacheWritePrice, pricing.CacheReadPrice,
|
||||
pricing.ImageInputPrice, pricing.ImageOutputPrice, pricing.PerRequestPrice, timePricingJSON,
|
||||
pricing.FastMultiplier, pricing.FlexMultiplier, pricing.ImageInputPrice, pricing.ImageOutputPrice,
|
||||
pricing.PerRequestPrice, timePricingJSON,
|
||||
).Scan(&pricing.ID, &pricing.CreatedAt, &pricing.UpdatedAt)
|
||||
if err != nil {
|
||||
return fmt.Errorf("insert model pricing: %w", err)
|
||||
@@ -288,10 +293,11 @@ func unmarshalChannelTimePricing(data []byte) (*service.ChannelTimePricing, erro
|
||||
func createIntervalExec(ctx context.Context, exec dbExec, iv *service.PricingInterval) error {
|
||||
return exec.QueryRowContext(ctx,
|
||||
`INSERT INTO channel_pricing_intervals
|
||||
(pricing_id, min_tokens, max_tokens, tier_label, input_price, output_price, cache_write_price, cache_read_price, per_request_price, sort_order)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10) RETURNING id, created_at, updated_at`,
|
||||
(pricing_id, min_tokens, max_tokens, tier_label, input_price, output_price, cache_write_price, cache_read_price, input_multiplier, output_multiplier, cache_write_multiplier, cache_read_multiplier, per_request_price, sort_order)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14) RETURNING id, created_at, updated_at`,
|
||||
iv.PricingID, iv.MinTokens, iv.MaxTokens, iv.TierLabel,
|
||||
iv.InputPrice, iv.OutputPrice, iv.CacheWritePrice, iv.CacheReadPrice,
|
||||
iv.InputMultiplier, iv.OutputMultiplier, iv.CacheWriteMultiplier, iv.CacheReadMultiplier,
|
||||
iv.PerRequestPrice, iv.SortOrder,
|
||||
).Scan(&iv.ID, &iv.CreatedAt, &iv.UpdatedAt)
|
||||
}
|
||||
|
||||
@@ -16,7 +16,7 @@ import (
|
||||
|
||||
var channelModelPricingTimePricingColumns = []string{
|
||||
"id", "channel_id", "platform", "models", "billing_mode", "input_price", "output_price",
|
||||
"cache_write_price", "cache_read_price", "image_input_price", "image_output_price",
|
||||
"cache_write_price", "cache_read_price", "fast_multiplier", "flex_multiplier", "image_input_price", "image_output_price",
|
||||
"per_request_price", "time_pricing", "created_at", "updated_at",
|
||||
}
|
||||
|
||||
@@ -33,7 +33,7 @@ func newChannelModelPricingTimePricingRepo(t *testing.T) (*channelRepository, sq
|
||||
func modelPricingTimePricingRow(timePricing any) *sqlmock.Rows {
|
||||
return sqlmock.NewRows(channelModelPricingTimePricingColumns).AddRow(
|
||||
int64(11), int64(7), "openai", `["gpt-5"]`, service.BillingModeToken,
|
||||
nil, nil, nil, nil, nil, nil, nil, timePricing,
|
||||
nil, nil, nil, nil, nil, nil, nil, nil, nil, timePricing,
|
||||
time.Date(2026, 8, 17, 0, 0, 0, 0, time.UTC), time.Date(2026, 8, 17, 1, 0, 0, 0, time.UTC),
|
||||
)
|
||||
}
|
||||
@@ -105,10 +105,10 @@ func TestChannelModelPricingTimePricingCreateAndUpdateRoundTrip(t *testing.T) {
|
||||
|
||||
t.Run("create writes JSON", func(t *testing.T) {
|
||||
repo, mock := newChannelModelPricingTimePricingRepo(t)
|
||||
mock.ExpectQuery(regexp.QuoteMeta("INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, image_input_price, image_output_price, per_request_price, time_pricing)")).
|
||||
mock.ExpectQuery(regexp.QuoteMeta("INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, fast_multiplier, flex_multiplier, image_input_price, image_output_price, per_request_price, time_pricing)")).
|
||||
WithArgs(
|
||||
int64(7), "openai", []byte(`["gpt-5"]`), service.BillingModeToken,
|
||||
nil, nil, nil, nil, nil, nil, nil, channelModelPricingTimePricingJSON,
|
||||
nil, nil, nil, nil, nil, nil, nil, nil, nil, channelModelPricingTimePricingJSON,
|
||||
).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"id", "created_at", "updated_at"}).AddRow(int64(11), time.Time{}, time.Time{}))
|
||||
|
||||
@@ -118,10 +118,10 @@ func TestChannelModelPricingTimePricingCreateAndUpdateRoundTrip(t *testing.T) {
|
||||
|
||||
t.Run("update writes JSON and entry ID", func(t *testing.T) {
|
||||
repo, mock := newChannelModelPricingTimePricingRepo(t)
|
||||
mock.ExpectExec(`(?s)UPDATE channel_model_pricing.*per_request_price = \$9, time_pricing = \$10, platform = \$11.*WHERE id = \$12`).
|
||||
mock.ExpectExec(`(?s)UPDATE channel_model_pricing.*per_request_price = \$11, time_pricing = \$12, platform = \$13.*WHERE id = \$14`).
|
||||
WithArgs(
|
||||
[]byte(`["gpt-5"]`), service.BillingModeToken,
|
||||
nil, nil, nil, nil, nil, nil, nil, channelModelPricingTimePricingJSON, "openai", int64(11),
|
||||
nil, nil, nil, nil, nil, nil, nil, nil, nil, channelModelPricingTimePricingJSON, "openai", int64(11),
|
||||
).
|
||||
WillReturnResult(sqlmock.NewResult(0, 1))
|
||||
|
||||
@@ -153,10 +153,10 @@ func TestChannelModelPricingTimePricingCreateAndUpdateWriteNullWhenDisabled(t *t
|
||||
|
||||
t.Run("create writes SQL NULL", func(t *testing.T) {
|
||||
repo, mock := newChannelModelPricingTimePricingRepo(t)
|
||||
mock.ExpectQuery(regexp.QuoteMeta("INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, image_input_price, image_output_price, per_request_price, time_pricing)")).
|
||||
mock.ExpectQuery(regexp.QuoteMeta("INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, fast_multiplier, flex_multiplier, image_input_price, image_output_price, per_request_price, time_pricing)")).
|
||||
WithArgs(
|
||||
int64(7), "openai", []byte(`["gpt-5"]`), service.BillingModeToken,
|
||||
nil, nil, nil, nil, nil, nil, nil, nil,
|
||||
nil, nil, nil, nil, nil, nil, nil, nil, nil, nil,
|
||||
).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"id", "created_at", "updated_at"}).AddRow(int64(11), time.Time{}, time.Time{}))
|
||||
|
||||
@@ -166,10 +166,10 @@ func TestChannelModelPricingTimePricingCreateAndUpdateWriteNullWhenDisabled(t *t
|
||||
|
||||
t.Run("update writes SQL NULL", func(t *testing.T) {
|
||||
repo, mock := newChannelModelPricingTimePricingRepo(t)
|
||||
mock.ExpectExec(`(?s)UPDATE channel_model_pricing.*per_request_price = \$9, time_pricing = \$10, platform = \$11.*WHERE id = \$12`).
|
||||
mock.ExpectExec(`(?s)UPDATE channel_model_pricing.*per_request_price = \$11, time_pricing = \$12, platform = \$13.*WHERE id = \$14`).
|
||||
WithArgs(
|
||||
[]byte(`["gpt-5"]`), service.BillingModeToken,
|
||||
nil, nil, nil, nil, nil, nil, nil, nil, "openai", int64(11),
|
||||
nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, "openai", int64(11),
|
||||
).
|
||||
WillReturnResult(sqlmock.NewResult(0, 1))
|
||||
|
||||
|
||||
@@ -68,7 +68,7 @@ func tryModelFilePricing(billingService *BillingService, model string, tokens Us
|
||||
return nil
|
||||
}
|
||||
normalizedTier := normalizeBillingServiceTier(serviceTier)
|
||||
if normalizedTier == "priority" || normalizedTier == "flex" ||
|
||||
if normalizedTier == "priority" || normalizedTier == "fast" || normalizedTier == "flex" ||
|
||||
billingService.shouldApplySessionLongContextPricing(tokens, pricing) {
|
||||
breakdown, err := billingService.CalculateCostWithServiceTier(model, tokens, 1, normalizedTier)
|
||||
if err != nil || breakdown == nil || breakdown.TotalCost <= 0 {
|
||||
|
||||
@@ -768,6 +768,26 @@ func TestResolveAccountStatsCost_FallsBackToLiteLLM(t *testing.T) {
|
||||
require.InDelta(t, 0.2, *result, 1e-12)
|
||||
}
|
||||
|
||||
func TestResolveAccountStatsCost_FallbackHonorsAnthropicFast(t *testing.T) {
|
||||
channel := &Channel{ID: 1, Status: StatusActive}
|
||||
cs := newTestChannelServiceForStats(t, channel, 10, "anthropic")
|
||||
bs := newTestBillingServiceWithPrices(map[string]*ModelPricing{
|
||||
"claude-opus-5": {
|
||||
InputPricePerToken: 5e-6,
|
||||
OutputPricePerToken: 25e-6,
|
||||
},
|
||||
})
|
||||
|
||||
result := resolveAccountStatsCost(
|
||||
context.Background(), cs, bs,
|
||||
1, 10, "claude-opus-5",
|
||||
UsageTokens{InputTokens: 1_000_000, OutputTokens: 1_000_000},
|
||||
1, 0, "fast",
|
||||
)
|
||||
require.NotNil(t, result)
|
||||
require.InDelta(t, 60, *result, 1e-12)
|
||||
}
|
||||
|
||||
func TestResolveAccountStatsCost_Gemini36FlashTierUsesFallbackPricing(t *testing.T) {
|
||||
channel := &Channel{
|
||||
ID: 1,
|
||||
|
||||
@@ -91,25 +91,27 @@ type BillingCache interface {
|
||||
|
||||
// ModelPricing 模型价格配置(per-token价格,与LiteLLM格式一致)
|
||||
type ModelPricing struct {
|
||||
InputPricePerToken float64 // 每token输入价格 (USD)
|
||||
InputPricePerTokenPriority float64 // priority service tier 下每token输入价格 (USD)
|
||||
ImageInputPricePerToken float64 // 图片输入 token 价格 (USD),用于多模态 embedding 等图文不同价场景;为 0 时回退到 InputPricePerToken
|
||||
OutputPricePerToken float64 // 每token输出价格 (USD)
|
||||
OutputPricePerTokenPriority float64 // priority service tier 下每token输出价格 (USD)
|
||||
CacheCreationPricePerToken float64 // 缓存创建每token价格 (USD)
|
||||
CacheCreationPricePerTokenPriority float64 // priority service tier 下缓存创建每token价格 (USD)
|
||||
CacheCreationPriceExplicit bool // 是否由渠道/区间定价显式设定(为 true 时即使 == 0 也不回退)
|
||||
CacheReadPricePerToken float64 // 缓存读取每token价格 (USD)
|
||||
CacheReadPricePerTokenPriority float64 // priority service tier 下缓存读取每token价格 (USD)
|
||||
CacheCreation5mPrice float64 // 5分钟缓存创建每token价格 (USD)
|
||||
CacheCreation1hPrice float64 // 1小时缓存创建每token价格 (USD)
|
||||
SupportsCacheBreakdown bool // 是否支持详细的缓存分类
|
||||
LongContextInputThreshold int // 超过阈值后按整次会话提升输入价格
|
||||
LongContextThresholdInclusive bool // 达到阈值即应用(xAI);默认保持严格大于以兼容既有模型
|
||||
LongContextInputMultiplier float64 // 长上下文整次会话输入倍率
|
||||
LongContextOutputMultiplier float64 // 长上下文整次会话输出倍率
|
||||
ImageOutputPricePerToken float64 // 图片输出 token 价格 (USD)
|
||||
ImageOutputPriceExplicit bool // 是否由渠道定价显式设定(为 true 时即使 == 0 也不回退)
|
||||
InputPricePerToken float64 // 每token输入价格 (USD)
|
||||
InputPricePerTokenPriority float64 // priority service tier 下每token输入价格 (USD)
|
||||
ImageInputPricePerToken float64 // 图片输入 token 价格 (USD),用于多模态 embedding 等图文不同价场景;为 0 时回退到 InputPricePerToken
|
||||
OutputPricePerToken float64 // 每token输出价格 (USD)
|
||||
OutputPricePerTokenPriority float64 // priority service tier 下每token输出价格 (USD)
|
||||
CacheCreationPricePerToken float64 // 缓存创建每token价格 (USD)
|
||||
CacheCreationPricePerTokenPriority float64 // priority service tier 下缓存创建每token价格 (USD)
|
||||
CacheCreationPriceExplicit bool // 是否由渠道/区间定价显式设定(为 true 时即使 == 0 也不回退)
|
||||
CacheReadPricePerToken float64 // 缓存读取每token价格 (USD)
|
||||
CacheReadPricePerTokenPriority float64 // priority service tier 下缓存读取每token价格 (USD)
|
||||
FastMultiplier *float64 // 渠道显式 Fast/priority 倍率;nil 时沿用模型目录行为
|
||||
FlexMultiplier *float64 // 渠道显式 Flex 倍率;nil 时沿用默认行为
|
||||
CacheCreation5mPrice float64 // 5分钟缓存创建每token价格 (USD)
|
||||
CacheCreation1hPrice float64 // 1小时缓存创建每token价格 (USD)
|
||||
SupportsCacheBreakdown bool // 是否支持详细的缓存分类
|
||||
LongContextInputThreshold int // 超过阈值后按整次会话提升输入价格
|
||||
LongContextThresholdInclusive bool // 达到阈值即应用(xAI);默认保持严格大于以兼容既有模型
|
||||
LongContextInputMultiplier float64 // 长上下文整次会话输入倍率
|
||||
LongContextOutputMultiplier float64 // 长上下文整次会话输出倍率
|
||||
ImageOutputPricePerToken float64 // 图片输出 token 价格 (USD)
|
||||
ImageOutputPriceExplicit bool // 是否由渠道定价显式设定(为 true 时即使 == 0 也不回退)
|
||||
}
|
||||
|
||||
const (
|
||||
@@ -123,7 +125,14 @@ func normalizeBillingServiceTier(serviceTier string) string {
|
||||
}
|
||||
|
||||
func usePriorityServiceTierPricing(serviceTier string, pricing *ModelPricing) bool {
|
||||
if pricing == nil || normalizeBillingServiceTier(serviceTier) != "priority" {
|
||||
if pricing == nil {
|
||||
return false
|
||||
}
|
||||
tier := normalizeBillingServiceTier(serviceTier)
|
||||
if tier != "priority" && tier != "fast" {
|
||||
return false
|
||||
}
|
||||
if pricing.FastMultiplier != nil {
|
||||
return false
|
||||
}
|
||||
return pricing.InputPricePerTokenPriority > 0 || pricing.OutputPricePerTokenPriority > 0 ||
|
||||
@@ -132,7 +141,7 @@ func usePriorityServiceTierPricing(serviceTier string, pricing *ModelPricing) bo
|
||||
|
||||
func serviceTierCostMultiplier(serviceTier string) float64 {
|
||||
switch normalizeBillingServiceTier(serviceTier) {
|
||||
case "priority":
|
||||
case "priority", "fast":
|
||||
return 2.0
|
||||
case "flex":
|
||||
return 0.5
|
||||
@@ -141,6 +150,34 @@ func serviceTierCostMultiplier(serviceTier string) float64 {
|
||||
}
|
||||
}
|
||||
|
||||
func configuredServiceTierMultiplier(serviceTier string, pricing *ModelPricing) float64 {
|
||||
if pricing != nil {
|
||||
switch normalizeBillingServiceTier(serviceTier) {
|
||||
case "priority", "fast":
|
||||
if pricing.FastMultiplier != nil {
|
||||
return *pricing.FastMultiplier
|
||||
}
|
||||
case "flex":
|
||||
if pricing.FlexMultiplier != nil {
|
||||
return *pricing.FlexMultiplier
|
||||
}
|
||||
}
|
||||
}
|
||||
return serviceTierCostMultiplier(serviceTier)
|
||||
}
|
||||
|
||||
func pricingWithPriorityMultiplier(base *ModelPricing, multiplier float64) *ModelPricing {
|
||||
if base == nil {
|
||||
return nil
|
||||
}
|
||||
cloned := *base
|
||||
cloned.InputPricePerTokenPriority = cloned.InputPricePerToken * multiplier
|
||||
cloned.OutputPricePerTokenPriority = cloned.OutputPricePerToken * multiplier
|
||||
cloned.CacheCreationPricePerTokenPriority = cloned.CacheCreationPricePerToken * multiplier
|
||||
cloned.CacheReadPricePerTokenPriority = cloned.CacheReadPricePerToken * multiplier
|
||||
return &cloned
|
||||
}
|
||||
|
||||
// UsageTokens 使用的token数量
|
||||
type UsageTokens struct {
|
||||
InputTokens int
|
||||
@@ -281,10 +318,10 @@ func (s *BillingService) initFallbackPricing() {
|
||||
// Claude 4.7 Opus (暂与4.6同价,待官方定价更新)
|
||||
s.fallbackPrices["claude-opus-4.7"] = s.fallbackPrices["claude-opus-4.6"]
|
||||
|
||||
// Claude 4.8 Opus / Claude Opus 5(官方同价:$5 输入 / $25 输出 per MTok)。
|
||||
// Claude 4.8 Opus / Claude Opus 5(标准 $5/$25,Fast $10/$50 per MTok)。
|
||||
// 缺少这两条时 getFallbackPricing 会掉到 claude-3-opus($15/$75),造成 3 倍超收。
|
||||
s.fallbackPrices["claude-opus-4.8"] = s.fallbackPrices["claude-opus-4.7"]
|
||||
s.fallbackPrices["claude-opus-5"] = s.fallbackPrices["claude-opus-4.8"]
|
||||
s.fallbackPrices["claude-opus-4.8"] = pricingWithPriorityMultiplier(s.fallbackPrices["claude-opus-4.7"], 2)
|
||||
s.fallbackPrices["claude-opus-5"] = pricingWithPriorityMultiplier(s.fallbackPrices["claude-opus-4.8"], 2)
|
||||
|
||||
// Gemini 3.1 Pro
|
||||
s.fallbackPrices["gemini-3.1-pro"] = &ModelPricing{
|
||||
@@ -320,9 +357,31 @@ func (s *BillingService) initFallbackPricing() {
|
||||
LongContextInputMultiplier: openAIGPT54LongContextInputMultiplier,
|
||||
LongContextOutputMultiplier: openAIGPT54LongContextOutputMultiplier,
|
||||
}
|
||||
// GPT-5.5 / GPT-5.5 Pro 暂无独立定价,回退到 GPT-5.4。
|
||||
s.fallbackPrices["gpt-5.5"] = s.fallbackPrices["gpt-5.4"]
|
||||
s.fallbackPrices["gpt-5.5-pro"] = s.fallbackPrices["gpt-5.4"]
|
||||
// OpenAI GPT-5.5 官方价格;Fast 为标准价 2.5 倍。
|
||||
// Source: https://platform.openai.com/docs/pricing
|
||||
s.fallbackPrices["gpt-5.5"] = pricingWithPriorityMultiplier(&ModelPricing{
|
||||
InputPricePerToken: 5e-6,
|
||||
OutputPricePerToken: 30e-6,
|
||||
// 官方未列独立 cache-write 价;内部出现 cache creation token 时按输入价兜底。
|
||||
CacheCreationPricePerToken: 5e-6,
|
||||
CacheReadPricePerToken: 0.5e-6,
|
||||
SupportsCacheBreakdown: false,
|
||||
LongContextInputThreshold: openAIGPT54LongContextInputThreshold,
|
||||
LongContextInputMultiplier: openAIGPT54LongContextInputMultiplier,
|
||||
LongContextOutputMultiplier: openAIGPT54LongContextOutputMultiplier,
|
||||
}, 2.5)
|
||||
// GPT-5.5 Pro 当前不提供 Fast;保留标准、Flex 和长上下文 fallback 价格。
|
||||
s.fallbackPrices["gpt-5.5-pro"] = &ModelPricing{
|
||||
InputPricePerToken: 30e-6,
|
||||
OutputPricePerToken: 180e-6,
|
||||
// 官方未列独立 cached-input/cache-write 价;内部出现对应 token 时按输入价兜底。
|
||||
CacheCreationPricePerToken: 30e-6,
|
||||
CacheReadPricePerToken: 30e-6,
|
||||
SupportsCacheBreakdown: false,
|
||||
LongContextInputThreshold: openAIGPT54LongContextInputThreshold,
|
||||
LongContextInputMultiplier: openAIGPT54LongContextInputMultiplier,
|
||||
LongContextOutputMultiplier: openAIGPT54LongContextOutputMultiplier,
|
||||
}
|
||||
|
||||
// OpenAI GPT-5.6 官方价格(USD/token)。缓存写入为输入价的 1.25 倍。
|
||||
s.fallbackPrices["gpt-5.6-sol"] = &ModelPricing{
|
||||
@@ -1009,25 +1068,9 @@ func (s *BillingService) GetModelPricingWithChannel(model string, channelPricing
|
||||
// 防止修改 fallbackPrices 中的共享指针
|
||||
cloned := *pricing
|
||||
pricing = &cloned
|
||||
if channelPricing.InputPrice != nil {
|
||||
pricing.InputPricePerToken = *channelPricing.InputPrice
|
||||
pricing.InputPricePerTokenPriority = *channelPricing.InputPrice
|
||||
}
|
||||
if channelPricing.OutputPrice != nil {
|
||||
pricing.OutputPricePerToken = *channelPricing.OutputPrice
|
||||
pricing.OutputPricePerTokenPriority = *channelPricing.OutputPrice
|
||||
}
|
||||
if channelPricing.CacheWritePrice != nil {
|
||||
pricing.CacheCreationPricePerToken = *channelPricing.CacheWritePrice
|
||||
pricing.CacheCreationPricePerTokenPriority = *channelPricing.CacheWritePrice
|
||||
pricing.CacheCreationPriceExplicit = true
|
||||
pricing.CacheCreation5mPrice = *channelPricing.CacheWritePrice
|
||||
pricing.CacheCreation1hPrice = *channelPricing.CacheWritePrice
|
||||
}
|
||||
if channelPricing.CacheReadPrice != nil {
|
||||
pricing.CacheReadPricePerToken = *channelPricing.CacheReadPrice
|
||||
pricing.CacheReadPricePerTokenPriority = *channelPricing.CacheReadPrice
|
||||
}
|
||||
applyChannelTokenPriceOverrides(pricing, channelPricing)
|
||||
pricing.FastMultiplier = channelPricing.FastMultiplier
|
||||
pricing.FlexMultiplier = channelPricing.FlexMultiplier
|
||||
if channelPricing.ImageOutputPrice != nil {
|
||||
pricing.ImageOutputPricePerToken = *channelPricing.ImageOutputPrice
|
||||
} else {
|
||||
@@ -1038,6 +1081,45 @@ func (s *BillingService) GetModelPricingWithChannel(model string, channelPricing
|
||||
return pricing, nil
|
||||
}
|
||||
|
||||
// channelTierOverridePrice applies a Standard-tier override while preserving
|
||||
// an explicit model-catalog Fast/Priority ratio. If the catalog has no tier
|
||||
// price, generic service-tier defaults remain responsible for the fallback.
|
||||
func channelTierOverridePrice(baseStandard, baseTier, channelStandard float64) float64 {
|
||||
if baseStandard > 0 && baseTier > 0 {
|
||||
return channelStandard * (baseTier / baseStandard)
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
func applyChannelTokenPriceOverrides(pricing *ModelPricing, channelPricing *ChannelModelPricing) {
|
||||
if pricing == nil || channelPricing == nil {
|
||||
return
|
||||
}
|
||||
if channelPricing.InputPrice != nil {
|
||||
priority := channelTierOverridePrice(pricing.InputPricePerToken, pricing.InputPricePerTokenPriority, *channelPricing.InputPrice)
|
||||
pricing.InputPricePerToken = *channelPricing.InputPrice
|
||||
pricing.InputPricePerTokenPriority = priority
|
||||
}
|
||||
if channelPricing.OutputPrice != nil {
|
||||
priority := channelTierOverridePrice(pricing.OutputPricePerToken, pricing.OutputPricePerTokenPriority, *channelPricing.OutputPrice)
|
||||
pricing.OutputPricePerToken = *channelPricing.OutputPrice
|
||||
pricing.OutputPricePerTokenPriority = priority
|
||||
}
|
||||
if channelPricing.CacheWritePrice != nil {
|
||||
priority := channelTierOverridePrice(pricing.CacheCreationPricePerToken, pricing.CacheCreationPricePerTokenPriority, *channelPricing.CacheWritePrice)
|
||||
pricing.CacheCreationPricePerToken = *channelPricing.CacheWritePrice
|
||||
pricing.CacheCreationPricePerTokenPriority = priority
|
||||
pricing.CacheCreationPriceExplicit = true
|
||||
pricing.CacheCreation5mPrice = *channelPricing.CacheWritePrice
|
||||
pricing.CacheCreation1hPrice = *channelPricing.CacheWritePrice
|
||||
}
|
||||
if channelPricing.CacheReadPrice != nil {
|
||||
priority := channelTierOverridePrice(pricing.CacheReadPricePerToken, pricing.CacheReadPricePerTokenPriority, *channelPricing.CacheReadPrice)
|
||||
pricing.CacheReadPricePerToken = *channelPricing.CacheReadPrice
|
||||
pricing.CacheReadPricePerTokenPriority = priority
|
||||
}
|
||||
}
|
||||
|
||||
// --- 统一计费入口 ---
|
||||
|
||||
// CostInput 统一计费输入
|
||||
@@ -1113,18 +1195,27 @@ func (s *BillingService) CalculateCostUnified(input CostInput) (*CostBreakdown,
|
||||
func (s *BillingService) calculateTokenCost(resolved *ResolvedPricing, input CostInput) (*CostBreakdown, error) {
|
||||
totalContext := input.Tokens.InputTokens + input.Tokens.CacheCreationTokens + input.Tokens.CacheReadTokens
|
||||
|
||||
pricing := input.Resolver.GetIntervalPricing(resolved, totalContext)
|
||||
// 分组开关是统一入口;账号 API 开关保留为额外开启能力,但 false 不否决分组配置。
|
||||
contextTierPricingEnabled := resolved.longContextPricingEnabled
|
||||
if input.LongContextBillingEnabled != nil && *input.LongContextBillingEnabled {
|
||||
contextTierPricingEnabled = true
|
||||
}
|
||||
|
||||
pricingContext := totalContext
|
||||
if !contextTierPricingEnabled {
|
||||
// 渠道可能显式配置了第一档,也可能只配置高上下文档。用 1 token
|
||||
// 选择最低档;未命中时自然回退到渠道基础价。
|
||||
pricingContext = 1
|
||||
}
|
||||
pricing := input.Resolver.GetIntervalPricing(resolved, pricingContext)
|
||||
if pricing == nil {
|
||||
return nil, fmt.Errorf("no pricing available for model: %s: %w", input.Model, ErrModelPricingUnavailable)
|
||||
}
|
||||
|
||||
pricing = s.applyModelSpecificPricingPolicy(input.Model, pricing)
|
||||
|
||||
// 长上下文定价仅在无区间定价时应用(区间定价已包含上下文分层)
|
||||
applyLongCtx := len(resolved.Intervals) == 0 && resolved.longContextPricingEnabled
|
||||
if input.LongContextBillingEnabled != nil {
|
||||
applyLongCtx = applyLongCtx && *input.LongContextBillingEnabled
|
||||
}
|
||||
// 官方长上下文阶梯仅在无区间定价时应用(区间定价已包含上下文分层)。
|
||||
applyLongCtx := len(resolved.Intervals) == 0 && contextTierPricingEnabled
|
||||
|
||||
breakdown := s.computeTokenBreakdown(pricing, input.Tokens, input.RateMultiplier, input.ServiceTier, applyLongCtx)
|
||||
applyCostBreakdownMultiplier(breakdown, resolvedChannelTimeMultiplier(resolved, input.PricingAt))
|
||||
@@ -1164,7 +1255,7 @@ func (s *BillingService) computeTokenBreakdown(
|
||||
cacheCreationPrice = pricing.CacheCreationPricePerTokenPriority
|
||||
}
|
||||
} else {
|
||||
tierMultiplier = serviceTierCostMultiplier(serviceTier)
|
||||
tierMultiplier = configuredServiceTierMultiplier(serviceTier, pricing)
|
||||
}
|
||||
|
||||
longContextPricingEligible := applyLongCtx && s.shouldApplySessionLongContextPricing(tokens, pricing)
|
||||
|
||||
@@ -215,7 +215,7 @@ func TestGetModelPricing_OpenAICompactAliasesFallback(t *testing.T) {
|
||||
cacheRead float64
|
||||
longContext int
|
||||
}{
|
||||
{model: "gpt5.5", inputPrice: 2.5e-6, outputPrice: 15e-6, cacheRead: 0.25e-6, longContext: 272000},
|
||||
{model: "gpt5.5", inputPrice: 5e-6, outputPrice: 30e-6, cacheRead: 0.5e-6, longContext: 272000},
|
||||
{model: "openai/gpt5.4", inputPrice: 2.5e-6, outputPrice: 15e-6, cacheRead: 0.25e-6, longContext: 272000},
|
||||
{model: "gpt5.4-mini", inputPrice: 7.5e-7, outputPrice: 4.5e-6, cacheRead: 7.5e-8, longContext: 0},
|
||||
{model: "gpt5.3codexspark", inputPrice: 1.5e-6, outputPrice: 12e-6, cacheRead: 0.15e-6, longContext: 0},
|
||||
@@ -293,14 +293,40 @@ func TestCalculateCost_OpenAIGPT55ProUsesGPT55PricingPolicy(t *testing.T) {
|
||||
cost, err := svc.CalculateCost("gpt-5.5-pro", tokens, 1.0)
|
||||
require.NoError(t, err)
|
||||
|
||||
expectedInput := float64(tokens.InputTokens) * 2.5e-6 * 2.0
|
||||
expectedOutput := float64(tokens.OutputTokens) * 15e-6 * 1.5
|
||||
expectedInput := float64(tokens.InputTokens) * 30e-6 * 2.0
|
||||
expectedOutput := float64(tokens.OutputTokens) * 180e-6 * 1.5
|
||||
require.InDelta(t, expectedInput, cost.InputCost, 1e-10)
|
||||
require.InDelta(t, expectedOutput, cost.OutputCost, 1e-10)
|
||||
require.InDelta(t, expectedInput+expectedOutput, cost.TotalCost, 1e-10)
|
||||
require.InDelta(t, expectedInput+expectedOutput, cost.ActualCost, 1e-10)
|
||||
}
|
||||
|
||||
func TestFallbackPricing_OpenAIGPT55UsesOfficialPrices(t *testing.T) {
|
||||
svc := newTestBillingService()
|
||||
|
||||
pricing, err := svc.GetModelPricing("gpt-5.5")
|
||||
require.NoError(t, err)
|
||||
require.InDelta(t, 5e-6, pricing.InputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 30e-6, pricing.OutputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 0.5e-6, pricing.CacheReadPricePerToken, 1e-12)
|
||||
require.InDelta(t, 5e-6, pricing.CacheCreationPricePerToken, 1e-12)
|
||||
require.InDelta(t, 12.5e-6, pricing.InputPricePerTokenPriority, 1e-12)
|
||||
require.InDelta(t, 75e-6, pricing.OutputPricePerTokenPriority, 1e-12)
|
||||
}
|
||||
|
||||
func TestFallbackPricing_OpenAIGPT55ProUsesOfficialPrices(t *testing.T) {
|
||||
svc := newTestBillingService()
|
||||
|
||||
pricing, err := svc.GetModelPricing("gpt-5.5-pro")
|
||||
require.NoError(t, err)
|
||||
require.InDelta(t, 30e-6, pricing.InputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 180e-6, pricing.OutputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 30e-6, pricing.CacheReadPricePerToken, 1e-12)
|
||||
require.InDelta(t, 30e-6, pricing.CacheCreationPricePerToken, 1e-12)
|
||||
require.Zero(t, pricing.InputPricePerTokenPriority)
|
||||
require.Zero(t, pricing.OutputPricePerTokenPriority)
|
||||
}
|
||||
|
||||
// 回归测试 #2293:长上下文计费触发时,cache_read_tokens 也应应用 LongContextInputMultiplier。
|
||||
// 修复前:CacheReadCost = tokens * 0.25e-6 (漏乘倍率,少计费用)。
|
||||
// 修复后:CacheReadCost = tokens * 0.25e-6 * LongContextInputMultiplier(=2.0)。
|
||||
@@ -1594,9 +1620,10 @@ func TestGetModelPricingWithChannel_OverrideInputPriceOnly(t *testing.T) {
|
||||
pricing, err := svc.GetModelPricingWithChannel("claude-sonnet-4", chPricing)
|
||||
require.NoError(t, err)
|
||||
|
||||
// InputPrice overridden (both normal and priority)
|
||||
// InputPrice overridden. claude-sonnet-4 has no catalog priority price, so
|
||||
// the priority slot is zeroed and serviceTierCostMultiplier owns the surcharge.
|
||||
require.InDelta(t, 99e-6, pricing.InputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 99e-6, pricing.InputPricePerTokenPriority, 1e-12)
|
||||
require.Zero(t, pricing.InputPricePerTokenPriority)
|
||||
|
||||
// OutputPrice unchanged (claude-sonnet-4 fallback = 15e-6)
|
||||
require.InDelta(t, 15e-6, pricing.OutputPricePerToken, 1e-12)
|
||||
@@ -1611,9 +1638,9 @@ func TestGetModelPricingWithChannel_OverrideOutputPriceOnly(t *testing.T) {
|
||||
pricing, err := svc.GetModelPricingWithChannel("claude-sonnet-4", chPricing)
|
||||
require.NoError(t, err)
|
||||
|
||||
// OutputPrice overridden
|
||||
// OutputPrice overridden; no catalog priority price to scale, so the slot is zeroed.
|
||||
require.InDelta(t, 88e-6, pricing.OutputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 88e-6, pricing.OutputPricePerTokenPriority, 1e-12)
|
||||
require.Zero(t, pricing.OutputPricePerTokenPriority)
|
||||
|
||||
// InputPrice unchanged (claude-sonnet-4 fallback = 3e-6)
|
||||
require.InDelta(t, 3e-6, pricing.InputPricePerToken, 1e-12)
|
||||
@@ -1633,15 +1660,18 @@ func TestGetModelPricingWithChannel_OverrideAllFields(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
|
||||
require.InDelta(t, 10e-6, pricing.InputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 10e-6, pricing.InputPricePerTokenPriority, 1e-12)
|
||||
require.InDelta(t, 20e-6, pricing.OutputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 20e-6, pricing.OutputPricePerTokenPriority, 1e-12)
|
||||
require.InDelta(t, 5e-6, pricing.CacheCreationPricePerToken, 1e-12)
|
||||
require.InDelta(t, 5e-6, pricing.CacheCreation5mPrice, 1e-12)
|
||||
require.InDelta(t, 5e-6, pricing.CacheCreation1hPrice, 1e-12)
|
||||
require.InDelta(t, 1e-6, pricing.CacheReadPricePerToken, 1e-12)
|
||||
require.InDelta(t, 1e-6, pricing.CacheReadPricePerTokenPriority, 1e-12)
|
||||
require.InDelta(t, 50e-6, pricing.ImageOutputPricePerToken, 1e-12)
|
||||
|
||||
// claude-sonnet-4 carries no catalog Fast/Priority tier, so every priority
|
||||
// slot stays zero and computeTokenBreakdown falls back to the 2x default.
|
||||
require.Zero(t, pricing.InputPricePerTokenPriority)
|
||||
require.Zero(t, pricing.OutputPricePerTokenPriority)
|
||||
require.Zero(t, pricing.CacheReadPricePerTokenPriority)
|
||||
}
|
||||
|
||||
func TestGetModelPricingWithChannel_CacheWritePriceAffects5mAnd1h(t *testing.T) {
|
||||
@@ -1668,9 +1698,27 @@ func TestGetModelPricingWithChannel_CacheReadPriceAffectsPriority(t *testing.T)
|
||||
pricing, err := svc.GetModelPricingWithChannel("claude-sonnet-4", chPricing)
|
||||
require.NoError(t, err)
|
||||
|
||||
// CacheReadPrice should set both normal and priority
|
||||
// CacheReadPrice sets the standard slot; the priority slot is zeroed because
|
||||
// claude-sonnet-4 has no catalog tier ratio to preserve.
|
||||
require.InDelta(t, 2e-6, pricing.CacheReadPricePerToken, 1e-12)
|
||||
require.InDelta(t, 2e-6, pricing.CacheReadPricePerTokenPriority, 1e-12)
|
||||
require.Zero(t, pricing.CacheReadPricePerTokenPriority)
|
||||
}
|
||||
|
||||
// 目录带 tier 价时,渠道覆盖必须按目录比例换算 priority 价,而不是归零。
|
||||
func TestGetModelPricingWithChannel_PreservesCatalogPriorityRatio(t *testing.T) {
|
||||
svc := newTestBillingService()
|
||||
|
||||
// gpt-5.4 目录价:input 2.5/5(2x),output 15/30(2x)。
|
||||
pricing, err := svc.GetModelPricingWithChannel("gpt-5.4", &ChannelModelPricing{
|
||||
InputPrice: testPtrFloat64(4e-6),
|
||||
OutputPrice: testPtrFloat64(30e-6),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
require.InDelta(t, 4e-6, pricing.InputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 8e-6, pricing.InputPricePerTokenPriority, 1e-12)
|
||||
require.InDelta(t, 30e-6, pricing.OutputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 60e-6, pricing.OutputPricePerTokenPriority, 1e-12)
|
||||
}
|
||||
|
||||
func TestGetModelPricingWithChannel_UnknownModelReturnsError(t *testing.T) {
|
||||
|
||||
@@ -96,6 +96,8 @@ type ChannelModelPricing struct {
|
||||
OutputPrice *float64 `json:"output_price"`
|
||||
CacheWritePrice *float64 `json:"cache_write_price"`
|
||||
CacheReadPrice *float64 `json:"cache_read_price"`
|
||||
FastMultiplier *float64 `json:"fast_multiplier"`
|
||||
FlexMultiplier *float64 `json:"flex_multiplier"`
|
||||
ImageInputPrice *float64 `json:"image_input_price"`
|
||||
ImageOutputPrice *float64 `json:"image_output_price"`
|
||||
PerRequestPrice *float64 `json:"per_request_price"`
|
||||
@@ -120,19 +122,23 @@ type ChannelTimePricingPeriod struct {
|
||||
|
||||
// PricingInterval 定价区间(token 区间 / 按次分层 / 图片分辨率分层)
|
||||
type PricingInterval struct {
|
||||
ID int64 `json:"id,omitempty"`
|
||||
PricingID int64 `json:"pricing_id,omitempty"`
|
||||
MinTokens int `json:"min_tokens"`
|
||||
MaxTokens *int `json:"max_tokens"`
|
||||
TierLabel string `json:"tier_label"`
|
||||
InputPrice *float64 `json:"input_price"`
|
||||
OutputPrice *float64 `json:"output_price"`
|
||||
CacheWritePrice *float64 `json:"cache_write_price"`
|
||||
CacheReadPrice *float64 `json:"cache_read_price"`
|
||||
PerRequestPrice *float64 `json:"per_request_price"`
|
||||
SortOrder int `json:"sort_order"`
|
||||
CreatedAt time.Time `json:"created_at,omitempty"`
|
||||
UpdatedAt time.Time `json:"updated_at,omitempty"`
|
||||
ID int64 `json:"id,omitempty"`
|
||||
PricingID int64 `json:"pricing_id,omitempty"`
|
||||
MinTokens int `json:"min_tokens"`
|
||||
MaxTokens *int `json:"max_tokens"`
|
||||
TierLabel string `json:"tier_label"`
|
||||
InputPrice *float64 `json:"input_price"`
|
||||
OutputPrice *float64 `json:"output_price"`
|
||||
CacheWritePrice *float64 `json:"cache_write_price"`
|
||||
CacheReadPrice *float64 `json:"cache_read_price"`
|
||||
InputMultiplier *float64 `json:"input_multiplier"`
|
||||
OutputMultiplier *float64 `json:"output_multiplier"`
|
||||
CacheWriteMultiplier *float64 `json:"cache_write_multiplier"`
|
||||
CacheReadMultiplier *float64 `json:"cache_read_multiplier"`
|
||||
PerRequestPrice *float64 `json:"per_request_price"`
|
||||
SortOrder int `json:"sort_order"`
|
||||
CreatedAt time.Time `json:"created_at,omitempty"`
|
||||
UpdatedAt time.Time `json:"updated_at,omitempty"`
|
||||
}
|
||||
|
||||
// IsActive 判断渠道是否启用
|
||||
@@ -358,7 +364,7 @@ func validateSingleInterval(iv *PricingInterval, idx int) error {
|
||||
return validateIntervalPrices(iv, idx)
|
||||
}
|
||||
|
||||
// validateIntervalPrices 校验区间内所有价格字段 >= 0
|
||||
// validateIntervalPrices 校验区间价格 >= 0、倍率 > 0。
|
||||
func validateIntervalPrices(iv *PricingInterval, idx int) error {
|
||||
prices := []struct {
|
||||
name string
|
||||
@@ -375,6 +381,20 @@ func validateIntervalPrices(iv *PricingInterval, idx int) error {
|
||||
return fmt.Errorf("interval #%d: %s must be >= 0", idx+1, p.name)
|
||||
}
|
||||
}
|
||||
multipliers := []struct {
|
||||
name string
|
||||
val *float64
|
||||
}{
|
||||
{"input_multiplier", iv.InputMultiplier},
|
||||
{"output_multiplier", iv.OutputMultiplier},
|
||||
{"cache_write_multiplier", iv.CacheWriteMultiplier},
|
||||
{"cache_read_multiplier", iv.CacheReadMultiplier},
|
||||
}
|
||||
for _, multiplier := range multipliers {
|
||||
if multiplier.val != nil && *multiplier.val <= 0 {
|
||||
return fmt.Errorf("interval #%d: %s must be > 0", idx+1, multiplier.name)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,288 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func pricingMultiplier(value float64) *float64 { return &value }
|
||||
|
||||
func TestConfiguredServiceTierMultiplier(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
serviceTier string
|
||||
pricing *ModelPricing
|
||||
want float64
|
||||
}{
|
||||
{name: "gpt-5.5 fast", serviceTier: "fast", pricing: &ModelPricing{FastMultiplier: pricingMultiplier(2.5)}, want: 2.5},
|
||||
{name: "priority alias", serviceTier: "priority", pricing: &ModelPricing{FastMultiplier: pricingMultiplier(2)}, want: 2},
|
||||
{name: "flex configured", serviceTier: "flex", pricing: &ModelPricing{FlexMultiplier: pricingMultiplier(0.4)}, want: 0.4},
|
||||
{name: "legacy fast default", serviceTier: "fast", pricing: &ModelPricing{}, want: 2},
|
||||
{name: "legacy flex default", serviceTier: "flex", pricing: &ModelPricing{}, want: 0.5},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
require.InDelta(t, tt.want, configuredServiceTierMultiplier(tt.serviceTier, tt.pricing), 1e-12)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfiguredServiceTierMultiplierAppliesToEveryTokenComponent(t *testing.T) {
|
||||
pricing := &ModelPricing{
|
||||
InputPricePerToken: 5e-6,
|
||||
OutputPricePerToken: 30e-6,
|
||||
CacheCreationPricePerToken: 6.25e-6,
|
||||
CacheReadPricePerToken: 0.5e-6,
|
||||
FastMultiplier: pricingMultiplier(2.5),
|
||||
FlexMultiplier: pricingMultiplier(0.5),
|
||||
}
|
||||
tokens := UsageTokens{
|
||||
InputTokens: 1_000_000,
|
||||
OutputTokens: 1_000_000,
|
||||
CacheCreationTokens: 1_000_000,
|
||||
CacheReadTokens: 1_000_000,
|
||||
}
|
||||
service := &BillingService{}
|
||||
|
||||
fast := service.computeTokenBreakdown(pricing, tokens, 1, "fast", false)
|
||||
require.InDelta(t, 12.5, fast.InputCost, 1e-12)
|
||||
require.InDelta(t, 75, fast.OutputCost, 1e-12)
|
||||
require.InDelta(t, 15.625, fast.CacheCreationCost, 1e-12)
|
||||
require.InDelta(t, 1.25, fast.CacheReadCost, 1e-12)
|
||||
|
||||
flex := service.computeTokenBreakdown(pricing, tokens, 1, "flex", false)
|
||||
require.InDelta(t, 2.5, flex.InputCost, 1e-12)
|
||||
require.InDelta(t, 15, flex.OutputCost, 1e-12)
|
||||
require.InDelta(t, 3.125, flex.CacheCreationCost, 1e-12)
|
||||
require.InDelta(t, 0.25, flex.CacheReadCost, 1e-12)
|
||||
}
|
||||
|
||||
func TestChannelOverridePreservesCatalogFastRatioByDefault(t *testing.T) {
|
||||
pricing := &ModelPricing{
|
||||
InputPricePerToken: 2,
|
||||
InputPricePerTokenPriority: 4,
|
||||
OutputPricePerToken: 6,
|
||||
OutputPricePerTokenPriority: 12,
|
||||
}
|
||||
applyChannelTokenPriceOverrides(pricing, &ChannelModelPricing{
|
||||
InputPrice: pricingMultiplier(3),
|
||||
OutputPrice: pricingMultiplier(9),
|
||||
})
|
||||
|
||||
require.InDelta(t, 3, pricing.InputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 6, pricing.InputPricePerTokenPriority, 1e-12)
|
||||
require.InDelta(t, 9, pricing.OutputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 18, pricing.OutputPricePerTokenPriority, 1e-12)
|
||||
}
|
||||
|
||||
func TestAnthropicFastUsesDefaultMultiplierWithoutCatalogTier(t *testing.T) {
|
||||
pricing := &ModelPricing{InputPricePerToken: 5e-6, OutputPricePerToken: 25e-6}
|
||||
cost := (&BillingService{}).computeTokenBreakdown(pricing, UsageTokens{
|
||||
InputTokens: 1_000_000, OutputTokens: 1_000_000,
|
||||
}, 1, "fast", false)
|
||||
|
||||
require.InDelta(t, 10, cost.InputCost, 1e-12)
|
||||
require.InDelta(t, 50, cost.OutputCost, 1e-12)
|
||||
}
|
||||
|
||||
func TestBuiltInModelFastDefaults(t *testing.T) {
|
||||
service := &BillingService{fallbackPrices: make(map[string]*ModelPricing)}
|
||||
service.initFallbackPricing()
|
||||
|
||||
for _, tt := range []struct {
|
||||
model string
|
||||
want float64
|
||||
}{
|
||||
{model: "gpt-5.5", want: 2.5},
|
||||
{model: "claude-opus-4.8", want: 2},
|
||||
{model: "claude-opus-5", want: 2},
|
||||
} {
|
||||
pricing := service.fallbackPrices[tt.model]
|
||||
require.NotNil(t, pricing)
|
||||
require.InDelta(t, tt.want, pricing.InputPricePerTokenPriority/pricing.InputPricePerToken, 1e-12)
|
||||
require.InDelta(t, tt.want, pricing.OutputPricePerTokenPriority/pricing.OutputPricePerToken, 1e-12)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIntervalMultipliersApplyToChannelBase(t *testing.T) {
|
||||
base := &ModelPricing{
|
||||
InputPricePerToken: 5,
|
||||
OutputPricePerToken: 30,
|
||||
CacheCreationPricePerToken: 6.25,
|
||||
CacheCreation5mPrice: 6.25,
|
||||
CacheCreation1hPrice: 6.25,
|
||||
CacheReadPricePerToken: 0.5,
|
||||
FastMultiplier: pricingMultiplier(2),
|
||||
FlexMultiplier: pricingMultiplier(0.5),
|
||||
}
|
||||
resolved := &ResolvedPricing{
|
||||
BasePricing: base,
|
||||
Intervals: []PricingInterval{{
|
||||
MinTokens: 272000,
|
||||
InputMultiplier: pricingMultiplier(2),
|
||||
OutputMultiplier: pricingMultiplier(1.5),
|
||||
CacheWriteMultiplier: pricingMultiplier(2),
|
||||
CacheReadMultiplier: pricingMultiplier(2),
|
||||
}},
|
||||
}
|
||||
|
||||
pricing := (&ModelPricingResolver{}).GetIntervalPricing(resolved, 272001)
|
||||
require.InDelta(t, 10, pricing.InputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 45, pricing.OutputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 12.5, pricing.CacheCreationPricePerToken, 1e-12)
|
||||
require.InDelta(t, 1, pricing.CacheReadPricePerToken, 1e-12)
|
||||
require.Same(t, base, (&ModelPricingResolver{}).GetIntervalPricing(resolved, 272000))
|
||||
}
|
||||
|
||||
func TestIntervalExplicitPriceTakesPrecedenceOverMultiplier(t *testing.T) {
|
||||
pricing := intervalToModelPricing(&PricingInterval{
|
||||
InputPrice: pricingMultiplier(7),
|
||||
InputMultiplier: pricingMultiplier(2),
|
||||
}, &ModelPricing{InputPricePerToken: 5}, nil)
|
||||
|
||||
require.InDelta(t, 7, pricing.InputPricePerToken, 1e-12)
|
||||
}
|
||||
|
||||
func TestIntervalPricePreservesDefaultFastRatio(t *testing.T) {
|
||||
pricing := intervalToModelPricing(&PricingInterval{
|
||||
InputPrice: pricingMultiplier(7),
|
||||
}, &ModelPricing{
|
||||
InputPricePerToken: 5,
|
||||
InputPricePerTokenPriority: 10,
|
||||
}, nil)
|
||||
|
||||
require.InDelta(t, 7, pricing.InputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 14, pricing.InputPricePerTokenPriority, 1e-12)
|
||||
}
|
||||
|
||||
func TestAnthropicSpeedServiceTier(t *testing.T) {
|
||||
account := &Account{Platform: PlatformAnthropic}
|
||||
|
||||
for _, model := range []string{"claude-opus-5", "claude-opus-4-8", "claude-opus-4.8"} {
|
||||
tier := anthropicSpeedServiceTier(account, "fast", model)
|
||||
require.NotNil(t, tier, "model %s should bill as fast", model)
|
||||
require.Equal(t, "fast", *tier)
|
||||
}
|
||||
|
||||
require.Nil(t, anthropicSpeedServiceTier(&Account{Platform: PlatformOpenAI}, "fast", "claude-opus-5"))
|
||||
require.Nil(t, anthropicSpeedServiceTier(account, "standard", "claude-opus-5"))
|
||||
}
|
||||
|
||||
// fast mode 不存在于这些模型/承载上,即便客户端传了 speed=fast 也不能计 2x。
|
||||
func TestAnthropicSpeedServiceTierRejectsUnsupportedTargets(t *testing.T) {
|
||||
account := &Account{Platform: PlatformAnthropic}
|
||||
|
||||
for _, model := range []string{
|
||||
"claude-opus-4-7", // fast mode 已被移除
|
||||
"claude-opus-4-6", //
|
||||
"claude-opus-4-5", // 不能被 "opus-5" 规则误判
|
||||
"claude-sonnet-5", // 非 Opus
|
||||
"claude-haiku-4-5", //
|
||||
"", //
|
||||
} {
|
||||
require.Nil(t, anthropicSpeedServiceTier(account, "fast", model),
|
||||
"model %q must not bill as fast", model)
|
||||
}
|
||||
|
||||
bedrock := &Account{Platform: PlatformAnthropic, Type: AccountTypeBedrock}
|
||||
require.Nil(t, anthropicSpeedServiceTier(bedrock, "fast", "claude-opus-5"))
|
||||
}
|
||||
|
||||
func TestAnthropicSpeedModelPrefersMappedUpstreamModel(t *testing.T) {
|
||||
parsed := &ParsedRequest{Model: "claude-opus-5"}
|
||||
require.Equal(t, "claude-opus-4-7", anthropicSpeedModel(parsed, &ForwardResult{
|
||||
UpstreamModel: "claude-opus-4-7",
|
||||
}))
|
||||
require.Equal(t, "claude-opus-5", anthropicSpeedModel(parsed, &ForwardResult{}))
|
||||
}
|
||||
|
||||
func TestMultiplierOnlyIntervalIsValid(t *testing.T) {
|
||||
require.NoError(t, ValidateIntervals([]PricingInterval{{
|
||||
MinTokens: 199999,
|
||||
InputMultiplier: pricingMultiplier(2),
|
||||
}}, BillingModeToken))
|
||||
require.NoError(t, checkIntervalsHavePrices(ChannelModelPricing{
|
||||
Models: []string{"grok-4.6"},
|
||||
Intervals: []PricingInterval{{
|
||||
MinTokens: 199999,
|
||||
InputMultiplier: pricingMultiplier(2),
|
||||
}},
|
||||
}))
|
||||
}
|
||||
|
||||
func TestChannelMultipliersMustBePositive(t *testing.T) {
|
||||
zero := 0.0
|
||||
require.Error(t, checkPricesNotNegative(ChannelModelPricing{FastMultiplier: &zero}))
|
||||
require.Error(t, checkPricesNotNegative(ChannelModelPricing{FlexMultiplier: &zero}))
|
||||
require.Error(t, ValidateIntervals([]PricingInterval{{
|
||||
MinTokens: 100,
|
||||
InputMultiplier: &zero,
|
||||
}}, BillingModeToken))
|
||||
}
|
||||
|
||||
func TestCalculateTokenCostContextTierEnablement(t *testing.T) {
|
||||
base := &ModelPricing{InputPricePerToken: 1e-6}
|
||||
resolved := &ResolvedPricing{
|
||||
BasePricing: base,
|
||||
Intervals: []PricingInterval{{
|
||||
MinTokens: 100,
|
||||
InputMultiplier: pricingMultiplier(2),
|
||||
}},
|
||||
}
|
||||
resolver := &ModelPricingResolver{}
|
||||
service := &BillingService{}
|
||||
tokens := UsageTokens{InputTokens: 200}
|
||||
|
||||
t.Run("group disabled uses base tier", func(t *testing.T) {
|
||||
resolved.longContextPricingEnabled = false
|
||||
cost, err := service.calculateTokenCost(resolved, CostInput{
|
||||
Model: "custom", Tokens: tokens, RateMultiplier: 1, Resolver: resolver,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.InDelta(t, 200e-6, cost.TotalCost, 1e-12)
|
||||
})
|
||||
|
||||
t.Run("group enabled uses interval", func(t *testing.T) {
|
||||
resolved.longContextPricingEnabled = true
|
||||
accountDisabled := false
|
||||
cost, err := service.calculateTokenCost(resolved, CostInput{
|
||||
Model: "custom", Tokens: tokens, RateMultiplier: 1, Resolver: resolver,
|
||||
LongContextBillingEnabled: &accountDisabled,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.InDelta(t, 400e-6, cost.TotalCost, 1e-12)
|
||||
})
|
||||
|
||||
t.Run("account enabled overrides disabled group", func(t *testing.T) {
|
||||
resolved.longContextPricingEnabled = false
|
||||
accountEnabled := true
|
||||
cost, err := service.calculateTokenCost(resolved, CostInput{
|
||||
Model: "custom", Tokens: tokens, RateMultiplier: 1, Resolver: resolver,
|
||||
LongContextBillingEnabled: &accountEnabled,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.InDelta(t, 400e-6, cost.TotalCost, 1e-12)
|
||||
})
|
||||
}
|
||||
|
||||
func TestCalculateTokenCostCombinesIntervalAndFastMultiplier(t *testing.T) {
|
||||
resolved := &ResolvedPricing{
|
||||
BasePricing: &ModelPricing{
|
||||
InputPricePerToken: 1e-6,
|
||||
FastMultiplier: pricingMultiplier(2.5),
|
||||
},
|
||||
Intervals: []PricingInterval{{
|
||||
MinTokens: 100,
|
||||
InputMultiplier: pricingMultiplier(2),
|
||||
}},
|
||||
longContextPricingEnabled: true,
|
||||
}
|
||||
cost, err := (&BillingService{}).calculateTokenCost(resolved, CostInput{
|
||||
Model: "custom", Tokens: UsageTokens{InputTokens: 200}, RateMultiplier: 1,
|
||||
ServiceTier: "fast", Resolver: &ModelPricingResolver{},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.InDelta(t, 1e-3, cost.TotalCost, 1e-12)
|
||||
}
|
||||
@@ -735,6 +735,17 @@ func checkPricesNotNegative(p ChannelModelPricing) error {
|
||||
return infraerrors.BadRequest("NEGATIVE_PRICE", fmt.Sprintf("%s must be >= 0", c.field))
|
||||
}
|
||||
}
|
||||
for _, c := range []struct {
|
||||
field string
|
||||
val *float64
|
||||
}{
|
||||
{"fast_multiplier", p.FastMultiplier},
|
||||
{"flex_multiplier", p.FlexMultiplier},
|
||||
} {
|
||||
if c.val != nil && *c.val <= 0 {
|
||||
return infraerrors.BadRequest("INVALID_MULTIPLIER", fmt.Sprintf("%s must be > 0", c.field))
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -742,7 +753,9 @@ func checkIntervalsHavePrices(p ChannelModelPricing) error {
|
||||
for _, iv := range p.Intervals {
|
||||
if iv.InputPrice == nil && iv.OutputPrice == nil &&
|
||||
iv.CacheWritePrice == nil && iv.CacheReadPrice == nil &&
|
||||
iv.PerRequestPrice == nil {
|
||||
iv.PerRequestPrice == nil && iv.InputMultiplier == nil &&
|
||||
iv.OutputMultiplier == nil && iv.CacheWriteMultiplier == nil &&
|
||||
iv.CacheReadMultiplier == nil {
|
||||
return infraerrors.BadRequest(
|
||||
"INTERVAL_MISSING_PRICE",
|
||||
fmt.Sprintf("interval [%d, %s] has no price fields set for model %v",
|
||||
|
||||
@@ -88,11 +88,21 @@ func sleepWithContext(ctx context.Context, d time.Duration) error {
|
||||
}
|
||||
|
||||
// Forward 转发请求到Claude API
|
||||
func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *Account, parsed *ParsedRequest) (*ForwardResult, error) {
|
||||
func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *Account, parsed *ParsedRequest) (result *ForwardResult, err error) {
|
||||
startTime := time.Now()
|
||||
if parsed == nil {
|
||||
return nil, fmt.Errorf("parse request: empty request")
|
||||
}
|
||||
// Anthropic Fast is requested with speed=fast rather than OpenAI's
|
||||
// service_tier. Attach it at this shared boundary so passthrough, OAuth and
|
||||
// partial-stream results all use the same billing and usage-log path.
|
||||
defer func() {
|
||||
if result != nil {
|
||||
if tier := anthropicSpeedServiceTier(account, parsed.Speed, anthropicSpeedModel(parsed, result)); tier != nil {
|
||||
result.ServiceTier = tier
|
||||
}
|
||||
}
|
||||
}()
|
||||
beginUpstreamResponseModelObservation(c)
|
||||
|
||||
// Web Search 模拟:纯 web_search 请求时,直接调用搜索 API 构造响应
|
||||
@@ -885,6 +895,50 @@ func (s *GatewayService) Forward(ctx context.Context, c *gin.Context, account *A
|
||||
}, nil
|
||||
}
|
||||
|
||||
func anthropicSpeedModel(parsed *ParsedRequest, result *ForwardResult) string {
|
||||
if result != nil {
|
||||
if upstreamModel := strings.TrimSpace(result.UpstreamModel); upstreamModel != "" {
|
||||
return upstreamModel
|
||||
}
|
||||
}
|
||||
if parsed == nil {
|
||||
return ""
|
||||
}
|
||||
return parsed.Model
|
||||
}
|
||||
|
||||
// anthropicSpeedServiceTier 把 Anthropic 的 speed=fast 归一成可计费的 "fast" tier。
|
||||
//
|
||||
// Fast mode 目前只在 Claude Opus 5 / Opus 4.8 上存在,且不支持 Bedrock 等第三方
|
||||
// 承载(Opus 4.7 的 fast mode 已被移除,传 speed=fast 会直接报错)。这里按模型和
|
||||
// 平台收紧,避免上游根本没跑 fast 时仍然按 2x 计费——宁可漏收也不能多收。
|
||||
//
|
||||
// 注:判据是请求参数而非响应里的 usage.speed。等 usage 解析链路统一暴露该字段后,
|
||||
// 应改为以响应为准。
|
||||
func anthropicSpeedServiceTier(account *Account, speed, model string) *string {
|
||||
if account == nil || account.Platform != PlatformAnthropic || speed != "fast" {
|
||||
return nil
|
||||
}
|
||||
if account.IsBedrock() || !modelSupportsAnthropicFastMode(model) {
|
||||
return nil
|
||||
}
|
||||
tier := "fast"
|
||||
return &tier
|
||||
}
|
||||
|
||||
// modelSupportsAnthropicFastMode 判断模型是否属于支持 fast mode 的 Opus 5 / Opus 4.8。
|
||||
func modelSupportsAnthropicFastMode(model string) bool {
|
||||
modelLower := strings.ToLower(strings.TrimSpace(model))
|
||||
if !strings.Contains(modelLower, "opus") {
|
||||
return false
|
||||
}
|
||||
// "opus-5" 必须先判:不能用裸 "5" 匹配,否则 claude-opus-4-5 会被误判。
|
||||
if strings.Contains(modelLower, "opus-5") || strings.Contains(modelLower, "opus5") {
|
||||
return true
|
||||
}
|
||||
return strings.Contains(modelLower, "4.8") || strings.Contains(modelLower, "4-8")
|
||||
}
|
||||
|
||||
// ResolveChannelMapping 委托渠道服务解析模型映射
|
||||
func (s *GatewayService) ResolveChannelMapping(ctx context.Context, groupID int64, model string) ChannelMappingResult {
|
||||
if s.channelService == nil {
|
||||
|
||||
@@ -117,6 +117,7 @@ func clearGatewayRequestDerivedState(parsed *ParsedRequest) {
|
||||
parsed.HasSystem = false
|
||||
parsed.ThinkingEnabled = false
|
||||
parsed.OutputEffort = ""
|
||||
parsed.Speed = ""
|
||||
parsed.MaxTokens = 0
|
||||
parsed.systemRange = missingJSONRange()
|
||||
parsed.messagesRange = missingJSONRange()
|
||||
@@ -224,6 +225,9 @@ func parseGatewayRequestCurrentBody(parsed *ParsedRequest, protocol string) erro
|
||||
parsed.ThinkingEnabled = thinkingType == "enabled" || thinkingType == "adaptive"
|
||||
|
||||
parsed.OutputEffort = strings.TrimSpace(gjson.Get(jsonStr, "output_config.effort").String())
|
||||
if protocol == domain.PlatformAnthropic {
|
||||
parsed.Speed = strings.ToLower(strings.TrimSpace(gjson.Get(jsonStr, "speed").String()))
|
||||
}
|
||||
|
||||
maxTokensResult := gjson.Get(jsonStr, "max_tokens")
|
||||
if maxTokensResult.Exists() && maxTokensResult.Type == gjson.Number {
|
||||
@@ -282,6 +286,7 @@ type ParsedRequest struct {
|
||||
HasSystem bool // 是否包含 system 字段(包含 null 也视为显式传入)
|
||||
ThinkingEnabled bool // 是否开启 thinking(部分平台会影响最终模型名)
|
||||
OutputEffort string // output_config.effort(Claude API 的推理强度控制)
|
||||
Speed string // Anthropic speed(当前可计费值为 "fast")
|
||||
MaxTokens int // max_tokens 值(用于探测请求拦截)
|
||||
SessionContext *SessionContext // 可选:请求上下文区分因子(nil 时行为不变)
|
||||
|
||||
|
||||
@@ -42,6 +42,22 @@ func TestParseGatewayRequest_ThinkingAdaptiveEnabled(t *testing.T) {
|
||||
require.True(t, parsed.ThinkingEnabled)
|
||||
}
|
||||
|
||||
func TestParseGatewayRequest_AnthropicFastSpeed(t *testing.T) {
|
||||
parsed, err := ParseGatewayRequest(
|
||||
NewRequestBodyRef([]byte(`{"model":"claude-opus-4-8","speed":" FAST "}`)),
|
||||
domain.PlatformAnthropic,
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "fast", parsed.Speed)
|
||||
|
||||
nonAnthropic, err := ParseGatewayRequest(
|
||||
NewRequestBodyRef([]byte(`{"model":"gpt-5.4","speed":"fast"}`)),
|
||||
"responses",
|
||||
)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, nonAnthropic.Speed)
|
||||
}
|
||||
|
||||
func TestParseGatewayRequest_MaxTokens(t *testing.T) {
|
||||
body := []byte(`{"model":"claude-haiku-4-5","max_tokens":1}`)
|
||||
parsed, err := ParseGatewayRequest(NewRequestBodyRef(body), "")
|
||||
|
||||
@@ -619,6 +619,9 @@ type ForwardResult struct {
|
||||
FirstTokenMs *int // 首字时间(流式请求)
|
||||
ClientDisconnect bool // 客户端是否在流式传输过程中断开
|
||||
ReasoningEffort *string
|
||||
// ServiceTier records the billable request tier. OpenAI uses service_tier;
|
||||
// Anthropic speed=fast is normalized to "fast".
|
||||
ServiceTier *string
|
||||
|
||||
// 图片生成计费字段(图片生成模型使用)
|
||||
ImageCount int // 生成的图片数量
|
||||
|
||||
@@ -1157,6 +1157,7 @@ func (s *GatewayService) calculateTokenCost(
|
||||
RequestCount: 1,
|
||||
RateMultiplier: multiplier,
|
||||
PricingAt: pricingAt,
|
||||
ServiceTier: optionalStringValue(result.ServiceTier),
|
||||
Resolver: s.resolver,
|
||||
Resolved: resolved,
|
||||
})
|
||||
@@ -1167,7 +1168,8 @@ func (s *GatewayService) calculateTokenCost(
|
||||
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, Resolver: s.resolver,
|
||||
Tokens: tokens, RequestCount: 1, RateMultiplier: multiplier, PricingAt: pricingAt,
|
||||
ServiceTier: optionalStringValue(result.ServiceTier), Resolver: s.resolver,
|
||||
})
|
||||
} else {
|
||||
cost, err = s.billingService.CalculateCost(billingModel, tokens, multiplier)
|
||||
@@ -1219,6 +1221,7 @@ func (s *GatewayService) buildRecordUsageLog(
|
||||
UpstreamModel: optionalTrimmedStringPtr(result.UpstreamModel),
|
||||
UpstreamResponseModel: optionalTrimmedStringPtr(result.UpstreamResponseModel),
|
||||
UpstreamModelMismatch: upstreamModelMismatch(sentModel, result.UpstreamResponseModel),
|
||||
ServiceTier: result.ServiceTier,
|
||||
ReasoningEffort: result.ReasoningEffort,
|
||||
InboundEndpoint: optionalTrimmedStringPtr(input.InboundEndpoint),
|
||||
UpstreamEndpoint: optionalTrimmedStringPtr(input.UpstreamEndpoint),
|
||||
|
||||
@@ -120,14 +120,8 @@ func (r *ModelPricingResolver) Resolve(ctx context.Context, input PricingInput)
|
||||
resolved.Source = PricingSourceChannel
|
||||
resolved.channelPricing = chPricing
|
||||
r.applyTokenOverrides(chPricing, resolved)
|
||||
if !longContextPricingEnabled {
|
||||
r.applyFirstTokenTier(resolved, chPricing)
|
||||
}
|
||||
} else if input.GroupID != nil && r.channelService != nil {
|
||||
r.applyChannelOverrides(ctx, *input.GroupID, input.Model, resolved)
|
||||
if resolved.Source == PricingSourceChannel && !longContextPricingEnabled {
|
||||
r.applyFirstTokenTier(resolved, resolved.channelPricing)
|
||||
}
|
||||
}
|
||||
|
||||
return resolved
|
||||
@@ -172,20 +166,6 @@ func matchGroupModelPricing(group *Group, model string) *ChannelModelPricing {
|
||||
return wildcard
|
||||
}
|
||||
|
||||
func (r *ModelPricingResolver) applyFirstTokenTier(resolved *ResolvedPricing, config *ChannelModelPricing) {
|
||||
if resolved == nil || len(resolved.Intervals) == 0 {
|
||||
return
|
||||
}
|
||||
first := resolved.Intervals[0]
|
||||
for _, interval := range resolved.Intervals[1:] {
|
||||
if interval.MinTokens < first.MinTokens {
|
||||
first = interval
|
||||
}
|
||||
}
|
||||
resolved.BasePricing = intervalToModelPricing(&first, resolved.SupportsCacheBreakdown, config)
|
||||
resolved.Intervals = nil
|
||||
}
|
||||
|
||||
// resolveBasePricing 从 LiteLLM 或 Fallback 获取基础定价
|
||||
func (r *ModelPricingResolver) resolveBasePricing(model string) (*ModelPricing, string) {
|
||||
pricing, err := r.billingService.GetModelPricing(model)
|
||||
@@ -221,31 +201,6 @@ func (r *ModelPricingResolver) applyChannelOverrides(ctx context.Context, groupI
|
||||
|
||||
// applyTokenOverrides 应用 token 模式的渠道覆盖
|
||||
func (r *ModelPricingResolver) applyTokenOverrides(chPricing *ChannelModelPricing, resolved *ResolvedPricing) {
|
||||
// 过滤掉所有价格字段都为空的无效 interval
|
||||
validIntervals := filterValidIntervals(chPricing.Intervals)
|
||||
|
||||
// 如果有有效的区间定价,使用区间
|
||||
if len(validIntervals) > 0 {
|
||||
resolved.Intervals = validIntervals
|
||||
// 区间不匹配时回退到 BasePricing,也需要覆盖图片价格
|
||||
if resolved.BasePricing == nil {
|
||||
resolved.BasePricing = &ModelPricing{}
|
||||
} else {
|
||||
// 防止修改 fallbackPrices 中的共享指针
|
||||
cloned := *resolved.BasePricing
|
||||
resolved.BasePricing = &cloned
|
||||
}
|
||||
if chPricing.ImageOutputPrice != nil {
|
||||
resolved.BasePricing.ImageOutputPricePerToken = *chPricing.ImageOutputPrice
|
||||
} else {
|
||||
resolved.BasePricing.ImageOutputPricePerToken = 0
|
||||
}
|
||||
resolved.BasePricing.ImageOutputPriceExplicit = true
|
||||
applyChannelImageInputPrice(chPricing, resolved.BasePricing)
|
||||
return
|
||||
}
|
||||
|
||||
// 否则用 flat 字段覆盖 BasePricing
|
||||
if resolved.BasePricing == nil {
|
||||
resolved.BasePricing = &ModelPricing{}
|
||||
} else {
|
||||
@@ -254,25 +209,9 @@ func (r *ModelPricingResolver) applyTokenOverrides(chPricing *ChannelModelPricin
|
||||
resolved.BasePricing = &cloned
|
||||
}
|
||||
|
||||
if chPricing.InputPrice != nil {
|
||||
resolved.BasePricing.InputPricePerToken = *chPricing.InputPrice
|
||||
resolved.BasePricing.InputPricePerTokenPriority = *chPricing.InputPrice
|
||||
}
|
||||
if chPricing.OutputPrice != nil {
|
||||
resolved.BasePricing.OutputPricePerToken = *chPricing.OutputPrice
|
||||
resolved.BasePricing.OutputPricePerTokenPriority = *chPricing.OutputPrice
|
||||
}
|
||||
if chPricing.CacheWritePrice != nil {
|
||||
resolved.BasePricing.CacheCreationPricePerToken = *chPricing.CacheWritePrice
|
||||
resolved.BasePricing.CacheCreationPricePerTokenPriority = *chPricing.CacheWritePrice
|
||||
resolved.BasePricing.CacheCreationPriceExplicit = true
|
||||
resolved.BasePricing.CacheCreation5mPrice = *chPricing.CacheWritePrice
|
||||
resolved.BasePricing.CacheCreation1hPrice = *chPricing.CacheWritePrice
|
||||
}
|
||||
if chPricing.CacheReadPrice != nil {
|
||||
resolved.BasePricing.CacheReadPricePerToken = *chPricing.CacheReadPrice
|
||||
resolved.BasePricing.CacheReadPricePerTokenPriority = *chPricing.CacheReadPrice
|
||||
}
|
||||
applyChannelTokenPriceOverrides(resolved.BasePricing, chPricing)
|
||||
resolved.BasePricing.FastMultiplier = chPricing.FastMultiplier
|
||||
resolved.BasePricing.FlexMultiplier = chPricing.FlexMultiplier
|
||||
// 渠道定价覆盖一切:显式配置则用配置值,未配置则归零(不回退到 LiteLLM)
|
||||
if chPricing.ImageOutputPrice != nil {
|
||||
resolved.BasePricing.ImageOutputPricePerToken = *chPricing.ImageOutputPrice
|
||||
@@ -281,6 +220,9 @@ func (r *ModelPricingResolver) applyTokenOverrides(chPricing *ChannelModelPricin
|
||||
}
|
||||
resolved.BasePricing.ImageOutputPriceExplicit = true
|
||||
applyChannelImageInputPrice(chPricing, resolved.BasePricing)
|
||||
|
||||
// 区间未命中时回退到上面已经应用渠道覆盖的基础价。
|
||||
resolved.Intervals = filterValidIntervals(chPricing.Intervals)
|
||||
}
|
||||
|
||||
// applyChannelImageInputPrice 应用渠道图片输入价:显式配置则用配置值;
|
||||
@@ -311,7 +253,9 @@ func filterValidIntervals(intervals []PricingInterval) []PricingInterval {
|
||||
for _, iv := range intervals {
|
||||
if iv.InputPrice != nil || iv.OutputPrice != nil ||
|
||||
iv.CacheWritePrice != nil || iv.CacheReadPrice != nil ||
|
||||
iv.PerRequestPrice != nil {
|
||||
iv.PerRequestPrice != nil || iv.InputMultiplier != nil ||
|
||||
iv.OutputMultiplier != nil || iv.CacheWriteMultiplier != nil ||
|
||||
iv.CacheReadMultiplier != nil {
|
||||
valid = append(valid, iv)
|
||||
}
|
||||
}
|
||||
@@ -330,32 +274,57 @@ func (r *ModelPricingResolver) GetIntervalPricing(resolved *ResolvedPricing, tot
|
||||
return resolved.BasePricing
|
||||
}
|
||||
|
||||
return intervalToModelPricing(iv, resolved.SupportsCacheBreakdown, resolved.channelPricing)
|
||||
pricing := intervalToModelPricing(iv, resolved.BasePricing, resolved.channelPricing)
|
||||
// BasePricing 为 nil(仅配置区间)时拷贝不到该标志,从 resolved 回填,
|
||||
// 保证 computeCacheCreationCost 的 5m/1h 分档判断不被区间路径吞掉。
|
||||
pricing.SupportsCacheBreakdown = resolved.SupportsCacheBreakdown
|
||||
return pricing
|
||||
}
|
||||
|
||||
// intervalToModelPricing 将区间定价转换为 ModelPricing
|
||||
func intervalToModelPricing(iv *PricingInterval, supportsCacheBreakdown bool, chPricing *ChannelModelPricing) *ModelPricing {
|
||||
pricing := &ModelPricing{
|
||||
SupportsCacheBreakdown: supportsCacheBreakdown,
|
||||
func intervalToModelPricing(iv *PricingInterval, base *ModelPricing, chPricing *ChannelModelPricing) *ModelPricing {
|
||||
pricing := &ModelPricing{}
|
||||
if base != nil {
|
||||
*pricing = *base
|
||||
}
|
||||
applyMultiplier := func(value float64, multiplier *float64) float64 {
|
||||
if multiplier == nil {
|
||||
return value
|
||||
}
|
||||
return value * *multiplier
|
||||
}
|
||||
if iv.InputPrice != nil {
|
||||
pricing.InputPricePerTokenPriority = channelTierOverridePrice(pricing.InputPricePerToken, pricing.InputPricePerTokenPriority, *iv.InputPrice)
|
||||
pricing.InputPricePerToken = *iv.InputPrice
|
||||
pricing.InputPricePerTokenPriority = *iv.InputPrice
|
||||
} else if iv.InputMultiplier != nil {
|
||||
pricing.InputPricePerToken = applyMultiplier(pricing.InputPricePerToken, iv.InputMultiplier)
|
||||
pricing.InputPricePerTokenPriority = applyMultiplier(pricing.InputPricePerTokenPriority, iv.InputMultiplier)
|
||||
}
|
||||
if iv.OutputPrice != nil {
|
||||
pricing.OutputPricePerTokenPriority = channelTierOverridePrice(pricing.OutputPricePerToken, pricing.OutputPricePerTokenPriority, *iv.OutputPrice)
|
||||
pricing.OutputPricePerToken = *iv.OutputPrice
|
||||
pricing.OutputPricePerTokenPriority = *iv.OutputPrice
|
||||
} else if iv.OutputMultiplier != nil {
|
||||
pricing.OutputPricePerToken = applyMultiplier(pricing.OutputPricePerToken, iv.OutputMultiplier)
|
||||
pricing.OutputPricePerTokenPriority = applyMultiplier(pricing.OutputPricePerTokenPriority, iv.OutputMultiplier)
|
||||
}
|
||||
if iv.CacheWritePrice != nil {
|
||||
pricing.CacheCreationPricePerTokenPriority = channelTierOverridePrice(pricing.CacheCreationPricePerToken, pricing.CacheCreationPricePerTokenPriority, *iv.CacheWritePrice)
|
||||
pricing.CacheCreationPricePerToken = *iv.CacheWritePrice
|
||||
pricing.CacheCreationPricePerTokenPriority = *iv.CacheWritePrice
|
||||
pricing.CacheCreationPriceExplicit = true
|
||||
pricing.CacheCreation5mPrice = *iv.CacheWritePrice
|
||||
pricing.CacheCreation1hPrice = *iv.CacheWritePrice
|
||||
} else if iv.CacheWriteMultiplier != nil {
|
||||
pricing.CacheCreationPricePerToken = applyMultiplier(pricing.CacheCreationPricePerToken, iv.CacheWriteMultiplier)
|
||||
pricing.CacheCreationPricePerTokenPriority = applyMultiplier(pricing.CacheCreationPricePerTokenPriority, iv.CacheWriteMultiplier)
|
||||
pricing.CacheCreation5mPrice = applyMultiplier(pricing.CacheCreation5mPrice, iv.CacheWriteMultiplier)
|
||||
pricing.CacheCreation1hPrice = applyMultiplier(pricing.CacheCreation1hPrice, iv.CacheWriteMultiplier)
|
||||
}
|
||||
if iv.CacheReadPrice != nil {
|
||||
pricing.CacheReadPricePerTokenPriority = channelTierOverridePrice(pricing.CacheReadPricePerToken, pricing.CacheReadPricePerTokenPriority, *iv.CacheReadPrice)
|
||||
pricing.CacheReadPricePerToken = *iv.CacheReadPrice
|
||||
pricing.CacheReadPricePerTokenPriority = *iv.CacheReadPrice
|
||||
} else if iv.CacheReadMultiplier != nil {
|
||||
pricing.CacheReadPricePerToken = applyMultiplier(pricing.CacheReadPricePerToken, iv.CacheReadMultiplier)
|
||||
pricing.CacheReadPricePerTokenPriority = applyMultiplier(pricing.CacheReadPricePerTokenPriority, iv.CacheReadMultiplier)
|
||||
}
|
||||
// 渠道定价存在时,ImageOutputPrice 显式覆盖;图片输入价用渠道级配置
|
||||
// (区间不携带图片输入价,与 image_output 一致)。
|
||||
|
||||
@@ -142,7 +142,7 @@ func TestGPT56ExplicitZeroCacheWritePriceIsPreserved(t *testing.T) {
|
||||
})
|
||||
|
||||
t.Run("interval price", func(t *testing.T) {
|
||||
pricing := intervalToModelPricing(&PricingInterval{CacheWritePrice: &zero}, false, nil)
|
||||
pricing := intervalToModelPricing(&PricingInterval{CacheWritePrice: &zero}, &ModelPricing{}, nil)
|
||||
require.True(t, pricing.CacheCreationPriceExplicit)
|
||||
|
||||
cost, err := bs.CalculateCostUnified(CostInput{
|
||||
@@ -261,9 +261,9 @@ func TestResolve_WithChannelOverride_TokenFlat(t *testing.T) {
|
||||
require.Equal(t, "channel", resolved.Source)
|
||||
require.NotNil(t, resolved.BasePricing)
|
||||
require.InDelta(t, 10e-6, resolved.BasePricing.InputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 10e-6, resolved.BasePricing.InputPricePerTokenPriority, 1e-12)
|
||||
require.Zero(t, resolved.BasePricing.InputPricePerTokenPriority)
|
||||
require.InDelta(t, 50e-6, resolved.BasePricing.OutputPricePerToken, 1e-12)
|
||||
require.InDelta(t, 50e-6, resolved.BasePricing.OutputPricePerTokenPriority, 1e-12)
|
||||
require.Zero(t, resolved.BasePricing.OutputPricePerTokenPriority)
|
||||
}
|
||||
|
||||
func TestResolve_WithChannelOverride_TokenPartialOverride(t *testing.T) {
|
||||
@@ -304,10 +304,12 @@ func TestResolve_WithChannelOverride_TokenWithIntervals(t *testing.T) {
|
||||
resolved := r.Resolve(context.Background(), PricingInput{
|
||||
Model: "claude-sonnet-4",
|
||||
GroupID: groupIDPtr(),
|
||||
Group: &Group{LongContextPricingEnabled: false},
|
||||
})
|
||||
|
||||
require.NotNil(t, resolved)
|
||||
require.Equal(t, "channel", resolved.Source)
|
||||
require.False(t, resolved.longContextPricingEnabled)
|
||||
require.Len(t, resolved.Intervals, 2)
|
||||
|
||||
// GetIntervalPricing should use channel intervals
|
||||
@@ -532,6 +534,7 @@ func TestGetIntervalPricing_ChannelIntervalsNoMatch(t *testing.T) {
|
||||
Platform: "anthropic",
|
||||
Models: []string{"claude-sonnet-4"},
|
||||
BillingMode: BillingModeToken,
|
||||
InputPrice: testPtrFloat64(4e-6),
|
||||
Intervals: []PricingInterval{
|
||||
// Only covers tokens > 50000
|
||||
{MinTokens: 50000, MaxTokens: testPtrInt(200000), InputPrice: testPtrFloat64(9e-6)},
|
||||
@@ -545,10 +548,10 @@ func TestGetIntervalPricing_ChannelIntervalsNoMatch(t *testing.T) {
|
||||
|
||||
// Token count 1000 doesn't match any interval (1000 <= 50000 minTokens)
|
||||
pricing := r.GetIntervalPricing(resolved, 1000)
|
||||
// Should fall back to BasePricing (from the billing service fallback)
|
||||
// Should fall back to BasePricing after applying the channel default.
|
||||
require.NotNil(t, pricing)
|
||||
require.Equal(t, resolved.BasePricing, pricing)
|
||||
require.InDelta(t, 3e-6, pricing.InputPricePerToken, 1e-12) // original base price
|
||||
require.InDelta(t, 4e-6, pricing.InputPricePerToken, 1e-12)
|
||||
}
|
||||
|
||||
// ===========================================================================
|
||||
@@ -689,6 +692,13 @@ func TestFilterValidIntervals(t *testing.T) {
|
||||
},
|
||||
wantLen: 1,
|
||||
},
|
||||
{
|
||||
name: "interval with only multiplier kept",
|
||||
intervals: []PricingInterval{
|
||||
{MinTokens: 272000, InputMultiplier: testPtrFloat64(2)},
|
||||
},
|
||||
wantLen: 1,
|
||||
},
|
||||
{
|
||||
name: "mixed valid and invalid",
|
||||
intervals: []PricingInterval{
|
||||
|
||||
@@ -1148,7 +1148,7 @@ func TestOpenAIGatewayServiceRecordUsage_GPT56SeparatesCacheWriteForBillingAndSt
|
||||
require.InDelta(t, usageRepo.lastLog.TotalCost*1.1, usageRepo.lastLog.ActualCost, 1e-12)
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingDisabledByDefault(t *testing.T) {
|
||||
func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingDisabledWhenGroupAndAccountOff(t *testing.T) {
|
||||
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
|
||||
userRepo := &openAIRecordUsageUserRepoStub{}
|
||||
subRepo := &openAIRecordUsageSubRepoStub{}
|
||||
@@ -1164,7 +1164,7 @@ func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingDisabledByDefaul
|
||||
Model: "gpt-5.4-2026-03-05",
|
||||
Duration: time.Second,
|
||||
},
|
||||
APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1014, true),
|
||||
APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1014, false),
|
||||
User: &User{ID: 2014},
|
||||
Account: &Account{ID: 3014, Platform: PlatformOpenAI},
|
||||
})
|
||||
@@ -1219,7 +1219,7 @@ func TestOpenAIGatewayServiceRecordUsage_Gpt54LongContextBillingEnabledPerAccoun
|
||||
require.True(t, usageRepo.lastLog.LongContextBillingApplied)
|
||||
}
|
||||
|
||||
func TestOpenAIGatewayServiceRecordUsage_GroupAndAccountLongContextMustBothAllow(t *testing.T) {
|
||||
func TestOpenAIGatewayServiceRecordUsage_GroupOrAccountLongContextAllows(t *testing.T) {
|
||||
tokens := OpenAIUsage{InputTokens: 300000, OutputTokens: 2000}
|
||||
baseInput := 300000 * 2.5e-6
|
||||
baseOutput := 2000 * 15e-6
|
||||
@@ -1234,9 +1234,9 @@ func TestOpenAIGatewayServiceRecordUsage_GroupAndAccountLongContextMustBothAllow
|
||||
Account: &Account{ID: 3020, Platform: PlatformOpenAI},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.False(t, usageRepo.lastLog.LongContextBillingApplied)
|
||||
require.InDelta(t, baseInput, usageRepo.lastLog.InputCost, 1e-10)
|
||||
require.InDelta(t, baseOutput, usageRepo.lastLog.OutputCost, 1e-10)
|
||||
require.True(t, usageRepo.lastLog.LongContextBillingApplied)
|
||||
require.InDelta(t, baseInput*2, usageRepo.lastLog.InputCost, 1e-10)
|
||||
require.InDelta(t, baseOutput*1.5, usageRepo.lastLog.OutputCost, 1e-10)
|
||||
})
|
||||
|
||||
t.Run("group off account on", func(t *testing.T) {
|
||||
@@ -1252,8 +1252,9 @@ func TestOpenAIGatewayServiceRecordUsage_GroupAndAccountLongContextMustBothAllow
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.False(t, usageRepo.lastLog.LongContextBillingApplied)
|
||||
require.InDelta(t, baseInput, usageRepo.lastLog.InputCost, 1e-10)
|
||||
require.True(t, usageRepo.lastLog.LongContextBillingApplied)
|
||||
require.InDelta(t, baseInput*2, usageRepo.lastLog.InputCost, 1e-10)
|
||||
require.InDelta(t, baseOutput*1.5, usageRepo.lastLog.OutputCost, 1e-10)
|
||||
})
|
||||
|
||||
t.Run("group on account on", func(t *testing.T) {
|
||||
@@ -1361,7 +1362,7 @@ func TestOpenAIGatewayServiceRecordUsage_SparkShadowUsesCurrentParentBillingSett
|
||||
Model: "gpt-5.4-2026-03-05",
|
||||
Duration: time.Second,
|
||||
},
|
||||
APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1016, true),
|
||||
APIKey: openAIRecordUsageAPIKeyWithGroup(svc, 1016, false),
|
||||
User: &User{ID: 2016},
|
||||
Account: &Account{
|
||||
ID: 3016,
|
||||
|
||||
@@ -128,7 +128,8 @@ type UsageLog struct {
|
||||
BillingTier *string
|
||||
// BillingMode 计费模式:token/image
|
||||
BillingMode *string
|
||||
// ServiceTier records the OpenAI service tier used for billing, e.g. "priority" / "flex".
|
||||
// ServiceTier records the billable request tier, e.g. OpenAI "priority" / "flex"
|
||||
// or Anthropic "fast".
|
||||
ServiceTier *string
|
||||
// ReasoningEffort is the request's reasoning effort level.
|
||||
// OpenAI: "low" / "medium" / "high" / "xhigh"; Claude: "low" / "medium" / "high" / "max".
|
||||
|
||||
@@ -10,6 +10,13 @@ func optionalTrimmedStringPtr(raw string) *string {
|
||||
return &trimmed
|
||||
}
|
||||
|
||||
func optionalStringValue(value *string) string {
|
||||
if value == nil {
|
||||
return ""
|
||||
}
|
||||
return strings.TrimSpace(*value)
|
||||
}
|
||||
|
||||
func forwardResultBillingModel(requestedModel, upstreamModel string) string {
|
||||
if trimmed := strings.TrimSpace(requestedModel); trimmed != "" {
|
||||
return trimmed
|
||||
|
||||
@@ -0,0 +1,56 @@
|
||||
ALTER TABLE channel_model_pricing
|
||||
ADD COLUMN IF NOT EXISTS fast_multiplier NUMERIC(12,6),
|
||||
ADD COLUMN IF NOT EXISTS flex_multiplier NUMERIC(12,6);
|
||||
|
||||
ALTER TABLE channel_pricing_intervals
|
||||
ADD COLUMN IF NOT EXISTS input_multiplier NUMERIC(12,6),
|
||||
ADD COLUMN IF NOT EXISTS output_multiplier NUMERIC(12,6),
|
||||
ADD COLUMN IF NOT EXISTS cache_write_multiplier NUMERIC(12,6),
|
||||
ADD COLUMN IF NOT EXISTS cache_read_multiplier NUMERIC(12,6);
|
||||
|
||||
DO $$
|
||||
BEGIN
|
||||
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'channel_model_pricing_fast_multiplier_positive' AND conrelid = 'channel_model_pricing'::regclass) THEN
|
||||
ALTER TABLE channel_model_pricing
|
||||
ADD CONSTRAINT channel_model_pricing_fast_multiplier_positive
|
||||
CHECK (fast_multiplier IS NULL OR fast_multiplier > 0);
|
||||
END IF;
|
||||
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'channel_model_pricing_flex_multiplier_positive' AND conrelid = 'channel_model_pricing'::regclass) THEN
|
||||
ALTER TABLE channel_model_pricing
|
||||
ADD CONSTRAINT channel_model_pricing_flex_multiplier_positive
|
||||
CHECK (flex_multiplier IS NULL OR flex_multiplier > 0);
|
||||
END IF;
|
||||
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'channel_pricing_intervals_input_multiplier_positive' AND conrelid = 'channel_pricing_intervals'::regclass) THEN
|
||||
ALTER TABLE channel_pricing_intervals
|
||||
ADD CONSTRAINT channel_pricing_intervals_input_multiplier_positive
|
||||
CHECK (input_multiplier IS NULL OR input_multiplier > 0);
|
||||
END IF;
|
||||
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'channel_pricing_intervals_output_multiplier_positive' AND conrelid = 'channel_pricing_intervals'::regclass) THEN
|
||||
ALTER TABLE channel_pricing_intervals
|
||||
ADD CONSTRAINT channel_pricing_intervals_output_multiplier_positive
|
||||
CHECK (output_multiplier IS NULL OR output_multiplier > 0);
|
||||
END IF;
|
||||
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'channel_pricing_intervals_cache_write_multiplier_positive' AND conrelid = 'channel_pricing_intervals'::regclass) THEN
|
||||
ALTER TABLE channel_pricing_intervals
|
||||
ADD CONSTRAINT channel_pricing_intervals_cache_write_multiplier_positive
|
||||
CHECK (cache_write_multiplier IS NULL OR cache_write_multiplier > 0);
|
||||
END IF;
|
||||
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'channel_pricing_intervals_cache_read_multiplier_positive' AND conrelid = 'channel_pricing_intervals'::regclass) THEN
|
||||
ALTER TABLE channel_pricing_intervals
|
||||
ADD CONSTRAINT channel_pricing_intervals_cache_read_multiplier_positive
|
||||
CHECK (cache_read_multiplier IS NULL OR cache_read_multiplier > 0);
|
||||
END IF;
|
||||
END $$;
|
||||
|
||||
COMMENT ON COLUMN channel_model_pricing.fast_multiplier IS
|
||||
'Fast/priority service tier multiplier applied to the selected standard channel price';
|
||||
COMMENT ON COLUMN channel_model_pricing.flex_multiplier IS
|
||||
'Flex service tier multiplier applied to the selected standard channel price';
|
||||
COMMENT ON COLUMN channel_pricing_intervals.input_multiplier IS
|
||||
'Interval input multiplier applied to the channel base input price when input_price is NULL';
|
||||
COMMENT ON COLUMN channel_pricing_intervals.output_multiplier IS
|
||||
'Interval output multiplier applied to the channel base output price when output_price is NULL';
|
||||
COMMENT ON COLUMN channel_pricing_intervals.cache_write_multiplier IS
|
||||
'Interval cache-write multiplier applied to the channel base cache-write price when cache_write_price is NULL';
|
||||
COMMENT ON COLUMN channel_pricing_intervals.cache_read_multiplier IS
|
||||
'Interval cache-read multiplier applied to the channel base cache-read price when cache_read_price is NULL';
|
||||
@@ -0,0 +1,43 @@
|
||||
package migrations
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestChannelPricingMultipliersMigration(t *testing.T) {
|
||||
content, err := FS.ReadFile("228_channel_pricing_multipliers.sql")
|
||||
require.NoError(t, err)
|
||||
|
||||
sql := strings.Join(strings.Fields(string(content)), " ")
|
||||
for _, column := range []string{
|
||||
"fast_multiplier NUMERIC(12,6)",
|
||||
"flex_multiplier NUMERIC(12,6)",
|
||||
"input_multiplier NUMERIC(12,6)",
|
||||
"output_multiplier NUMERIC(12,6)",
|
||||
"cache_write_multiplier NUMERIC(12,6)",
|
||||
"cache_read_multiplier NUMERIC(12,6)",
|
||||
} {
|
||||
require.Contains(t, sql, "ADD COLUMN IF NOT EXISTS "+column)
|
||||
}
|
||||
|
||||
constraints := []struct {
|
||||
table string
|
||||
name string
|
||||
column string
|
||||
}{
|
||||
{"channel_model_pricing", "channel_model_pricing_fast_multiplier_positive", "fast_multiplier"},
|
||||
{"channel_model_pricing", "channel_model_pricing_flex_multiplier_positive", "flex_multiplier"},
|
||||
{"channel_pricing_intervals", "channel_pricing_intervals_input_multiplier_positive", "input_multiplier"},
|
||||
{"channel_pricing_intervals", "channel_pricing_intervals_output_multiplier_positive", "output_multiplier"},
|
||||
{"channel_pricing_intervals", "channel_pricing_intervals_cache_write_multiplier_positive", "cache_write_multiplier"},
|
||||
{"channel_pricing_intervals", "channel_pricing_intervals_cache_read_multiplier_positive", "cache_read_multiplier"},
|
||||
}
|
||||
for _, constraint := range constraints {
|
||||
require.Contains(t, sql, "conname = '"+constraint.name+"' AND conrelid = '"+constraint.table+"'::regclass")
|
||||
require.Contains(t, sql, "ALTER TABLE "+constraint.table+" ADD CONSTRAINT "+constraint.name+
|
||||
" CHECK ("+constraint.column+" IS NULL OR "+constraint.column+" > 0)")
|
||||
}
|
||||
}
|
||||
@@ -17,6 +17,10 @@ export interface PricingInterval {
|
||||
output_price: number | null
|
||||
cache_write_price: number | null
|
||||
cache_read_price: number | null
|
||||
input_multiplier: number | null
|
||||
output_multiplier: number | null
|
||||
cache_write_multiplier: number | null
|
||||
cache_read_multiplier: number | null
|
||||
per_request_price: number | null
|
||||
sort_order: number
|
||||
}
|
||||
@@ -41,6 +45,8 @@ export interface ChannelModelPricing {
|
||||
output_price: number | null
|
||||
cache_write_price: number | null
|
||||
cache_read_price: number | null
|
||||
fast_multiplier?: number | null
|
||||
flex_multiplier?: number | null
|
||||
image_input_price: number | null
|
||||
image_output_price: number | null
|
||||
per_request_price: number | null
|
||||
|
||||
@@ -2958,6 +2958,7 @@
|
||||
<!-- OpenAI WS Mode 三态(off/ctx_pool/passthrough) -->
|
||||
<div
|
||||
v-if="form.platform === 'openai' && (accountCategory === 'oauth-based' || accountCategory === 'apikey')"
|
||||
data-testid="create-openai-ws-mode"
|
||||
class="border-t border-gray-200 pt-4 dark:border-dark-600"
|
||||
>
|
||||
<div class="flex items-center justify-between">
|
||||
@@ -3044,9 +3045,9 @@
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- OpenAI OAuth Codex 官方客户端限制开关 -->
|
||||
<!-- OpenAI API 长上下文计费开关 -->
|
||||
<div
|
||||
v-if="form.platform === 'openai' && (accountCategory === 'oauth-based' || accountCategory === 'apikey')"
|
||||
v-if="form.platform === 'openai' && !hideAccountLongContextBilling && (accountCategory === 'oauth-based' || accountCategory === 'apikey')"
|
||||
class="border-t border-gray-200 pt-4 dark:border-dark-600"
|
||||
>
|
||||
<div class="flex items-center justify-between gap-4">
|
||||
@@ -3740,6 +3741,7 @@ import Toggle from '@/components/common/Toggle.vue'
|
||||
import GrokBaseUrlPresets from '@/components/account/GrokBaseUrlPresets.vue'
|
||||
import CnBaseUrlPresets from '@/components/account/CnBaseUrlPresets.vue'
|
||||
import HeaderOverrideEditor from '@/components/account/HeaderOverrideEditor.vue'
|
||||
import { allSelectedGroupsEnableLongContextPricing } from '@/components/account/longContextBilling'
|
||||
import {
|
||||
applyAntigravityProjectID,
|
||||
applyHeaderOverride,
|
||||
@@ -3857,6 +3859,10 @@ const emit = defineEmits<{
|
||||
|
||||
const appStore = useAppStore()
|
||||
|
||||
const hideAccountLongContextBilling = computed(() => {
|
||||
return allSelectedGroupsEnableLongContextPricing(form.group_ids, props.groups)
|
||||
})
|
||||
|
||||
// OAuth composables
|
||||
const oauth = useAccountOAuth() // For Anthropic OAuth
|
||||
const openaiOAuth = useOpenAIOAuth() // For OpenAI OAuth
|
||||
|
||||
@@ -1966,7 +1966,7 @@
|
||||
|
||||
<!-- OpenAI API 长上下文计费开关 -->
|
||||
<div
|
||||
v-if="account?.platform === 'openai' && !isSparkShadow && (account?.type === 'oauth' || account?.type === 'setup-token' || account?.type === 'apikey')"
|
||||
v-if="account?.platform === 'openai' && !isSparkShadow && !hideAccountLongContextBilling && (account?.type === 'oauth' || account?.type === 'setup-token' || account?.type === 'apikey')"
|
||||
class="border-t border-gray-200 pt-4 dark:border-dark-600"
|
||||
>
|
||||
<div class="flex items-center justify-between gap-4">
|
||||
@@ -2811,6 +2811,7 @@ import {
|
||||
} from '@/components/account/credentialsBuilder'
|
||||
import { formatDateTime, formatDateTimeLocalInput, parseDateTimeLocalInput } from '@/utils/format'
|
||||
import { createStableObjectKeyResolver } from '@/utils/stableObjectKey'
|
||||
import { allSelectedGroupsEnableLongContextPricing } from '@/components/account/longContextBilling'
|
||||
import { VERTEX_LOCATION_OPTIONS } from '@/constants/account'
|
||||
import {
|
||||
OPENAI_WS_MODE_CTX_POOL,
|
||||
@@ -2851,6 +2852,10 @@ const authStore = useAuthStore()
|
||||
// 故隐藏代理选择器。
|
||||
const isSparkShadow = computed(() => props.account?.parent_account_id != null)
|
||||
|
||||
const hideAccountLongContextBilling = computed(() => {
|
||||
return allSelectedGroupsEnableLongContextPricing(form.group_ids, props.groups)
|
||||
})
|
||||
|
||||
const handleOllamaCloudUsageUpdated = (state: OllamaCloudUsageState) => {
|
||||
if (props.account) emit('updated', { ...props.account, ollama_cloud_usage: state })
|
||||
}
|
||||
|
||||
@@ -7,11 +7,13 @@ const {
|
||||
probeUpstreamBillingMock,
|
||||
importCodexSessionMock,
|
||||
createOpenAICodexPATMock,
|
||||
authIsSimpleMode,
|
||||
} = vi.hoisted(() => ({
|
||||
createAccountMock: vi.fn(),
|
||||
probeUpstreamBillingMock: vi.fn(),
|
||||
importCodexSessionMock: vi.fn(),
|
||||
createOpenAICodexPATMock: vi.fn(),
|
||||
authIsSimpleMode: { value: true },
|
||||
}))
|
||||
|
||||
vi.mock('@/stores/app', () => ({
|
||||
@@ -23,7 +25,11 @@ vi.mock('@/stores/app', () => ({
|
||||
}))
|
||||
|
||||
vi.mock('@/stores/auth', () => ({
|
||||
useAuthStore: () => ({ isSimpleMode: true }),
|
||||
useAuthStore: () => ({
|
||||
get isSimpleMode() {
|
||||
return authIsSimpleMode.value
|
||||
},
|
||||
}),
|
||||
}))
|
||||
|
||||
vi.mock('@/api/admin', () => ({
|
||||
@@ -84,9 +90,29 @@ const OAuthAuthorizationFlowStub = defineComponent({
|
||||
`,
|
||||
})
|
||||
|
||||
function mountModal() {
|
||||
const GroupSelectorStub = defineComponent({
|
||||
name: 'GroupSelector',
|
||||
props: {
|
||||
modelValue: {
|
||||
type: Array,
|
||||
default: () => [],
|
||||
},
|
||||
},
|
||||
emits: ['update:modelValue'],
|
||||
template: `
|
||||
<button
|
||||
type="button"
|
||||
data-testid="select-pricing-groups"
|
||||
@click="$emit('update:modelValue', [1, 2])"
|
||||
>
|
||||
groups
|
||||
</button>
|
||||
`,
|
||||
})
|
||||
|
||||
function mountModal(groups: any[] = []) {
|
||||
return mount(CreateAccountModal, {
|
||||
props: { show: true, proxies: [], groups: [] },
|
||||
props: { show: true, proxies: [], groups },
|
||||
global: {
|
||||
stubs: {
|
||||
BaseDialog: BaseDialogStub,
|
||||
@@ -97,7 +123,7 @@ function mountModal() {
|
||||
PlatformIcon: true,
|
||||
ProxySelector: true,
|
||||
ProxyAdBanner: true,
|
||||
GroupSelector: true,
|
||||
GroupSelector: GroupSelectorStub,
|
||||
ModelWhitelistSelector: true,
|
||||
QuotaLimitCard: true,
|
||||
},
|
||||
@@ -147,6 +173,7 @@ async function openCodexImportStep(toggleClicks = 0) {
|
||||
|
||||
describe('CreateAccountModal OpenAI long-context billing', () => {
|
||||
beforeEach(() => {
|
||||
authIsSimpleMode.value = true
|
||||
createAccountMock.mockReset().mockResolvedValue({ id: 42, platform: 'openai', type: 'apikey' })
|
||||
probeUpstreamBillingMock.mockReset().mockResolvedValue({})
|
||||
importCodexSessionMock.mockReset().mockResolvedValue({
|
||||
@@ -160,6 +187,34 @@ describe('CreateAccountModal OpenAI long-context billing', () => {
|
||||
createOpenAICodexPATMock.mockReset().mockResolvedValue({})
|
||||
})
|
||||
|
||||
it('hides only the redundant account toggle when every selected group enables tier pricing', async () => {
|
||||
authIsSimpleMode.value = false
|
||||
const wrapper = mountModal([
|
||||
{ id: 1, long_context_pricing_enabled: true },
|
||||
{ id: 2, long_context_pricing_enabled: true },
|
||||
])
|
||||
|
||||
await selectButtonByText(wrapper, 'OpenAI')
|
||||
await wrapper.get('[data-testid="select-pricing-groups"]').trigger('click')
|
||||
|
||||
expect(wrapper.find('[data-testid="openai-long-context-billing-toggle"]').exists()).toBe(false)
|
||||
expect(wrapper.find('[data-testid="create-openai-ws-mode"]').exists()).toBe(true)
|
||||
})
|
||||
|
||||
it('keeps the account toggle when any selected group disables tier pricing', async () => {
|
||||
authIsSimpleMode.value = false
|
||||
const wrapper = mountModal([
|
||||
{ id: 1, long_context_pricing_enabled: true },
|
||||
{ id: 2, long_context_pricing_enabled: false },
|
||||
])
|
||||
|
||||
await selectButtonByText(wrapper, 'OpenAI')
|
||||
await wrapper.get('[data-testid="select-pricing-groups"]').trigger('click')
|
||||
|
||||
expect(wrapper.find('[data-testid="openai-long-context-billing-toggle"]').exists()).toBe(true)
|
||||
expect(wrapper.find('[data-testid="create-openai-ws-mode"]').exists()).toBe(true)
|
||||
})
|
||||
|
||||
it('sends false explicitly for normal OpenAI account creation by default', async () => {
|
||||
await submitApiKeyAccount('openai')
|
||||
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
import { describe, expect, it } from 'vitest'
|
||||
|
||||
import { allSelectedGroupsEnableLongContextPricing } from '../longContextBilling'
|
||||
|
||||
const group = (id: number, enabled: boolean) => ({
|
||||
id,
|
||||
long_context_pricing_enabled: enabled,
|
||||
}) as any
|
||||
|
||||
describe('allSelectedGroupsEnableLongContextPricing', () => {
|
||||
it('hides the account override when every selected group enables tier pricing', () => {
|
||||
expect(allSelectedGroupsEnableLongContextPricing(
|
||||
[1, 2],
|
||||
[group(1, true), group(2, true)]
|
||||
)).toBe(true)
|
||||
})
|
||||
|
||||
it('keeps the account override when any selected group disables tier pricing', () => {
|
||||
expect(allSelectedGroupsEnableLongContextPricing(
|
||||
[1, 2],
|
||||
[group(1, true), group(2, false)]
|
||||
)).toBe(false)
|
||||
})
|
||||
|
||||
it('keeps the account override when selection is empty or group data is incomplete', () => {
|
||||
expect(allSelectedGroupsEnableLongContextPricing([], [group(1, true)])).toBe(false)
|
||||
expect(allSelectedGroupsEnableLongContextPricing([1, 2], [group(1, true)])).toBe(false)
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,11 @@
|
||||
import type { AdminGroup } from '@/types'
|
||||
|
||||
export function allSelectedGroupsEnableLongContextPricing(
|
||||
groupIds: number[],
|
||||
groups: AdminGroup[]
|
||||
): boolean {
|
||||
if (groupIds.length === 0) return false
|
||||
const selectedGroups = groups.filter(group => groupIds.includes(group.id))
|
||||
return selectedGroups.length === groupIds.length &&
|
||||
selectedGroups.every(group => group.long_context_pricing_enabled === true)
|
||||
}
|
||||
@@ -3,35 +3,59 @@
|
||||
:class="isEmpty ? 'border-red-400 bg-red-50 dark:border-red-500 dark:bg-red-950/20' : 'border-gray-200 bg-white dark:border-dark-500 dark:bg-dark-700'">
|
||||
<!-- Token mode: context range + prices ($/MTok) -->
|
||||
<template v-if="mode === 'token'">
|
||||
<div class="w-20">
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.minTokens') }}</label>
|
||||
<input :value="interval.min_tokens" @input="emitField('min_tokens', toInt(($event.target as HTMLInputElement).value))"
|
||||
type="number" min="0" class="input mt-0.5 text-xs" />
|
||||
</div>
|
||||
<div class="w-20">
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.maxTokens') }} <span class="text-gray-300">{{ t('admin.channels.form.inclusive') }}</span></label>
|
||||
<input :value="interval.max_tokens ?? ''" @input="emitField('max_tokens', toIntOrNull(($event.target as HTMLInputElement).value))"
|
||||
type="number" min="0" class="input mt-0.5 text-xs" :placeholder="'∞'" />
|
||||
</div>
|
||||
<div class="flex-1">
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.inputPrice') }} <span v-if="isEmpty" class="text-red-500">*</span> <span class="text-gray-300">$/M</span></label>
|
||||
<input :value="interval.input_price" @input="emitField('input_price', ($event.target as HTMLInputElement).value)"
|
||||
type="number" step="any" min="0" class="input mt-0.5 text-xs" />
|
||||
</div>
|
||||
<div class="flex-1">
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.outputPrice') }} <span v-if="isEmpty" class="text-red-500">*</span> <span class="text-gray-300">$/M</span></label>
|
||||
<input :value="interval.output_price" @input="emitField('output_price', ($event.target as HTMLInputElement).value)"
|
||||
type="number" step="any" min="0" class="input mt-0.5 text-xs" />
|
||||
</div>
|
||||
<div class="flex-1">
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.cacheWritePriceShort') }} <span class="text-gray-300">$/M</span></label>
|
||||
<input :value="interval.cache_write_price" @input="emitField('cache_write_price', ($event.target as HTMLInputElement).value)"
|
||||
type="number" step="any" min="0" class="input mt-0.5 text-xs" />
|
||||
</div>
|
||||
<div class="flex-1">
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.cacheReadPriceShort') }} <span class="text-gray-300">$/M</span></label>
|
||||
<input :value="interval.cache_read_price" @input="emitField('cache_read_price', ($event.target as HTMLInputElement).value)"
|
||||
type="number" step="any" min="0" class="input mt-0.5 text-xs" />
|
||||
<div class="grid min-w-0 flex-1 grid-cols-2 gap-2 sm:grid-cols-4 xl:grid-cols-6">
|
||||
<div>
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.minTokens') }}</label>
|
||||
<input :value="interval.min_tokens" @input="emitField('min_tokens', toInt(($event.target as HTMLInputElement).value))"
|
||||
type="number" min="0" class="input mt-0.5 text-xs" />
|
||||
</div>
|
||||
<div>
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.maxTokens') }} <span class="text-gray-300">{{ t('admin.channels.form.inclusive') }}</span></label>
|
||||
<input :value="interval.max_tokens ?? ''" @input="emitField('max_tokens', toIntOrNull(($event.target as HTMLInputElement).value))"
|
||||
type="number" min="0" class="input mt-0.5 text-xs" :placeholder="'∞'" />
|
||||
</div>
|
||||
<div>
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.inputPrice') }} <span v-if="isEmpty" class="text-red-500">*</span> <span class="text-gray-300">$/M</span></label>
|
||||
<input :value="interval.input_price" @input="emitField('input_price', ($event.target as HTMLInputElement).value)"
|
||||
type="number" step="any" min="0" class="input mt-0.5 text-xs" />
|
||||
</div>
|
||||
<div>
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.outputPrice') }} <span v-if="isEmpty" class="text-red-500">*</span> <span class="text-gray-300">$/M</span></label>
|
||||
<input :value="interval.output_price" @input="emitField('output_price', ($event.target as HTMLInputElement).value)"
|
||||
type="number" step="any" min="0" class="input mt-0.5 text-xs" />
|
||||
</div>
|
||||
<div>
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.cacheWritePriceShort') }} <span class="text-gray-300">$/M</span></label>
|
||||
<input :value="interval.cache_write_price" @input="emitField('cache_write_price', ($event.target as HTMLInputElement).value)"
|
||||
type="number" step="any" min="0" class="input mt-0.5 text-xs" />
|
||||
</div>
|
||||
<div>
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.cacheReadPriceShort') }} <span class="text-gray-300">$/M</span></label>
|
||||
<input :value="interval.cache_read_price" @input="emitField('cache_read_price', ($event.target as HTMLInputElement).value)"
|
||||
type="number" step="any" min="0" class="input mt-0.5 text-xs" />
|
||||
</div>
|
||||
<template v-if="enableMultipliers">
|
||||
<div>
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.inputMultiplier') }}</label>
|
||||
<input :value="interval.input_multiplier" @input="emitField('input_multiplier', ($event.target as HTMLInputElement).value)"
|
||||
type="number" step="any" min="0.000001" class="input mt-0.5 text-xs" />
|
||||
</div>
|
||||
<div>
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.outputMultiplier') }}</label>
|
||||
<input :value="interval.output_multiplier" @input="emitField('output_multiplier', ($event.target as HTMLInputElement).value)"
|
||||
type="number" step="any" min="0.000001" class="input mt-0.5 text-xs" />
|
||||
</div>
|
||||
<div>
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.cacheWriteMultiplier') }}</label>
|
||||
<input :value="interval.cache_write_multiplier" @input="emitField('cache_write_multiplier', ($event.target as HTMLInputElement).value)"
|
||||
type="number" step="any" min="0.000001" class="input mt-0.5 text-xs" />
|
||||
</div>
|
||||
<div>
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.cacheReadMultiplier') }}</label>
|
||||
<input :value="interval.cache_read_multiplier" @input="emitField('cache_read_multiplier', ($event.target as HTMLInputElement).value)"
|
||||
type="number" step="any" min="0.000001" class="input mt-0.5 text-xs" />
|
||||
</div>
|
||||
</template>
|
||||
</div>
|
||||
</template>
|
||||
|
||||
@@ -79,6 +103,7 @@ const { t } = useI18n()
|
||||
const props = defineProps<{
|
||||
interval: IntervalFormEntry
|
||||
mode: BillingMode
|
||||
enableMultipliers?: boolean
|
||||
}>()
|
||||
|
||||
const emit = defineEmits<{
|
||||
@@ -93,6 +118,10 @@ const isEmpty = computed(() => {
|
||||
(iv.output_price == null || iv.output_price === '') &&
|
||||
(iv.cache_write_price == null || iv.cache_write_price === '') &&
|
||||
(iv.cache_read_price == null || iv.cache_read_price === '') &&
|
||||
(iv.input_multiplier == null || iv.input_multiplier === '') &&
|
||||
(iv.output_multiplier == null || iv.output_multiplier === '') &&
|
||||
(iv.cache_write_multiplier == null || iv.cache_write_multiplier === '') &&
|
||||
(iv.cache_read_multiplier == null || iv.cache_read_multiplier === '') &&
|
||||
(iv.per_request_price == null || iv.per_request_price === '')
|
||||
})
|
||||
|
||||
|
||||
@@ -139,7 +139,20 @@
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Token intervals (channel-only; group long-context uses official presets) -->
|
||||
<div v-if="enableTierMultipliers" class="mt-3 grid max-w-md grid-cols-2 gap-2">
|
||||
<div>
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.fastMultiplier') }}</label>
|
||||
<input :value="entry.fast_multiplier" @input="emitField('fast_multiplier', ($event.target as HTMLInputElement).value)"
|
||||
type="number" step="any" min="0.000001" class="input mt-0.5 text-sm" :placeholder="t('admin.channels.form.multiplierPlaceholder')" />
|
||||
</div>
|
||||
<div>
|
||||
<label class="text-xs text-gray-400">{{ t('admin.channels.form.flexMultiplier') }}</label>
|
||||
<input :value="entry.flex_multiplier" @input="emitField('flex_multiplier', ($event.target as HTMLInputElement).value)"
|
||||
type="number" step="any" min="0.000001" class="input mt-0.5 text-sm" :placeholder="t('admin.channels.form.multiplierPlaceholder')" />
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<!-- Channel token intervals; the group long-context toggle controls whether tiers apply. -->
|
||||
<div v-if="!hideTokenIntervals" class="mt-3">
|
||||
<div class="flex items-center justify-between">
|
||||
<label class="text-xs font-medium text-gray-500 dark:text-gray-400">
|
||||
@@ -156,6 +169,7 @@
|
||||
:key="idx"
|
||||
:interval="iv"
|
||||
:mode="entry.billing_mode"
|
||||
:enable-multipliers="enableTierMultipliers"
|
||||
@update="updateInterval(idx, $event)"
|
||||
@remove="removeInterval(idx)"
|
||||
/>
|
||||
@@ -262,9 +276,11 @@ const props = withDefaults(defineProps<{
|
||||
platform?: string
|
||||
hideTokenIntervals?: boolean
|
||||
enableTimePricing?: boolean
|
||||
enableTierMultipliers?: boolean
|
||||
}>(), {
|
||||
hideTokenIntervals: false,
|
||||
enableTimePricing: false,
|
||||
enableTierMultipliers: false,
|
||||
})
|
||||
|
||||
const emit = defineEmits<{
|
||||
@@ -297,6 +313,8 @@ function addInterval() {
|
||||
min_tokens: 0, max_tokens: null, tier_label: '',
|
||||
input_price: null, output_price: null, cache_write_price: null,
|
||||
cache_read_price: null, per_request_price: null,
|
||||
input_multiplier: null, output_multiplier: null,
|
||||
cache_write_multiplier: null, cache_read_multiplier: null,
|
||||
sort_order: intervals.length
|
||||
})
|
||||
emit('update', { ...props.entry, intervals })
|
||||
@@ -311,6 +329,8 @@ function addMediaTier() {
|
||||
min_tokens: 0, max_tokens: null, tier_label: labels[intervals.length] || '',
|
||||
input_price: null, output_price: null, cache_write_price: null,
|
||||
cache_read_price: null, per_request_price: null,
|
||||
input_multiplier: null, output_multiplier: null,
|
||||
cache_write_multiplier: null, cache_read_multiplier: null,
|
||||
sort_order: intervals.length
|
||||
})
|
||||
emit('update', { ...props.entry, intervals })
|
||||
|
||||
@@ -16,6 +16,8 @@ function createEntry(billingMode: PricingFormEntry['billing_mode'] = 'token'): P
|
||||
output_price: null,
|
||||
cache_write_price: null,
|
||||
cache_read_price: null,
|
||||
fast_multiplier: null,
|
||||
flex_multiplier: null,
|
||||
image_input_price: null,
|
||||
image_output_price: null,
|
||||
per_request_price: null,
|
||||
@@ -69,3 +71,16 @@ describe('PricingEntryCard time pricing visibility', () => {
|
||||
expect(entry.time_pricing.periods).toHaveLength(1)
|
||||
})
|
||||
})
|
||||
|
||||
describe('PricingEntryCard service tier multipliers', () => {
|
||||
it('shows Fast and Flex controls only when explicitly enabled', () => {
|
||||
const hidden = shallowMount(PricingEntryCard, { props: { entry: createEntry() } })
|
||||
expect(hidden.text()).not.toContain('admin.channels.form.fastMultiplier')
|
||||
|
||||
const shown = shallowMount(PricingEntryCard, {
|
||||
props: { entry: createEntry(), enableTierMultipliers: true },
|
||||
})
|
||||
expect(shown.text()).toContain('admin.channels.form.fastMultiplier')
|
||||
expect(shown.text()).toContain('admin.channels.form.flexMultiplier')
|
||||
})
|
||||
})
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
import { describe, expect, it } from 'vitest'
|
||||
import {
|
||||
apiIntervalsToForm,
|
||||
apiTimePricingToForm,
|
||||
createDefaultTimePricingForm,
|
||||
formIntervalsToAPI,
|
||||
formTimePricingToAPI,
|
||||
isValidPositiveMultiplier,
|
||||
validateIntervals,
|
||||
validateTimePricing,
|
||||
type IntervalFormEntry,
|
||||
@@ -10,6 +13,51 @@ import {
|
||||
type TimePricingPeriodFormEntry,
|
||||
} from '../types'
|
||||
|
||||
describe('interval multiplier conversion', () => {
|
||||
it('preserves component multipliers without MTok conversion', () => {
|
||||
const form = apiIntervalsToForm([{
|
||||
min_tokens: 272000,
|
||||
max_tokens: null,
|
||||
tier_label: '',
|
||||
input_price: null,
|
||||
output_price: null,
|
||||
cache_write_price: null,
|
||||
cache_read_price: null,
|
||||
input_multiplier: 2,
|
||||
output_multiplier: 1.5,
|
||||
cache_write_multiplier: 2,
|
||||
cache_read_multiplier: 2,
|
||||
per_request_price: null,
|
||||
sort_order: 0,
|
||||
}])
|
||||
|
||||
expect(form[0].input_multiplier).toBe(2)
|
||||
expect(form[0].output_multiplier).toBe(1.5)
|
||||
expect(formIntervalsToAPI(form)[0]).toMatchObject({
|
||||
input_multiplier: 2,
|
||||
output_multiplier: 1.5,
|
||||
cache_write_multiplier: 2,
|
||||
cache_read_multiplier: 2,
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
describe('positive multiplier validation', () => {
|
||||
it('accepts empty and positive values but rejects zero and negative values', () => {
|
||||
expect(isValidPositiveMultiplier(null)).toBe(true)
|
||||
expect(isValidPositiveMultiplier('')).toBe(true)
|
||||
expect(isValidPositiveMultiplier('0.5')).toBe(true)
|
||||
expect(isValidPositiveMultiplier(0)).toBe(false)
|
||||
expect(isValidPositiveMultiplier(-1)).toBe(false)
|
||||
})
|
||||
|
||||
it('rejects a zero interval multiplier', () => {
|
||||
expect(validateIntervals([
|
||||
makeInterval({ min_tokens: 100, input_multiplier: 0 }),
|
||||
], 'token', t)).toContain('multiplierPositive')
|
||||
})
|
||||
})
|
||||
|
||||
function makeInterval(over: Partial<IntervalFormEntry>): IntervalFormEntry {
|
||||
return {
|
||||
min_tokens: 0,
|
||||
@@ -19,6 +67,10 @@ function makeInterval(over: Partial<IntervalFormEntry>): IntervalFormEntry {
|
||||
output_price: null,
|
||||
cache_write_price: null,
|
||||
cache_read_price: null,
|
||||
input_multiplier: null,
|
||||
output_multiplier: null,
|
||||
cache_write_multiplier: null,
|
||||
cache_read_multiplier: null,
|
||||
per_request_price: null,
|
||||
sort_order: 0,
|
||||
...over,
|
||||
|
||||
@@ -10,6 +10,10 @@ export interface IntervalFormEntry {
|
||||
output_price: number | string | null
|
||||
cache_write_price: number | string | null
|
||||
cache_read_price: number | string | null
|
||||
input_multiplier: number | string | null
|
||||
output_multiplier: number | string | null
|
||||
cache_write_multiplier: number | string | null
|
||||
cache_read_multiplier: number | string | null
|
||||
per_request_price: number | string | null
|
||||
sort_order: number
|
||||
}
|
||||
@@ -21,6 +25,8 @@ export interface PricingFormEntry {
|
||||
output_price: number | string | null
|
||||
cache_write_price: number | string | null
|
||||
cache_read_price: number | string | null
|
||||
fast_multiplier?: number | string | null
|
||||
flex_multiplier?: number | string | null
|
||||
image_input_price: number | string | null
|
||||
image_output_price: number | string | null
|
||||
per_request_price: number | string | null
|
||||
@@ -164,6 +170,12 @@ export function toNullableNumber(val: number | string | null | undefined): numbe
|
||||
return isNaN(num) ? null : num
|
||||
}
|
||||
|
||||
export function isValidPositiveMultiplier(val: number | string | null | undefined): boolean {
|
||||
if (val === null || val === undefined || val === '') return true
|
||||
const multiplier = Number(val)
|
||||
return Number.isFinite(multiplier) && multiplier > 0
|
||||
}
|
||||
|
||||
/** 前端显示值($/MTok) → 后端存储值(per-token) */
|
||||
export function mTokToPerToken(val: number | string | null | undefined): number | null {
|
||||
const num = toNullableNumber(val)
|
||||
@@ -186,6 +198,10 @@ export function apiIntervalsToForm(intervals: PricingInterval[]): IntervalFormEn
|
||||
output_price: perTokenToMTok(iv.output_price),
|
||||
cache_write_price: perTokenToMTok(iv.cache_write_price),
|
||||
cache_read_price: perTokenToMTok(iv.cache_read_price),
|
||||
input_multiplier: iv.input_multiplier,
|
||||
output_multiplier: iv.output_multiplier,
|
||||
cache_write_multiplier: iv.cache_write_multiplier,
|
||||
cache_read_multiplier: iv.cache_read_multiplier,
|
||||
per_request_price: iv.per_request_price,
|
||||
sort_order: iv.sort_order
|
||||
}))
|
||||
@@ -200,6 +216,10 @@ export function formIntervalsToAPI(intervals: IntervalFormEntry[]): PricingInter
|
||||
output_price: mTokToPerToken(iv.output_price),
|
||||
cache_write_price: mTokToPerToken(iv.cache_write_price),
|
||||
cache_read_price: mTokToPerToken(iv.cache_read_price),
|
||||
input_multiplier: toNullableNumber(iv.input_multiplier),
|
||||
output_multiplier: toNullableNumber(iv.output_multiplier),
|
||||
cache_write_multiplier: toNullableNumber(iv.cache_write_multiplier),
|
||||
cache_read_multiplier: toNullableNumber(iv.cache_read_multiplier),
|
||||
per_request_price: toNullableNumber(iv.per_request_price),
|
||||
sort_order: iv.sort_order
|
||||
}))
|
||||
@@ -332,6 +352,20 @@ function validateIntervalPrices(iv: IntervalFormEntry, idx: number, t: Translate
|
||||
)
|
||||
}
|
||||
}
|
||||
const multipliers: [string, number | string | null][] = [
|
||||
['inputMultiplier', iv.input_multiplier],
|
||||
['outputMultiplier', iv.output_multiplier],
|
||||
['cacheWriteMultiplier', iv.cache_write_multiplier],
|
||||
['cacheReadMultiplier', iv.cache_read_multiplier],
|
||||
]
|
||||
for (const [key, val] of multipliers) {
|
||||
if (!isValidPositiveMultiplier(val)) {
|
||||
return intervalValidationMessage(t, 'multiplierPositive', {
|
||||
index,
|
||||
field: intervalPriceLabel(t, key),
|
||||
})
|
||||
}
|
||||
}
|
||||
return null
|
||||
}
|
||||
|
||||
|
||||
@@ -70,6 +70,7 @@ export default {
|
||||
maxPositive: 'Interval #{index}: maximum token count ({value}) must be greater than 0',
|
||||
maxGreaterThanMin: 'Interval #{index}: maximum token count ({max}) must be greater than minimum token count ({min})',
|
||||
negativePrice: 'Interval #{index}: {field} cannot be negative',
|
||||
multiplierPositive: 'Interval #{index}: {field} must be greater than 0',
|
||||
unboundedLast: 'Interval #{index}: an unbounded interval (empty maximum token count) must be last',
|
||||
overlap: 'Intervals #{previousIndex} and #{currentIndex} overlap: previous upper bound ({previousMax}) is greater than current lower bound ({currentMin})',
|
||||
price: {
|
||||
@@ -128,6 +129,14 @@ export default {
|
||||
imageTokenPrice: 'Image Output',
|
||||
imageOutputPrice: 'Image Output Price',
|
||||
pricePlaceholder: 'Default',
|
||||
fastMultiplier: 'Fast Multiplier',
|
||||
flexMultiplier: 'Flex Multiplier',
|
||||
multiplierPlaceholder: 'Not set',
|
||||
multiplierPositive: 'Fast/Flex multipliers must be greater than 0',
|
||||
inputMultiplier: 'Input Mult.',
|
||||
outputMultiplier: 'Output Mult.',
|
||||
cacheWriteMultiplier: 'Cache Write Mult.',
|
||||
cacheReadMultiplier: 'Cache Read Mult.',
|
||||
intervals: 'Context Intervals (optional)',
|
||||
timePricing: 'Time-based pricing (optional)',
|
||||
timezone: 'Time zone',
|
||||
|
||||
@@ -1011,7 +1011,7 @@ export default {
|
||||
title: 'Per-model group pricing',
|
||||
description: 'Overrides channel and built-in prices for matching models. Long-context tiers come from official presets — do not enter custom intervals. Use per-request tiers such as realtime, tts, and stt for audio.',
|
||||
longContext: 'Enable long-context tier pricing',
|
||||
longContextHint: 'When checked, official/preset long-context tiers apply. When unchecked, token models stay on the first-tier base rate.',
|
||||
longContextHint: 'When checked, channel intervals or official preset tiers apply. Otherwise the first tier is used unless the account explicitly enables long-context billing.',
|
||||
add: 'Add model price'
|
||||
},
|
||||
voicePricing: {
|
||||
|
||||
@@ -70,6 +70,7 @@ export default {
|
||||
maxPositive: '区间 #{index}:最大 token 数({value})必须大于 0',
|
||||
maxGreaterThanMin: '区间 #{index}:最大 token 数({max})必须大于最小 token 数({min})',
|
||||
negativePrice: '区间 #{index}:{field}不能为负数',
|
||||
multiplierPositive: '区间 #{index}:{field}必须大于 0',
|
||||
unboundedLast: '区间 #{index}:无上限区间(最大 token 数为空)必须放在最后',
|
||||
overlap: '区间 #{previousIndex} 和 #{currentIndex} 重叠:前一个上界({previousMax})大于当前下界({currentMin})',
|
||||
price: {
|
||||
@@ -128,6 +129,14 @@ export default {
|
||||
imageTokenPrice: '图片输出',
|
||||
imageOutputPrice: '图片输出价格',
|
||||
pricePlaceholder: '默认',
|
||||
fastMultiplier: 'Fast 倍率',
|
||||
flexMultiplier: 'Flex 倍率',
|
||||
multiplierPlaceholder: '未配置',
|
||||
multiplierPositive: 'Fast/Flex 倍率必须大于 0',
|
||||
inputMultiplier: '输入倍率',
|
||||
outputMultiplier: '输出倍率',
|
||||
cacheWriteMultiplier: '缓存写倍率',
|
||||
cacheReadMultiplier: '缓存读倍率',
|
||||
intervals: '上下文区间定价(可选)',
|
||||
timePricing: '时间段定价(可选)',
|
||||
timezone: '时区',
|
||||
|
||||
@@ -1008,7 +1008,7 @@ export default {
|
||||
title: '分组逐模型定价',
|
||||
description: '匹配模型后覆盖渠道和内置价格。长上下文阶梯沿用官方/预设价卡,无需再手填区间。音频可用按次层级配置 realtime、tts、stt。',
|
||||
longContext: '启用长上下文阶梯定价',
|
||||
longContextHint: '勾选后按官方/预设阶梯计费;关闭则始终按第一档基础价。',
|
||||
longContextHint: '勾选后按渠道区间或官方预设阶梯计费;关闭后默认按第一档,账号显式开启时除外。',
|
||||
add: '添加模型价格'
|
||||
},
|
||||
voicePricing: {
|
||||
|
||||
@@ -448,6 +448,7 @@
|
||||
:entry="entry"
|
||||
:platform="section.platform"
|
||||
enable-time-pricing
|
||||
enable-tier-multipliers
|
||||
@update="updatePricingEntry(sIdx, idx, $event)"
|
||||
@remove="removePricingEntry(sIdx, idx)"
|
||||
/>
|
||||
@@ -633,7 +634,7 @@ import { extractApiErrorMessage } from '@/utils/apiError'
|
||||
import { adminAPI } from '@/api/admin'
|
||||
import type { Channel, ChannelModelPricing, CreateChannelRequest, UpdateChannelRequest, AccountStatsPricingRule } from '@/api/admin/channels'
|
||||
import type { PricingFormEntry } from '@/components/admin/channel/types'
|
||||
import { apiIntervalsToForm, apiTimePricingToForm, createDefaultTimePricingForm, findModelConflict, formIntervalsToAPI, formTimePricingToAPI, mTokToPerToken, perTokenToMTok, validateIntervals, validateTimePricing } from '@/components/admin/channel/types'
|
||||
import { apiIntervalsToForm, apiTimePricingToForm, createDefaultTimePricingForm, findModelConflict, formIntervalsToAPI, formTimePricingToAPI, isValidPositiveMultiplier, mTokToPerToken, perTokenToMTok, validateIntervals, validateTimePricing } from '@/components/admin/channel/types'
|
||||
import type { AdminGroup, GroupPlatform } from '@/types'
|
||||
import type { Column } from '@/components/common/types'
|
||||
import { platformTextClass, platformBadgeLightClass } from '@/utils/platformColors'
|
||||
@@ -861,6 +862,8 @@ function addPricingEntry(sectionIdx: number) {
|
||||
output_price: null,
|
||||
cache_write_price: null,
|
||||
cache_read_price: null,
|
||||
fast_multiplier: null,
|
||||
flex_multiplier: null,
|
||||
image_input_price: null,
|
||||
image_output_price: null,
|
||||
per_request_price: null,
|
||||
@@ -895,6 +898,8 @@ async function syncLatestModels(sectionIdx: number) {
|
||||
output_price: null,
|
||||
cache_write_price: null,
|
||||
cache_read_price: null,
|
||||
fast_multiplier: null,
|
||||
flex_multiplier: null,
|
||||
image_input_price: null,
|
||||
image_output_price: null,
|
||||
per_request_price: null,
|
||||
@@ -1120,6 +1125,8 @@ function formToAPI(): { group_ids: number[], model_pricing: ChannelModelPricing[
|
||||
output_price: mTokToPerToken(entry.output_price),
|
||||
cache_write_price: mTokToPerToken(entry.cache_write_price),
|
||||
cache_read_price: mTokToPerToken(entry.cache_read_price),
|
||||
fast_multiplier: entry.fast_multiplier != null && entry.fast_multiplier !== '' ? Number(entry.fast_multiplier) : null,
|
||||
flex_multiplier: entry.flex_multiplier != null && entry.flex_multiplier !== '' ? Number(entry.flex_multiplier) : null,
|
||||
image_input_price: mTokToPerToken(entry.image_input_price),
|
||||
image_output_price: mTokToPerToken(entry.image_output_price),
|
||||
per_request_price: entry.per_request_price != null && entry.per_request_price !== '' ? Number(entry.per_request_price) : null,
|
||||
@@ -1219,6 +1226,8 @@ function apiToForm(channel: Channel): PlatformSection[] {
|
||||
output_price: perTokenToMTok(p.output_price),
|
||||
cache_write_price: perTokenToMTok(p.cache_write_price),
|
||||
cache_read_price: perTokenToMTok(p.cache_read_price),
|
||||
fast_multiplier: p.fast_multiplier,
|
||||
flex_multiplier: p.flex_multiplier,
|
||||
image_input_price: perTokenToMTok(p.image_input_price),
|
||||
image_output_price: perTokenToMTok(p.image_output_price),
|
||||
per_request_price: p.per_request_price,
|
||||
@@ -1524,6 +1533,14 @@ async function handleSubmit() {
|
||||
// 校验区间合法性(范围、重叠等)
|
||||
for (const section of form.platforms.filter(s => s.enabled)) {
|
||||
for (const entry of section.model_pricing) {
|
||||
if (!isValidPositiveMultiplier(entry.fast_multiplier) ||
|
||||
!isValidPositiveMultiplier(entry.flex_multiplier)) {
|
||||
const platformLabel = t('admin.groups.platforms.' + section.platform, section.platform)
|
||||
const modelLabel = entry.models.join(', ') || t('admin.channels.form.unnamed')
|
||||
appStore.showError(`${platformLabel} - ${modelLabel}: ${t('admin.channels.form.multiplierPositive')}`)
|
||||
activeTab.value = section.platform
|
||||
return
|
||||
}
|
||||
if (!entry.intervals || entry.intervals.length === 0) continue
|
||||
const intervalErr = validateIntervals(entry.intervals, entry.billing_mode, t)
|
||||
if (intervalErr) {
|
||||
|
||||
Reference in New Issue
Block a user