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/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_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/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)") + } +}