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:
Wesley Liddick
2026-08-20 10:13:17 +08:00
committed by GitHub
39 changed files with 1303 additions and 280 deletions
@@ -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,
+144 -53
View File
@@ -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) {
+34 -14
View File
@@ -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)
}
+14 -1
View File
@@ -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",
+55 -1
View File
@@ -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,
+2 -1
View File
@@ -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)")
}
}
+6
View File
@@ -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: {
+18 -1
View File
@@ -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) {