渠道定价:持久化服务层级与区间倍率

This commit is contained in:
IanShaw
2026-08-19 06:35:11 -07:00
parent 32a0d9ba2d
commit fce90ecf89
8 changed files with 294 additions and 84 deletions
@@ -64,6 +64,8 @@ type channelModelPricingRequest struct {
OutputPrice *float64 `json:"output_price" binding:"omitempty,min=0"`
CacheWritePrice *float64 `json:"cache_write_price" binding:"omitempty,min=0"`
CacheReadPrice *float64 `json:"cache_read_price" binding:"omitempty,min=0"`
FastMultiplier *float64 `json:"fast_multiplier" binding:"omitempty,gt=0"`
FlexMultiplier *float64 `json:"flex_multiplier" binding:"omitempty,gt=0"`
ImageInputPrice *float64 `json:"image_input_price" binding:"omitempty,min=0"`
ImageOutputPrice *float64 `json:"image_output_price" binding:"omitempty,min=0"`
PerRequestPrice *float64 `json:"per_request_price" binding:"omitempty,min=0"`
@@ -83,15 +85,19 @@ type channelTimePricingPeriodRequest struct {
}
type pricingIntervalRequest struct {
MinTokens int `json:"min_tokens"`
MaxTokens *int `json:"max_tokens"`
TierLabel string `json:"tier_label"`
InputPrice *float64 `json:"input_price"`
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheReadPrice *float64 `json:"cache_read_price"`
PerRequestPrice *float64 `json:"per_request_price"`
SortOrder int `json:"sort_order"`
MinTokens int `json:"min_tokens"`
MaxTokens *int `json:"max_tokens"`
TierLabel string `json:"tier_label"`
InputPrice *float64 `json:"input_price"`
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheReadPrice *float64 `json:"cache_read_price"`
InputMultiplier *float64 `json:"input_multiplier" binding:"omitempty,gt=0"`
OutputMultiplier *float64 `json:"output_multiplier" binding:"omitempty,gt=0"`
CacheWriteMultiplier *float64 `json:"cache_write_multiplier" binding:"omitempty,gt=0"`
CacheReadMultiplier *float64 `json:"cache_read_multiplier" binding:"omitempty,gt=0"`
PerRequestPrice *float64 `json:"per_request_price"`
SortOrder int `json:"sort_order"`
}
type accountStatsPricingRuleRequest struct {
@@ -128,6 +134,8 @@ type channelModelPricingResponse struct {
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheReadPrice *float64 `json:"cache_read_price"`
FastMultiplier *float64 `json:"fast_multiplier"`
FlexMultiplier *float64 `json:"flex_multiplier"`
ImageInputPrice *float64 `json:"image_input_price"`
ImageOutputPrice *float64 `json:"image_output_price"`
PerRequestPrice *float64 `json:"per_request_price"`
@@ -147,16 +155,20 @@ type channelTimePricingPeriodResponse struct {
}
type pricingIntervalResponse struct {
ID int64 `json:"id"`
MinTokens int `json:"min_tokens"`
MaxTokens *int `json:"max_tokens"`
TierLabel string `json:"tier_label,omitempty"`
InputPrice *float64 `json:"input_price"`
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheReadPrice *float64 `json:"cache_read_price"`
PerRequestPrice *float64 `json:"per_request_price"`
SortOrder int `json:"sort_order"`
ID int64 `json:"id"`
MinTokens int `json:"min_tokens"`
MaxTokens *int `json:"max_tokens"`
TierLabel string `json:"tier_label,omitempty"`
InputPrice *float64 `json:"input_price"`
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheReadPrice *float64 `json:"cache_read_price"`
InputMultiplier *float64 `json:"input_multiplier"`
OutputMultiplier *float64 `json:"output_multiplier"`
CacheWriteMultiplier *float64 `json:"cache_write_multiplier"`
CacheReadMultiplier *float64 `json:"cache_read_multiplier"`
PerRequestPrice *float64 `json:"per_request_price"`
SortOrder int `json:"sort_order"`
}
type accountStatsPricingRuleResponse struct {
@@ -248,6 +260,8 @@ func pricingToResponse(p *service.ChannelModelPricing) channelModelPricingRespon
OutputPrice: p.OutputPrice,
CacheWritePrice: p.CacheWritePrice,
CacheReadPrice: p.CacheReadPrice,
FastMultiplier: p.FastMultiplier,
FlexMultiplier: p.FlexMultiplier,
ImageInputPrice: p.ImageInputPrice,
ImageOutputPrice: p.ImageOutputPrice,
PerRequestPrice: p.PerRequestPrice,
@@ -273,20 +287,24 @@ func timePricingToResponse(value *service.ChannelTimePricing) *channelTimePricin
func intervalToResponse(iv service.PricingInterval) pricingIntervalResponse {
return pricingIntervalResponse{
ID: iv.ID,
MinTokens: iv.MinTokens,
MaxTokens: iv.MaxTokens,
TierLabel: iv.TierLabel,
InputPrice: iv.InputPrice,
OutputPrice: iv.OutputPrice,
CacheWritePrice: iv.CacheWritePrice,
CacheReadPrice: iv.CacheReadPrice,
PerRequestPrice: iv.PerRequestPrice,
SortOrder: iv.SortOrder,
ID: iv.ID,
MinTokens: iv.MinTokens,
MaxTokens: iv.MaxTokens,
TierLabel: iv.TierLabel,
InputPrice: iv.InputPrice,
OutputPrice: iv.OutputPrice,
CacheWritePrice: iv.CacheWritePrice,
CacheReadPrice: iv.CacheReadPrice,
InputMultiplier: iv.InputMultiplier,
OutputMultiplier: iv.OutputMultiplier,
CacheWriteMultiplier: iv.CacheWriteMultiplier,
CacheReadMultiplier: iv.CacheReadMultiplier,
PerRequestPrice: iv.PerRequestPrice,
SortOrder: iv.SortOrder,
}
}
func pricingRequestToService(reqs []channelModelPricingRequest) []service.ChannelModelPricing {
func pricingRequestToService(reqs []channelModelPricingRequest, allowChannelMultipliers bool) []service.ChannelModelPricing {
result := make([]service.ChannelModelPricing, 0, len(reqs))
for _, r := range reqs {
billingMode := service.BillingMode(r.BillingMode)
@@ -296,18 +314,34 @@ func pricingRequestToService(reqs []channelModelPricingRequest) []service.Channe
platform := r.Platform
intervals := make([]service.PricingInterval, 0, len(r.Intervals))
for _, iv := range r.Intervals {
var inputMultiplier, outputMultiplier, cacheWriteMultiplier, cacheReadMultiplier *float64
if allowChannelMultipliers {
inputMultiplier = iv.InputMultiplier
outputMultiplier = iv.OutputMultiplier
cacheWriteMultiplier = iv.CacheWriteMultiplier
cacheReadMultiplier = iv.CacheReadMultiplier
}
intervals = append(intervals, service.PricingInterval{
MinTokens: iv.MinTokens,
MaxTokens: iv.MaxTokens,
TierLabel: iv.TierLabel,
InputPrice: iv.InputPrice,
OutputPrice: iv.OutputPrice,
CacheWritePrice: iv.CacheWritePrice,
CacheReadPrice: iv.CacheReadPrice,
PerRequestPrice: iv.PerRequestPrice,
SortOrder: iv.SortOrder,
MinTokens: iv.MinTokens,
MaxTokens: iv.MaxTokens,
TierLabel: iv.TierLabel,
InputPrice: iv.InputPrice,
OutputPrice: iv.OutputPrice,
CacheWritePrice: iv.CacheWritePrice,
CacheReadPrice: iv.CacheReadPrice,
InputMultiplier: inputMultiplier,
OutputMultiplier: outputMultiplier,
CacheWriteMultiplier: cacheWriteMultiplier,
CacheReadMultiplier: cacheReadMultiplier,
PerRequestPrice: iv.PerRequestPrice,
SortOrder: iv.SortOrder,
})
}
var fastMultiplier, flexMultiplier *float64
if allowChannelMultipliers {
fastMultiplier = r.FastMultiplier
flexMultiplier = r.FlexMultiplier
}
result = append(result, service.ChannelModelPricing{
Platform: platform,
Models: r.Models,
@@ -316,6 +350,8 @@ func pricingRequestToService(reqs []channelModelPricingRequest) []service.Channe
OutputPrice: r.OutputPrice,
CacheWritePrice: r.CacheWritePrice,
CacheReadPrice: r.CacheReadPrice,
FastMultiplier: fastMultiplier,
FlexMultiplier: flexMultiplier,
ImageInputPrice: r.ImageInputPrice,
ImageOutputPrice: r.ImageOutputPrice,
PerRequestPrice: r.PerRequestPrice,
@@ -346,7 +382,7 @@ func accountStatsPricingRuleRequestToService(r accountStatsPricingRuleRequest) s
Name: r.Name,
GroupIDs: r.GroupIDs,
AccountIDs: r.AccountIDs,
Pricing: pricingRequestToService(r.Pricing),
Pricing: pricingRequestToService(r.Pricing, false),
}
}
@@ -407,7 +443,7 @@ func (h *ChannelHandler) Create(c *gin.Context) {
return
}
pricing := pricingRequestToService(req.ModelPricing)
pricing := pricingRequestToService(req.ModelPricing, true)
// Main model_pricing requires a platform; default to anthropic for backward compatibility.
for i := range pricing {
if pricing[i].Platform == "" {
@@ -481,7 +517,7 @@ func (h *ChannelHandler) Update(c *gin.Context) {
ApplyPricingToAccountStats: req.ApplyPricingToAccountStats,
}
if req.ModelPricing != nil {
pricing := pricingRequestToService(*req.ModelPricing)
pricing := pricingRequestToService(*req.ModelPricing, true)
for i := range pricing {
if pricing[i].Platform == "" {
pricing[i].Platform = service.PlatformAnthropic
@@ -305,7 +305,7 @@ func TestPricingRequestToService_Defaults(t *testing.T) {
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
result := pricingRequestToService([]channelModelPricingRequest{tt.req})
result := pricingRequestToService([]channelModelPricingRequest{tt.req}, true)
require.Len(t, result, 1)
switch tt.wantField {
case "BillingMode":
@@ -332,7 +332,7 @@ func TestPricingRequestToService_WithAllFields(t *testing.T) {
},
}
result := pricingRequestToService(reqs)
result := pricingRequestToService(reqs, true)
require.Len(t, result, 1)
r := result[0]
require.Equal(t, "openai", r.Platform)
@@ -373,7 +373,7 @@ func TestPricingRequestToService_WithIntervals(t *testing.T) {
},
}
result := pricingRequestToService(reqs)
result := pricingRequestToService(reqs, true)
require.Len(t, result, 1)
require.Len(t, result[0].Intervals, 2)
@@ -396,7 +396,7 @@ func TestPricingRequestToService_WithIntervals(t *testing.T) {
}
func TestPricingRequestToService_EmptySlice(t *testing.T) {
result := pricingRequestToService([]channelModelPricingRequest{})
result := pricingRequestToService([]channelModelPricingRequest{}, true)
require.NotNil(t, result)
require.Empty(t, result)
}
@@ -410,7 +410,7 @@ func TestPricingRequestToService_NilPriceFields(t *testing.T) {
},
}
result := pricingRequestToService(reqs)
result := pricingRequestToService(reqs, true)
require.Len(t, result, 1)
r := result[0]
require.Nil(t, r.InputPrice)
@@ -433,16 +433,52 @@ func TestPricingRequestToService_TimePricing(t *testing.T) {
},
}
got := pricingRequestToService([]channelModelPricingRequest{req})
got := pricingRequestToService([]channelModelPricingRequest{req}, true)
require.Equal(t, "Asia/Shanghai", got[0].TimePricing.Timezone)
require.Equal(t, 2.0, got[0].TimePricing.Periods[0].Multiplier)
}
func TestPricingRequestToService_TimePricingNil(t *testing.T) {
got := pricingRequestToService([]channelModelPricingRequest{{Models: []string{"gpt-5"}}})
got := pricingRequestToService([]channelModelPricingRequest{{Models: []string{"gpt-5"}}}, true)
require.Nil(t, got[0].TimePricing)
}
// 账号成本统计规则不支持倍率:allowChannelMultipliers=false 时必须丢弃,
// 避免渠道倍率意外污染账号成本口径。
func TestPricingRequestToService_MultipliersGatedByFlag(t *testing.T) {
req := channelModelPricingRequest{
Models: []string{"gpt-5"},
BillingMode: "token",
FastMultiplier: float64Ptr(2.5),
FlexMultiplier: float64Ptr(0.5),
Intervals: []pricingIntervalRequest{{
MinTokens: 272000,
InputMultiplier: float64Ptr(2),
OutputMultiplier: float64Ptr(1.5),
CacheWriteMultiplier: float64Ptr(2),
CacheReadMultiplier: float64Ptr(2),
}},
}
allowed := pricingRequestToService([]channelModelPricingRequest{req}, true)
require.Equal(t, float64Ptr(2.5), allowed[0].FastMultiplier)
require.Equal(t, float64Ptr(0.5), allowed[0].FlexMultiplier)
require.Equal(t, float64Ptr(2), allowed[0].Intervals[0].InputMultiplier)
require.Equal(t, float64Ptr(1.5), allowed[0].Intervals[0].OutputMultiplier)
require.Equal(t, float64Ptr(2), allowed[0].Intervals[0].CacheWriteMultiplier)
require.Equal(t, float64Ptr(2), allowed[0].Intervals[0].CacheReadMultiplier)
dropped := pricingRequestToService([]channelModelPricingRequest{req}, false)
require.Nil(t, dropped[0].FastMultiplier)
require.Nil(t, dropped[0].FlexMultiplier)
require.Nil(t, dropped[0].Intervals[0].InputMultiplier)
require.Nil(t, dropped[0].Intervals[0].OutputMultiplier)
require.Nil(t, dropped[0].Intervals[0].CacheWriteMultiplier)
require.Nil(t, dropped[0].Intervals[0].CacheReadMultiplier)
// 非倍率字段不受开关影响
require.Equal(t, 272000, dropped[0].Intervals[0].MinTokens)
}
func TestPricingToResponse_TimePricing(t *testing.T) {
got := pricingToResponse(&service.ChannelModelPricing{
BillingMode: service.BillingModeToken,
@@ -16,7 +16,7 @@ import (
func (r *channelRepository) ListModelPricing(ctx context.Context, channelID int64) ([]service.ChannelModelPricing, error) {
rows, err := r.db.QueryContext(ctx,
`SELECT id, channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, image_input_price, image_output_price, per_request_price, time_pricing, created_at, updated_at
`SELECT id, channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, fast_multiplier, flex_multiplier, image_input_price, image_output_price, per_request_price, time_pricing, created_at, updated_at
FROM channel_model_pricing WHERE channel_id = $1 ORDER BY id`, channelID,
)
if err != nil {
@@ -61,10 +61,11 @@ func (r *channelRepository) UpdateModelPricing(ctx context.Context, pricing *ser
}
result, err := r.db.ExecContext(ctx,
`UPDATE channel_model_pricing
SET models = $1, billing_mode = $2, input_price = $3, output_price = $4, cache_write_price = $5, cache_read_price = $6, image_input_price = $7, image_output_price = $8, per_request_price = $9, time_pricing = $10, platform = $11, updated_at = NOW()
WHERE id = $12`,
SET models = $1, billing_mode = $2, input_price = $3, output_price = $4, cache_write_price = $5, cache_read_price = $6, fast_multiplier = $7, flex_multiplier = $8, image_input_price = $9, image_output_price = $10, per_request_price = $11, time_pricing = $12, platform = $13, updated_at = NOW()
WHERE id = $14`,
modelsJSON, billingMode, pricing.InputPrice, pricing.OutputPrice, pricing.CacheWritePrice, pricing.CacheReadPrice,
pricing.ImageInputPrice, pricing.ImageOutputPrice, pricing.PerRequestPrice, timePricingJSON, pricing.Platform, pricing.ID,
pricing.FastMultiplier, pricing.FlexMultiplier, pricing.ImageInputPrice, pricing.ImageOutputPrice, pricing.PerRequestPrice,
timePricingJSON, pricing.Platform, pricing.ID,
)
if err != nil {
return fmt.Errorf("update model pricing: %w", err)
@@ -95,7 +96,7 @@ func (r *channelRepository) ReplaceModelPricing(ctx context.Context, channelID i
// batchLoadModelPricing 批量加载多个渠道的模型定价(含区间)
func (r *channelRepository) batchLoadModelPricing(ctx context.Context, channelIDs []int64) (map[int64][]service.ChannelModelPricing, error) {
rows, err := r.db.QueryContext(ctx,
`SELECT id, channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, image_input_price, image_output_price, per_request_price, time_pricing, created_at, updated_at
`SELECT id, channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, fast_multiplier, flex_multiplier, image_input_price, image_output_price, per_request_price, time_pricing, created_at, updated_at
FROM channel_model_pricing WHERE channel_id = ANY($1) ORDER BY channel_id, id`,
pq.Array(channelIDs),
)
@@ -136,6 +137,7 @@ func (r *channelRepository) batchLoadIntervals(ctx context.Context, pricingIDs [
rows, err := r.db.QueryContext(ctx,
`SELECT id, pricing_id, min_tokens, max_tokens, tier_label,
input_price, output_price, cache_write_price, cache_read_price,
input_multiplier, output_multiplier, cache_write_multiplier, cache_read_multiplier,
per_request_price, sort_order, created_at, updated_at
FROM channel_pricing_intervals
WHERE pricing_id = ANY($1) ORDER BY pricing_id, sort_order, id`,
@@ -152,6 +154,7 @@ func (r *channelRepository) batchLoadIntervals(ctx context.Context, pricingIDs [
if err := rows.Scan(
&iv.ID, &iv.PricingID, &iv.MinTokens, &iv.MaxTokens, &iv.TierLabel,
&iv.InputPrice, &iv.OutputPrice, &iv.CacheWritePrice, &iv.CacheReadPrice,
&iv.InputMultiplier, &iv.OutputMultiplier, &iv.CacheWriteMultiplier, &iv.CacheReadMultiplier,
&iv.PerRequestPrice, &iv.SortOrder, &iv.CreatedAt, &iv.UpdatedAt,
); err != nil {
return nil, fmt.Errorf("scan interval: %w", err)
@@ -177,6 +180,7 @@ func scanModelPricingRows(rows *sql.Rows) ([]service.ChannelModelPricing, []int6
if err := rows.Scan(
&p.ID, &p.ChannelID, &p.Platform, &modelsJSON, &p.BillingMode,
&p.InputPrice, &p.OutputPrice, &p.CacheWritePrice, &p.CacheReadPrice,
&p.FastMultiplier, &p.FlexMultiplier,
&p.ImageInputPrice, &p.ImageOutputPrice, &p.PerRequestPrice, &timePricingJSON, &p.CreatedAt, &p.UpdatedAt,
); err != nil {
return nil, nil, fmt.Errorf("scan model pricing: %w", err)
@@ -243,11 +247,12 @@ func createModelPricingExec(ctx context.Context, exec dbExec, pricing *service.C
platform = "anthropic"
}
err = exec.QueryRowContext(ctx,
`INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, image_input_price, image_output_price, per_request_price, time_pricing)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12) RETURNING id, created_at, updated_at`,
`INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, fast_multiplier, flex_multiplier, image_input_price, image_output_price, per_request_price, time_pricing)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14) RETURNING id, created_at, updated_at`,
pricing.ChannelID, platform, modelsJSON, billingMode,
pricing.InputPrice, pricing.OutputPrice, pricing.CacheWritePrice, pricing.CacheReadPrice,
pricing.ImageInputPrice, pricing.ImageOutputPrice, pricing.PerRequestPrice, timePricingJSON,
pricing.FastMultiplier, pricing.FlexMultiplier, pricing.ImageInputPrice, pricing.ImageOutputPrice,
pricing.PerRequestPrice, timePricingJSON,
).Scan(&pricing.ID, &pricing.CreatedAt, &pricing.UpdatedAt)
if err != nil {
return fmt.Errorf("insert model pricing: %w", err)
@@ -288,10 +293,11 @@ func unmarshalChannelTimePricing(data []byte) (*service.ChannelTimePricing, erro
func createIntervalExec(ctx context.Context, exec dbExec, iv *service.PricingInterval) error {
return exec.QueryRowContext(ctx,
`INSERT INTO channel_pricing_intervals
(pricing_id, min_tokens, max_tokens, tier_label, input_price, output_price, cache_write_price, cache_read_price, per_request_price, sort_order)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10) RETURNING id, created_at, updated_at`,
(pricing_id, min_tokens, max_tokens, tier_label, input_price, output_price, cache_write_price, cache_read_price, input_multiplier, output_multiplier, cache_write_multiplier, cache_read_multiplier, per_request_price, sort_order)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14) RETURNING id, created_at, updated_at`,
iv.PricingID, iv.MinTokens, iv.MaxTokens, iv.TierLabel,
iv.InputPrice, iv.OutputPrice, iv.CacheWritePrice, iv.CacheReadPrice,
iv.InputMultiplier, iv.OutputMultiplier, iv.CacheWriteMultiplier, iv.CacheReadMultiplier,
iv.PerRequestPrice, iv.SortOrder,
).Scan(&iv.ID, &iv.CreatedAt, &iv.UpdatedAt)
}
@@ -16,7 +16,7 @@ import (
var channelModelPricingTimePricingColumns = []string{
"id", "channel_id", "platform", "models", "billing_mode", "input_price", "output_price",
"cache_write_price", "cache_read_price", "image_input_price", "image_output_price",
"cache_write_price", "cache_read_price", "fast_multiplier", "flex_multiplier", "image_input_price", "image_output_price",
"per_request_price", "time_pricing", "created_at", "updated_at",
}
@@ -33,7 +33,7 @@ func newChannelModelPricingTimePricingRepo(t *testing.T) (*channelRepository, sq
func modelPricingTimePricingRow(timePricing any) *sqlmock.Rows {
return sqlmock.NewRows(channelModelPricingTimePricingColumns).AddRow(
int64(11), int64(7), "openai", `["gpt-5"]`, service.BillingModeToken,
nil, nil, nil, nil, nil, nil, nil, timePricing,
nil, nil, nil, nil, nil, nil, nil, nil, nil, timePricing,
time.Date(2026, 8, 17, 0, 0, 0, 0, time.UTC), time.Date(2026, 8, 17, 1, 0, 0, 0, time.UTC),
)
}
@@ -105,10 +105,10 @@ func TestChannelModelPricingTimePricingCreateAndUpdateRoundTrip(t *testing.T) {
t.Run("create writes JSON", func(t *testing.T) {
repo, mock := newChannelModelPricingTimePricingRepo(t)
mock.ExpectQuery(regexp.QuoteMeta("INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, image_input_price, image_output_price, per_request_price, time_pricing)")).
mock.ExpectQuery(regexp.QuoteMeta("INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, fast_multiplier, flex_multiplier, image_input_price, image_output_price, per_request_price, time_pricing)")).
WithArgs(
int64(7), "openai", []byte(`["gpt-5"]`), service.BillingModeToken,
nil, nil, nil, nil, nil, nil, nil, channelModelPricingTimePricingJSON,
nil, nil, nil, nil, nil, nil, nil, nil, nil, channelModelPricingTimePricingJSON,
).
WillReturnRows(sqlmock.NewRows([]string{"id", "created_at", "updated_at"}).AddRow(int64(11), time.Time{}, time.Time{}))
@@ -118,10 +118,10 @@ func TestChannelModelPricingTimePricingCreateAndUpdateRoundTrip(t *testing.T) {
t.Run("update writes JSON and entry ID", func(t *testing.T) {
repo, mock := newChannelModelPricingTimePricingRepo(t)
mock.ExpectExec(`(?s)UPDATE channel_model_pricing.*per_request_price = \$9, time_pricing = \$10, platform = \$11.*WHERE id = \$12`).
mock.ExpectExec(`(?s)UPDATE channel_model_pricing.*per_request_price = \$11, time_pricing = \$12, platform = \$13.*WHERE id = \$14`).
WithArgs(
[]byte(`["gpt-5"]`), service.BillingModeToken,
nil, nil, nil, nil, nil, nil, nil, channelModelPricingTimePricingJSON, "openai", int64(11),
nil, nil, nil, nil, nil, nil, nil, nil, nil, channelModelPricingTimePricingJSON, "openai", int64(11),
).
WillReturnResult(sqlmock.NewResult(0, 1))
@@ -153,10 +153,10 @@ func TestChannelModelPricingTimePricingCreateAndUpdateWriteNullWhenDisabled(t *t
t.Run("create writes SQL NULL", func(t *testing.T) {
repo, mock := newChannelModelPricingTimePricingRepo(t)
mock.ExpectQuery(regexp.QuoteMeta("INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, image_input_price, image_output_price, per_request_price, time_pricing)")).
mock.ExpectQuery(regexp.QuoteMeta("INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, fast_multiplier, flex_multiplier, image_input_price, image_output_price, per_request_price, time_pricing)")).
WithArgs(
int64(7), "openai", []byte(`["gpt-5"]`), service.BillingModeToken,
nil, nil, nil, nil, nil, nil, nil, nil,
nil, nil, nil, nil, nil, nil, nil, nil, nil, nil,
).
WillReturnRows(sqlmock.NewRows([]string{"id", "created_at", "updated_at"}).AddRow(int64(11), time.Time{}, time.Time{}))
@@ -166,10 +166,10 @@ func TestChannelModelPricingTimePricingCreateAndUpdateWriteNullWhenDisabled(t *t
t.Run("update writes SQL NULL", func(t *testing.T) {
repo, mock := newChannelModelPricingTimePricingRepo(t)
mock.ExpectExec(`(?s)UPDATE channel_model_pricing.*per_request_price = \$9, time_pricing = \$10, platform = \$11.*WHERE id = \$12`).
mock.ExpectExec(`(?s)UPDATE channel_model_pricing.*per_request_price = \$11, time_pricing = \$12, platform = \$13.*WHERE id = \$14`).
WithArgs(
[]byte(`["gpt-5"]`), service.BillingModeToken,
nil, nil, nil, nil, nil, nil, nil, nil, "openai", int64(11),
nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, "openai", int64(11),
).
WillReturnResult(sqlmock.NewResult(0, 1))
+34 -14
View File
@@ -96,6 +96,8 @@ type ChannelModelPricing struct {
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheReadPrice *float64 `json:"cache_read_price"`
FastMultiplier *float64 `json:"fast_multiplier"`
FlexMultiplier *float64 `json:"flex_multiplier"`
ImageInputPrice *float64 `json:"image_input_price"`
ImageOutputPrice *float64 `json:"image_output_price"`
PerRequestPrice *float64 `json:"per_request_price"`
@@ -120,19 +122,23 @@ type ChannelTimePricingPeriod struct {
// PricingInterval 定价区间(token 区间 / 按次分层 / 图片分辨率分层)
type PricingInterval struct {
ID int64 `json:"id,omitempty"`
PricingID int64 `json:"pricing_id,omitempty"`
MinTokens int `json:"min_tokens"`
MaxTokens *int `json:"max_tokens"`
TierLabel string `json:"tier_label"`
InputPrice *float64 `json:"input_price"`
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheReadPrice *float64 `json:"cache_read_price"`
PerRequestPrice *float64 `json:"per_request_price"`
SortOrder int `json:"sort_order"`
CreatedAt time.Time `json:"created_at,omitempty"`
UpdatedAt time.Time `json:"updated_at,omitempty"`
ID int64 `json:"id,omitempty"`
PricingID int64 `json:"pricing_id,omitempty"`
MinTokens int `json:"min_tokens"`
MaxTokens *int `json:"max_tokens"`
TierLabel string `json:"tier_label"`
InputPrice *float64 `json:"input_price"`
OutputPrice *float64 `json:"output_price"`
CacheWritePrice *float64 `json:"cache_write_price"`
CacheReadPrice *float64 `json:"cache_read_price"`
InputMultiplier *float64 `json:"input_multiplier"`
OutputMultiplier *float64 `json:"output_multiplier"`
CacheWriteMultiplier *float64 `json:"cache_write_multiplier"`
CacheReadMultiplier *float64 `json:"cache_read_multiplier"`
PerRequestPrice *float64 `json:"per_request_price"`
SortOrder int `json:"sort_order"`
CreatedAt time.Time `json:"created_at,omitempty"`
UpdatedAt time.Time `json:"updated_at,omitempty"`
}
// IsActive 判断渠道是否启用
@@ -358,7 +364,7 @@ func validateSingleInterval(iv *PricingInterval, idx int) error {
return validateIntervalPrices(iv, idx)
}
// validateIntervalPrices 校验区间内所有价格字段 >= 0
// validateIntervalPrices 校验区间价格 >= 0、倍率 > 0。
func validateIntervalPrices(iv *PricingInterval, idx int) error {
prices := []struct {
name string
@@ -375,6 +381,20 @@ func validateIntervalPrices(iv *PricingInterval, idx int) error {
return fmt.Errorf("interval #%d: %s must be >= 0", idx+1, p.name)
}
}
multipliers := []struct {
name string
val *float64
}{
{"input_multiplier", iv.InputMultiplier},
{"output_multiplier", iv.OutputMultiplier},
{"cache_write_multiplier", iv.CacheWriteMultiplier},
{"cache_read_multiplier", iv.CacheReadMultiplier},
}
for _, multiplier := range multipliers {
if multiplier.val != nil && *multiplier.val <= 0 {
return fmt.Errorf("interval #%d: %s must be > 0", idx+1, multiplier.name)
}
}
return nil
}
+14 -1
View File
@@ -735,6 +735,17 @@ func checkPricesNotNegative(p ChannelModelPricing) error {
return infraerrors.BadRequest("NEGATIVE_PRICE", fmt.Sprintf("%s must be >= 0", c.field))
}
}
for _, c := range []struct {
field string
val *float64
}{
{"fast_multiplier", p.FastMultiplier},
{"flex_multiplier", p.FlexMultiplier},
} {
if c.val != nil && *c.val <= 0 {
return infraerrors.BadRequest("INVALID_MULTIPLIER", fmt.Sprintf("%s must be > 0", c.field))
}
}
return nil
}
@@ -742,7 +753,9 @@ func checkIntervalsHavePrices(p ChannelModelPricing) error {
for _, iv := range p.Intervals {
if iv.InputPrice == nil && iv.OutputPrice == nil &&
iv.CacheWritePrice == nil && iv.CacheReadPrice == nil &&
iv.PerRequestPrice == nil {
iv.PerRequestPrice == nil && iv.InputMultiplier == nil &&
iv.OutputMultiplier == nil && iv.CacheWriteMultiplier == nil &&
iv.CacheReadMultiplier == nil {
return infraerrors.BadRequest(
"INTERVAL_MISSING_PRICE",
fmt.Sprintf("interval [%d, %s] has no price fields set for model %v",
@@ -0,0 +1,56 @@
ALTER TABLE channel_model_pricing
ADD COLUMN IF NOT EXISTS fast_multiplier NUMERIC(12,6),
ADD COLUMN IF NOT EXISTS flex_multiplier NUMERIC(12,6);
ALTER TABLE channel_pricing_intervals
ADD COLUMN IF NOT EXISTS input_multiplier NUMERIC(12,6),
ADD COLUMN IF NOT EXISTS output_multiplier NUMERIC(12,6),
ADD COLUMN IF NOT EXISTS cache_write_multiplier NUMERIC(12,6),
ADD COLUMN IF NOT EXISTS cache_read_multiplier NUMERIC(12,6);
DO $$
BEGIN
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'channel_model_pricing_fast_multiplier_positive' AND conrelid = 'channel_model_pricing'::regclass) THEN
ALTER TABLE channel_model_pricing
ADD CONSTRAINT channel_model_pricing_fast_multiplier_positive
CHECK (fast_multiplier IS NULL OR fast_multiplier > 0);
END IF;
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'channel_model_pricing_flex_multiplier_positive' AND conrelid = 'channel_model_pricing'::regclass) THEN
ALTER TABLE channel_model_pricing
ADD CONSTRAINT channel_model_pricing_flex_multiplier_positive
CHECK (flex_multiplier IS NULL OR flex_multiplier > 0);
END IF;
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'channel_pricing_intervals_input_multiplier_positive' AND conrelid = 'channel_pricing_intervals'::regclass) THEN
ALTER TABLE channel_pricing_intervals
ADD CONSTRAINT channel_pricing_intervals_input_multiplier_positive
CHECK (input_multiplier IS NULL OR input_multiplier > 0);
END IF;
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'channel_pricing_intervals_output_multiplier_positive' AND conrelid = 'channel_pricing_intervals'::regclass) THEN
ALTER TABLE channel_pricing_intervals
ADD CONSTRAINT channel_pricing_intervals_output_multiplier_positive
CHECK (output_multiplier IS NULL OR output_multiplier > 0);
END IF;
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'channel_pricing_intervals_cache_write_multiplier_positive' AND conrelid = 'channel_pricing_intervals'::regclass) THEN
ALTER TABLE channel_pricing_intervals
ADD CONSTRAINT channel_pricing_intervals_cache_write_multiplier_positive
CHECK (cache_write_multiplier IS NULL OR cache_write_multiplier > 0);
END IF;
IF NOT EXISTS (SELECT 1 FROM pg_constraint WHERE conname = 'channel_pricing_intervals_cache_read_multiplier_positive' AND conrelid = 'channel_pricing_intervals'::regclass) THEN
ALTER TABLE channel_pricing_intervals
ADD CONSTRAINT channel_pricing_intervals_cache_read_multiplier_positive
CHECK (cache_read_multiplier IS NULL OR cache_read_multiplier > 0);
END IF;
END $$;
COMMENT ON COLUMN channel_model_pricing.fast_multiplier IS
'Fast/priority service tier multiplier applied to the selected standard channel price';
COMMENT ON COLUMN channel_model_pricing.flex_multiplier IS
'Flex service tier multiplier applied to the selected standard channel price';
COMMENT ON COLUMN channel_pricing_intervals.input_multiplier IS
'Interval input multiplier applied to the channel base input price when input_price is NULL';
COMMENT ON COLUMN channel_pricing_intervals.output_multiplier IS
'Interval output multiplier applied to the channel base output price when output_price is NULL';
COMMENT ON COLUMN channel_pricing_intervals.cache_write_multiplier IS
'Interval cache-write multiplier applied to the channel base cache-write price when cache_write_price is NULL';
COMMENT ON COLUMN channel_pricing_intervals.cache_read_multiplier IS
'Interval cache-read multiplier applied to the channel base cache-read price when cache_read_price is NULL';
@@ -0,0 +1,43 @@
package migrations
import (
"strings"
"testing"
"github.com/stretchr/testify/require"
)
func TestChannelPricingMultipliersMigration(t *testing.T) {
content, err := FS.ReadFile("228_channel_pricing_multipliers.sql")
require.NoError(t, err)
sql := strings.Join(strings.Fields(string(content)), " ")
for _, column := range []string{
"fast_multiplier NUMERIC(12,6)",
"flex_multiplier NUMERIC(12,6)",
"input_multiplier NUMERIC(12,6)",
"output_multiplier NUMERIC(12,6)",
"cache_write_multiplier NUMERIC(12,6)",
"cache_read_multiplier NUMERIC(12,6)",
} {
require.Contains(t, sql, "ADD COLUMN IF NOT EXISTS "+column)
}
constraints := []struct {
table string
name string
column string
}{
{"channel_model_pricing", "channel_model_pricing_fast_multiplier_positive", "fast_multiplier"},
{"channel_model_pricing", "channel_model_pricing_flex_multiplier_positive", "flex_multiplier"},
{"channel_pricing_intervals", "channel_pricing_intervals_input_multiplier_positive", "input_multiplier"},
{"channel_pricing_intervals", "channel_pricing_intervals_output_multiplier_positive", "output_multiplier"},
{"channel_pricing_intervals", "channel_pricing_intervals_cache_write_multiplier_positive", "cache_write_multiplier"},
{"channel_pricing_intervals", "channel_pricing_intervals_cache_read_multiplier_positive", "cache_read_multiplier"},
}
for _, constraint := range constraints {
require.Contains(t, sql, "conname = '"+constraint.name+"' AND conrelid = '"+constraint.table+"'::regclass")
require.Contains(t, sql, "ALTER TABLE "+constraint.table+" ADD CONSTRAINT "+constraint.name+
" CHECK ("+constraint.column+" IS NULL OR "+constraint.column+" > 0)")
}
}