diff --git a/backend/internal/handler/admin/channel_handler.go b/backend/internal/handler/admin/channel_handler.go index dfce024fa3..dc0e5b07e2 100644 --- a/backend/internal/handler/admin/channel_handler.go +++ b/backend/internal/handler/admin/channel_handler.go @@ -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 diff --git a/backend/internal/handler/admin/channel_handler_test.go b/backend/internal/handler/admin/channel_handler_test.go index 6c5ddedc92..75aec7db31 100644 --- a/backend/internal/handler/admin/channel_handler_test.go +++ b/backend/internal/handler/admin/channel_handler_test.go @@ -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, diff --git a/backend/internal/repository/channel_repo_pricing.go b/backend/internal/repository/channel_repo_pricing.go index b7a9fd6b02..e25642085a 100644 --- a/backend/internal/repository/channel_repo_pricing.go +++ b/backend/internal/repository/channel_repo_pricing.go @@ -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) } diff --git a/backend/internal/repository/channel_repo_pricing_time_test.go b/backend/internal/repository/channel_repo_pricing_time_test.go index d7af1886e7..38acdbdc73 100644 --- a/backend/internal/repository/channel_repo_pricing_time_test.go +++ b/backend/internal/repository/channel_repo_pricing_time_test.go @@ -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)) diff --git a/backend/internal/service/account_stats_pricing.go b/backend/internal/service/account_stats_pricing.go index df9e8e05aa..8d5bc144fc 100644 --- a/backend/internal/service/account_stats_pricing.go +++ b/backend/internal/service/account_stats_pricing.go @@ -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 { diff --git a/backend/internal/service/account_stats_pricing_test.go b/backend/internal/service/account_stats_pricing_test.go index 48336a5834..1bd28896fc 100644 --- a/backend/internal/service/account_stats_pricing_test.go +++ b/backend/internal/service/account_stats_pricing_test.go @@ -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, diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index 6ead8da021..af8c694495 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -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) diff --git a/backend/internal/service/billing_service_test.go b/backend/internal/service/billing_service_test.go index 80d07af44b..1e801d8e95 100644 --- a/backend/internal/service/billing_service_test.go +++ b/backend/internal/service/billing_service_test.go @@ -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) { diff --git a/backend/internal/service/channel.go b/backend/internal/service/channel.go index 5345693eb0..0fe4e3b07c 100644 --- a/backend/internal/service/channel.go +++ b/backend/internal/service/channel.go @@ -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 } diff --git a/backend/internal/service/channel_pricing_multipliers_test.go b/backend/internal/service/channel_pricing_multipliers_test.go new file mode 100644 index 0000000000..5b38fddec2 --- /dev/null +++ b/backend/internal/service/channel_pricing_multipliers_test.go @@ -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) +} diff --git a/backend/internal/service/channel_service.go b/backend/internal/service/channel_service.go index ee732dbeef..b3fe7dc63a 100644 --- a/backend/internal/service/channel_service.go +++ b/backend/internal/service/channel_service.go @@ -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", diff --git a/backend/internal/service/gateway_forward.go b/backend/internal/service/gateway_forward.go index 863e4d2e83..d8a105c7d8 100644 --- a/backend/internal/service/gateway_forward.go +++ b/backend/internal/service/gateway_forward.go @@ -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 { diff --git a/backend/internal/service/gateway_request.go b/backend/internal/service/gateway_request.go index 4ffd5ee8e9..111b807f2f 100644 --- a/backend/internal/service/gateway_request.go +++ b/backend/internal/service/gateway_request.go @@ -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 时行为不变) diff --git a/backend/internal/service/gateway_request_test.go b/backend/internal/service/gateway_request_test.go index d4e010a5e8..6074c02991 100644 --- a/backend/internal/service/gateway_request_test.go +++ b/backend/internal/service/gateway_request_test.go @@ -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), "") diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 2578d59b1d..20e7e5b9e6 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -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 // 生成的图片数量 diff --git a/backend/internal/service/gateway_usage_billing.go b/backend/internal/service/gateway_usage_billing.go index c9671328be..e1fac3d6cf 100644 --- a/backend/internal/service/gateway_usage_billing.go +++ b/backend/internal/service/gateway_usage_billing.go @@ -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), diff --git a/backend/internal/service/model_pricing_resolver.go b/backend/internal/service/model_pricing_resolver.go index 52c6b930c5..62283e81b7 100644 --- a/backend/internal/service/model_pricing_resolver.go +++ b/backend/internal/service/model_pricing_resolver.go @@ -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 一致)。 diff --git a/backend/internal/service/model_pricing_resolver_test.go b/backend/internal/service/model_pricing_resolver_test.go index 11613fa9eb..3fb3e1c63f 100644 --- a/backend/internal/service/model_pricing_resolver_test.go +++ b/backend/internal/service/model_pricing_resolver_test.go @@ -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{ diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go index 21e5cd4096..9cca810400 100644 --- a/backend/internal/service/openai_gateway_record_usage_test.go +++ b/backend/internal/service/openai_gateway_record_usage_test.go @@ -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, diff --git a/backend/internal/service/usage_log.go b/backend/internal/service/usage_log.go index 7a41555ece..b6dc2a9d0e 100644 --- a/backend/internal/service/usage_log.go +++ b/backend/internal/service/usage_log.go @@ -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". diff --git a/backend/internal/service/usage_log_helpers.go b/backend/internal/service/usage_log_helpers.go index b431b50aef..deb102e736 100644 --- a/backend/internal/service/usage_log_helpers.go +++ b/backend/internal/service/usage_log_helpers.go @@ -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 diff --git a/backend/migrations/228_channel_pricing_multipliers.sql b/backend/migrations/228_channel_pricing_multipliers.sql new file mode 100644 index 0000000000..33322ba1a3 --- /dev/null +++ b/backend/migrations/228_channel_pricing_multipliers.sql @@ -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'; diff --git a/backend/migrations/channel_pricing_multipliers_migration_test.go b/backend/migrations/channel_pricing_multipliers_migration_test.go new file mode 100644 index 0000000000..96ff90f53b --- /dev/null +++ b/backend/migrations/channel_pricing_multipliers_migration_test.go @@ -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)") + } +} diff --git a/frontend/src/api/admin/channels.ts b/frontend/src/api/admin/channels.ts index 6556417ced..9ce50bcd0a 100644 --- a/frontend/src/api/admin/channels.ts +++ b/frontend/src/api/admin/channels.ts @@ -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 diff --git a/frontend/src/components/account/CreateAccountModal.vue b/frontend/src/components/account/CreateAccountModal.vue index 55a0831ae1..47ee8c27d5 100644 --- a/frontend/src/components/account/CreateAccountModal.vue +++ b/frontend/src/components/account/CreateAccountModal.vue @@ -2958,6 +2958,7 @@