mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 12:04:00 +08:00
渠道定价:持久化服务层级与区间倍率
This commit is contained in:
@@ -64,6 +64,8 @@ type channelModelPricingRequest struct {
|
||||
OutputPrice *float64 `json:"output_price" binding:"omitempty,min=0"`
|
||||
CacheWritePrice *float64 `json:"cache_write_price" binding:"omitempty,min=0"`
|
||||
CacheReadPrice *float64 `json:"cache_read_price" binding:"omitempty,min=0"`
|
||||
FastMultiplier *float64 `json:"fast_multiplier" binding:"omitempty,gt=0"`
|
||||
FlexMultiplier *float64 `json:"flex_multiplier" binding:"omitempty,gt=0"`
|
||||
ImageInputPrice *float64 `json:"image_input_price" binding:"omitempty,min=0"`
|
||||
ImageOutputPrice *float64 `json:"image_output_price" binding:"omitempty,min=0"`
|
||||
PerRequestPrice *float64 `json:"per_request_price" binding:"omitempty,min=0"`
|
||||
@@ -83,15 +85,19 @@ type channelTimePricingPeriodRequest struct {
|
||||
}
|
||||
|
||||
type pricingIntervalRequest struct {
|
||||
MinTokens int `json:"min_tokens"`
|
||||
MaxTokens *int `json:"max_tokens"`
|
||||
TierLabel string `json:"tier_label"`
|
||||
InputPrice *float64 `json:"input_price"`
|
||||
OutputPrice *float64 `json:"output_price"`
|
||||
CacheWritePrice *float64 `json:"cache_write_price"`
|
||||
CacheReadPrice *float64 `json:"cache_read_price"`
|
||||
PerRequestPrice *float64 `json:"per_request_price"`
|
||||
SortOrder int `json:"sort_order"`
|
||||
MinTokens int `json:"min_tokens"`
|
||||
MaxTokens *int `json:"max_tokens"`
|
||||
TierLabel string `json:"tier_label"`
|
||||
InputPrice *float64 `json:"input_price"`
|
||||
OutputPrice *float64 `json:"output_price"`
|
||||
CacheWritePrice *float64 `json:"cache_write_price"`
|
||||
CacheReadPrice *float64 `json:"cache_read_price"`
|
||||
InputMultiplier *float64 `json:"input_multiplier" binding:"omitempty,gt=0"`
|
||||
OutputMultiplier *float64 `json:"output_multiplier" binding:"omitempty,gt=0"`
|
||||
CacheWriteMultiplier *float64 `json:"cache_write_multiplier" binding:"omitempty,gt=0"`
|
||||
CacheReadMultiplier *float64 `json:"cache_read_multiplier" binding:"omitempty,gt=0"`
|
||||
PerRequestPrice *float64 `json:"per_request_price"`
|
||||
SortOrder int `json:"sort_order"`
|
||||
}
|
||||
|
||||
type accountStatsPricingRuleRequest struct {
|
||||
@@ -128,6 +134,8 @@ type channelModelPricingResponse struct {
|
||||
OutputPrice *float64 `json:"output_price"`
|
||||
CacheWritePrice *float64 `json:"cache_write_price"`
|
||||
CacheReadPrice *float64 `json:"cache_read_price"`
|
||||
FastMultiplier *float64 `json:"fast_multiplier"`
|
||||
FlexMultiplier *float64 `json:"flex_multiplier"`
|
||||
ImageInputPrice *float64 `json:"image_input_price"`
|
||||
ImageOutputPrice *float64 `json:"image_output_price"`
|
||||
PerRequestPrice *float64 `json:"per_request_price"`
|
||||
@@ -147,16 +155,20 @@ type channelTimePricingPeriodResponse struct {
|
||||
}
|
||||
|
||||
type pricingIntervalResponse struct {
|
||||
ID int64 `json:"id"`
|
||||
MinTokens int `json:"min_tokens"`
|
||||
MaxTokens *int `json:"max_tokens"`
|
||||
TierLabel string `json:"tier_label,omitempty"`
|
||||
InputPrice *float64 `json:"input_price"`
|
||||
OutputPrice *float64 `json:"output_price"`
|
||||
CacheWritePrice *float64 `json:"cache_write_price"`
|
||||
CacheReadPrice *float64 `json:"cache_read_price"`
|
||||
PerRequestPrice *float64 `json:"per_request_price"`
|
||||
SortOrder int `json:"sort_order"`
|
||||
ID int64 `json:"id"`
|
||||
MinTokens int `json:"min_tokens"`
|
||||
MaxTokens *int `json:"max_tokens"`
|
||||
TierLabel string `json:"tier_label,omitempty"`
|
||||
InputPrice *float64 `json:"input_price"`
|
||||
OutputPrice *float64 `json:"output_price"`
|
||||
CacheWritePrice *float64 `json:"cache_write_price"`
|
||||
CacheReadPrice *float64 `json:"cache_read_price"`
|
||||
InputMultiplier *float64 `json:"input_multiplier"`
|
||||
OutputMultiplier *float64 `json:"output_multiplier"`
|
||||
CacheWriteMultiplier *float64 `json:"cache_write_multiplier"`
|
||||
CacheReadMultiplier *float64 `json:"cache_read_multiplier"`
|
||||
PerRequestPrice *float64 `json:"per_request_price"`
|
||||
SortOrder int `json:"sort_order"`
|
||||
}
|
||||
|
||||
type accountStatsPricingRuleResponse struct {
|
||||
@@ -248,6 +260,8 @@ func pricingToResponse(p *service.ChannelModelPricing) channelModelPricingRespon
|
||||
OutputPrice: p.OutputPrice,
|
||||
CacheWritePrice: p.CacheWritePrice,
|
||||
CacheReadPrice: p.CacheReadPrice,
|
||||
FastMultiplier: p.FastMultiplier,
|
||||
FlexMultiplier: p.FlexMultiplier,
|
||||
ImageInputPrice: p.ImageInputPrice,
|
||||
ImageOutputPrice: p.ImageOutputPrice,
|
||||
PerRequestPrice: p.PerRequestPrice,
|
||||
@@ -273,20 +287,24 @@ func timePricingToResponse(value *service.ChannelTimePricing) *channelTimePricin
|
||||
|
||||
func intervalToResponse(iv service.PricingInterval) pricingIntervalResponse {
|
||||
return pricingIntervalResponse{
|
||||
ID: iv.ID,
|
||||
MinTokens: iv.MinTokens,
|
||||
MaxTokens: iv.MaxTokens,
|
||||
TierLabel: iv.TierLabel,
|
||||
InputPrice: iv.InputPrice,
|
||||
OutputPrice: iv.OutputPrice,
|
||||
CacheWritePrice: iv.CacheWritePrice,
|
||||
CacheReadPrice: iv.CacheReadPrice,
|
||||
PerRequestPrice: iv.PerRequestPrice,
|
||||
SortOrder: iv.SortOrder,
|
||||
ID: iv.ID,
|
||||
MinTokens: iv.MinTokens,
|
||||
MaxTokens: iv.MaxTokens,
|
||||
TierLabel: iv.TierLabel,
|
||||
InputPrice: iv.InputPrice,
|
||||
OutputPrice: iv.OutputPrice,
|
||||
CacheWritePrice: iv.CacheWritePrice,
|
||||
CacheReadPrice: iv.CacheReadPrice,
|
||||
InputMultiplier: iv.InputMultiplier,
|
||||
OutputMultiplier: iv.OutputMultiplier,
|
||||
CacheWriteMultiplier: iv.CacheWriteMultiplier,
|
||||
CacheReadMultiplier: iv.CacheReadMultiplier,
|
||||
PerRequestPrice: iv.PerRequestPrice,
|
||||
SortOrder: iv.SortOrder,
|
||||
}
|
||||
}
|
||||
|
||||
func pricingRequestToService(reqs []channelModelPricingRequest) []service.ChannelModelPricing {
|
||||
func pricingRequestToService(reqs []channelModelPricingRequest, allowChannelMultipliers bool) []service.ChannelModelPricing {
|
||||
result := make([]service.ChannelModelPricing, 0, len(reqs))
|
||||
for _, r := range reqs {
|
||||
billingMode := service.BillingMode(r.BillingMode)
|
||||
@@ -296,18 +314,34 @@ func pricingRequestToService(reqs []channelModelPricingRequest) []service.Channe
|
||||
platform := r.Platform
|
||||
intervals := make([]service.PricingInterval, 0, len(r.Intervals))
|
||||
for _, iv := range r.Intervals {
|
||||
var inputMultiplier, outputMultiplier, cacheWriteMultiplier, cacheReadMultiplier *float64
|
||||
if allowChannelMultipliers {
|
||||
inputMultiplier = iv.InputMultiplier
|
||||
outputMultiplier = iv.OutputMultiplier
|
||||
cacheWriteMultiplier = iv.CacheWriteMultiplier
|
||||
cacheReadMultiplier = iv.CacheReadMultiplier
|
||||
}
|
||||
intervals = append(intervals, service.PricingInterval{
|
||||
MinTokens: iv.MinTokens,
|
||||
MaxTokens: iv.MaxTokens,
|
||||
TierLabel: iv.TierLabel,
|
||||
InputPrice: iv.InputPrice,
|
||||
OutputPrice: iv.OutputPrice,
|
||||
CacheWritePrice: iv.CacheWritePrice,
|
||||
CacheReadPrice: iv.CacheReadPrice,
|
||||
PerRequestPrice: iv.PerRequestPrice,
|
||||
SortOrder: iv.SortOrder,
|
||||
MinTokens: iv.MinTokens,
|
||||
MaxTokens: iv.MaxTokens,
|
||||
TierLabel: iv.TierLabel,
|
||||
InputPrice: iv.InputPrice,
|
||||
OutputPrice: iv.OutputPrice,
|
||||
CacheWritePrice: iv.CacheWritePrice,
|
||||
CacheReadPrice: iv.CacheReadPrice,
|
||||
InputMultiplier: inputMultiplier,
|
||||
OutputMultiplier: outputMultiplier,
|
||||
CacheWriteMultiplier: cacheWriteMultiplier,
|
||||
CacheReadMultiplier: cacheReadMultiplier,
|
||||
PerRequestPrice: iv.PerRequestPrice,
|
||||
SortOrder: iv.SortOrder,
|
||||
})
|
||||
}
|
||||
var fastMultiplier, flexMultiplier *float64
|
||||
if allowChannelMultipliers {
|
||||
fastMultiplier = r.FastMultiplier
|
||||
flexMultiplier = r.FlexMultiplier
|
||||
}
|
||||
result = append(result, service.ChannelModelPricing{
|
||||
Platform: platform,
|
||||
Models: r.Models,
|
||||
@@ -316,6 +350,8 @@ func pricingRequestToService(reqs []channelModelPricingRequest) []service.Channe
|
||||
OutputPrice: r.OutputPrice,
|
||||
CacheWritePrice: r.CacheWritePrice,
|
||||
CacheReadPrice: r.CacheReadPrice,
|
||||
FastMultiplier: fastMultiplier,
|
||||
FlexMultiplier: flexMultiplier,
|
||||
ImageInputPrice: r.ImageInputPrice,
|
||||
ImageOutputPrice: r.ImageOutputPrice,
|
||||
PerRequestPrice: r.PerRequestPrice,
|
||||
@@ -346,7 +382,7 @@ func accountStatsPricingRuleRequestToService(r accountStatsPricingRuleRequest) s
|
||||
Name: r.Name,
|
||||
GroupIDs: r.GroupIDs,
|
||||
AccountIDs: r.AccountIDs,
|
||||
Pricing: pricingRequestToService(r.Pricing),
|
||||
Pricing: pricingRequestToService(r.Pricing, false),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -407,7 +443,7 @@ func (h *ChannelHandler) Create(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
pricing := pricingRequestToService(req.ModelPricing)
|
||||
pricing := pricingRequestToService(req.ModelPricing, true)
|
||||
// Main model_pricing requires a platform; default to anthropic for backward compatibility.
|
||||
for i := range pricing {
|
||||
if pricing[i].Platform == "" {
|
||||
@@ -481,7 +517,7 @@ func (h *ChannelHandler) Update(c *gin.Context) {
|
||||
ApplyPricingToAccountStats: req.ApplyPricingToAccountStats,
|
||||
}
|
||||
if req.ModelPricing != nil {
|
||||
pricing := pricingRequestToService(*req.ModelPricing)
|
||||
pricing := pricingRequestToService(*req.ModelPricing, true)
|
||||
for i := range pricing {
|
||||
if pricing[i].Platform == "" {
|
||||
pricing[i].Platform = service.PlatformAnthropic
|
||||
|
||||
@@ -305,7 +305,7 @@ func TestPricingRequestToService_Defaults(t *testing.T) {
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := pricingRequestToService([]channelModelPricingRequest{tt.req})
|
||||
result := pricingRequestToService([]channelModelPricingRequest{tt.req}, true)
|
||||
require.Len(t, result, 1)
|
||||
switch tt.wantField {
|
||||
case "BillingMode":
|
||||
@@ -332,7 +332,7 @@ func TestPricingRequestToService_WithAllFields(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
result := pricingRequestToService(reqs)
|
||||
result := pricingRequestToService(reqs, true)
|
||||
require.Len(t, result, 1)
|
||||
r := result[0]
|
||||
require.Equal(t, "openai", r.Platform)
|
||||
@@ -373,7 +373,7 @@ func TestPricingRequestToService_WithIntervals(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
result := pricingRequestToService(reqs)
|
||||
result := pricingRequestToService(reqs, true)
|
||||
require.Len(t, result, 1)
|
||||
require.Len(t, result[0].Intervals, 2)
|
||||
|
||||
@@ -396,7 +396,7 @@ func TestPricingRequestToService_WithIntervals(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestPricingRequestToService_EmptySlice(t *testing.T) {
|
||||
result := pricingRequestToService([]channelModelPricingRequest{})
|
||||
result := pricingRequestToService([]channelModelPricingRequest{}, true)
|
||||
require.NotNil(t, result)
|
||||
require.Empty(t, result)
|
||||
}
|
||||
@@ -410,7 +410,7 @@ func TestPricingRequestToService_NilPriceFields(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
result := pricingRequestToService(reqs)
|
||||
result := pricingRequestToService(reqs, true)
|
||||
require.Len(t, result, 1)
|
||||
r := result[0]
|
||||
require.Nil(t, r.InputPrice)
|
||||
@@ -433,16 +433,52 @@ func TestPricingRequestToService_TimePricing(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
got := pricingRequestToService([]channelModelPricingRequest{req})
|
||||
got := pricingRequestToService([]channelModelPricingRequest{req}, true)
|
||||
require.Equal(t, "Asia/Shanghai", got[0].TimePricing.Timezone)
|
||||
require.Equal(t, 2.0, got[0].TimePricing.Periods[0].Multiplier)
|
||||
}
|
||||
|
||||
func TestPricingRequestToService_TimePricingNil(t *testing.T) {
|
||||
got := pricingRequestToService([]channelModelPricingRequest{{Models: []string{"gpt-5"}}})
|
||||
got := pricingRequestToService([]channelModelPricingRequest{{Models: []string{"gpt-5"}}}, true)
|
||||
require.Nil(t, got[0].TimePricing)
|
||||
}
|
||||
|
||||
// 账号成本统计规则不支持倍率:allowChannelMultipliers=false 时必须丢弃,
|
||||
// 避免渠道倍率意外污染账号成本口径。
|
||||
func TestPricingRequestToService_MultipliersGatedByFlag(t *testing.T) {
|
||||
req := channelModelPricingRequest{
|
||||
Models: []string{"gpt-5"},
|
||||
BillingMode: "token",
|
||||
FastMultiplier: float64Ptr(2.5),
|
||||
FlexMultiplier: float64Ptr(0.5),
|
||||
Intervals: []pricingIntervalRequest{{
|
||||
MinTokens: 272000,
|
||||
InputMultiplier: float64Ptr(2),
|
||||
OutputMultiplier: float64Ptr(1.5),
|
||||
CacheWriteMultiplier: float64Ptr(2),
|
||||
CacheReadMultiplier: float64Ptr(2),
|
||||
}},
|
||||
}
|
||||
|
||||
allowed := pricingRequestToService([]channelModelPricingRequest{req}, true)
|
||||
require.Equal(t, float64Ptr(2.5), allowed[0].FastMultiplier)
|
||||
require.Equal(t, float64Ptr(0.5), allowed[0].FlexMultiplier)
|
||||
require.Equal(t, float64Ptr(2), allowed[0].Intervals[0].InputMultiplier)
|
||||
require.Equal(t, float64Ptr(1.5), allowed[0].Intervals[0].OutputMultiplier)
|
||||
require.Equal(t, float64Ptr(2), allowed[0].Intervals[0].CacheWriteMultiplier)
|
||||
require.Equal(t, float64Ptr(2), allowed[0].Intervals[0].CacheReadMultiplier)
|
||||
|
||||
dropped := pricingRequestToService([]channelModelPricingRequest{req}, false)
|
||||
require.Nil(t, dropped[0].FastMultiplier)
|
||||
require.Nil(t, dropped[0].FlexMultiplier)
|
||||
require.Nil(t, dropped[0].Intervals[0].InputMultiplier)
|
||||
require.Nil(t, dropped[0].Intervals[0].OutputMultiplier)
|
||||
require.Nil(t, dropped[0].Intervals[0].CacheWriteMultiplier)
|
||||
require.Nil(t, dropped[0].Intervals[0].CacheReadMultiplier)
|
||||
// 非倍率字段不受开关影响
|
||||
require.Equal(t, 272000, dropped[0].Intervals[0].MinTokens)
|
||||
}
|
||||
|
||||
func TestPricingToResponse_TimePricing(t *testing.T) {
|
||||
got := pricingToResponse(&service.ChannelModelPricing{
|
||||
BillingMode: service.BillingModeToken,
|
||||
|
||||
@@ -16,7 +16,7 @@ import (
|
||||
|
||||
func (r *channelRepository) ListModelPricing(ctx context.Context, channelID int64) ([]service.ChannelModelPricing, error) {
|
||||
rows, err := r.db.QueryContext(ctx,
|
||||
`SELECT id, channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, image_input_price, image_output_price, per_request_price, time_pricing, created_at, updated_at
|
||||
`SELECT id, channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, fast_multiplier, flex_multiplier, image_input_price, image_output_price, per_request_price, time_pricing, created_at, updated_at
|
||||
FROM channel_model_pricing WHERE channel_id = $1 ORDER BY id`, channelID,
|
||||
)
|
||||
if err != nil {
|
||||
@@ -61,10 +61,11 @@ func (r *channelRepository) UpdateModelPricing(ctx context.Context, pricing *ser
|
||||
}
|
||||
result, err := r.db.ExecContext(ctx,
|
||||
`UPDATE channel_model_pricing
|
||||
SET models = $1, billing_mode = $2, input_price = $3, output_price = $4, cache_write_price = $5, cache_read_price = $6, image_input_price = $7, image_output_price = $8, per_request_price = $9, time_pricing = $10, platform = $11, updated_at = NOW()
|
||||
WHERE id = $12`,
|
||||
SET models = $1, billing_mode = $2, input_price = $3, output_price = $4, cache_write_price = $5, cache_read_price = $6, fast_multiplier = $7, flex_multiplier = $8, image_input_price = $9, image_output_price = $10, per_request_price = $11, time_pricing = $12, platform = $13, updated_at = NOW()
|
||||
WHERE id = $14`,
|
||||
modelsJSON, billingMode, pricing.InputPrice, pricing.OutputPrice, pricing.CacheWritePrice, pricing.CacheReadPrice,
|
||||
pricing.ImageInputPrice, pricing.ImageOutputPrice, pricing.PerRequestPrice, timePricingJSON, pricing.Platform, pricing.ID,
|
||||
pricing.FastMultiplier, pricing.FlexMultiplier, pricing.ImageInputPrice, pricing.ImageOutputPrice, pricing.PerRequestPrice,
|
||||
timePricingJSON, pricing.Platform, pricing.ID,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("update model pricing: %w", err)
|
||||
@@ -95,7 +96,7 @@ func (r *channelRepository) ReplaceModelPricing(ctx context.Context, channelID i
|
||||
// batchLoadModelPricing 批量加载多个渠道的模型定价(含区间)
|
||||
func (r *channelRepository) batchLoadModelPricing(ctx context.Context, channelIDs []int64) (map[int64][]service.ChannelModelPricing, error) {
|
||||
rows, err := r.db.QueryContext(ctx,
|
||||
`SELECT id, channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, image_input_price, image_output_price, per_request_price, time_pricing, created_at, updated_at
|
||||
`SELECT id, channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, fast_multiplier, flex_multiplier, image_input_price, image_output_price, per_request_price, time_pricing, created_at, updated_at
|
||||
FROM channel_model_pricing WHERE channel_id = ANY($1) ORDER BY channel_id, id`,
|
||||
pq.Array(channelIDs),
|
||||
)
|
||||
@@ -136,6 +137,7 @@ func (r *channelRepository) batchLoadIntervals(ctx context.Context, pricingIDs [
|
||||
rows, err := r.db.QueryContext(ctx,
|
||||
`SELECT id, pricing_id, min_tokens, max_tokens, tier_label,
|
||||
input_price, output_price, cache_write_price, cache_read_price,
|
||||
input_multiplier, output_multiplier, cache_write_multiplier, cache_read_multiplier,
|
||||
per_request_price, sort_order, created_at, updated_at
|
||||
FROM channel_pricing_intervals
|
||||
WHERE pricing_id = ANY($1) ORDER BY pricing_id, sort_order, id`,
|
||||
@@ -152,6 +154,7 @@ func (r *channelRepository) batchLoadIntervals(ctx context.Context, pricingIDs [
|
||||
if err := rows.Scan(
|
||||
&iv.ID, &iv.PricingID, &iv.MinTokens, &iv.MaxTokens, &iv.TierLabel,
|
||||
&iv.InputPrice, &iv.OutputPrice, &iv.CacheWritePrice, &iv.CacheReadPrice,
|
||||
&iv.InputMultiplier, &iv.OutputMultiplier, &iv.CacheWriteMultiplier, &iv.CacheReadMultiplier,
|
||||
&iv.PerRequestPrice, &iv.SortOrder, &iv.CreatedAt, &iv.UpdatedAt,
|
||||
); err != nil {
|
||||
return nil, fmt.Errorf("scan interval: %w", err)
|
||||
@@ -177,6 +180,7 @@ func scanModelPricingRows(rows *sql.Rows) ([]service.ChannelModelPricing, []int6
|
||||
if err := rows.Scan(
|
||||
&p.ID, &p.ChannelID, &p.Platform, &modelsJSON, &p.BillingMode,
|
||||
&p.InputPrice, &p.OutputPrice, &p.CacheWritePrice, &p.CacheReadPrice,
|
||||
&p.FastMultiplier, &p.FlexMultiplier,
|
||||
&p.ImageInputPrice, &p.ImageOutputPrice, &p.PerRequestPrice, &timePricingJSON, &p.CreatedAt, &p.UpdatedAt,
|
||||
); err != nil {
|
||||
return nil, nil, fmt.Errorf("scan model pricing: %w", err)
|
||||
@@ -243,11 +247,12 @@ func createModelPricingExec(ctx context.Context, exec dbExec, pricing *service.C
|
||||
platform = "anthropic"
|
||||
}
|
||||
err = exec.QueryRowContext(ctx,
|
||||
`INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, image_input_price, image_output_price, per_request_price, time_pricing)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12) RETURNING id, created_at, updated_at`,
|
||||
`INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, fast_multiplier, flex_multiplier, image_input_price, image_output_price, per_request_price, time_pricing)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14) RETURNING id, created_at, updated_at`,
|
||||
pricing.ChannelID, platform, modelsJSON, billingMode,
|
||||
pricing.InputPrice, pricing.OutputPrice, pricing.CacheWritePrice, pricing.CacheReadPrice,
|
||||
pricing.ImageInputPrice, pricing.ImageOutputPrice, pricing.PerRequestPrice, timePricingJSON,
|
||||
pricing.FastMultiplier, pricing.FlexMultiplier, pricing.ImageInputPrice, pricing.ImageOutputPrice,
|
||||
pricing.PerRequestPrice, timePricingJSON,
|
||||
).Scan(&pricing.ID, &pricing.CreatedAt, &pricing.UpdatedAt)
|
||||
if err != nil {
|
||||
return fmt.Errorf("insert model pricing: %w", err)
|
||||
@@ -288,10 +293,11 @@ func unmarshalChannelTimePricing(data []byte) (*service.ChannelTimePricing, erro
|
||||
func createIntervalExec(ctx context.Context, exec dbExec, iv *service.PricingInterval) error {
|
||||
return exec.QueryRowContext(ctx,
|
||||
`INSERT INTO channel_pricing_intervals
|
||||
(pricing_id, min_tokens, max_tokens, tier_label, input_price, output_price, cache_write_price, cache_read_price, per_request_price, sort_order)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10) RETURNING id, created_at, updated_at`,
|
||||
(pricing_id, min_tokens, max_tokens, tier_label, input_price, output_price, cache_write_price, cache_read_price, input_multiplier, output_multiplier, cache_write_multiplier, cache_read_multiplier, per_request_price, sort_order)
|
||||
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14) RETURNING id, created_at, updated_at`,
|
||||
iv.PricingID, iv.MinTokens, iv.MaxTokens, iv.TierLabel,
|
||||
iv.InputPrice, iv.OutputPrice, iv.CacheWritePrice, iv.CacheReadPrice,
|
||||
iv.InputMultiplier, iv.OutputMultiplier, iv.CacheWriteMultiplier, iv.CacheReadMultiplier,
|
||||
iv.PerRequestPrice, iv.SortOrder,
|
||||
).Scan(&iv.ID, &iv.CreatedAt, &iv.UpdatedAt)
|
||||
}
|
||||
|
||||
@@ -16,7 +16,7 @@ import (
|
||||
|
||||
var channelModelPricingTimePricingColumns = []string{
|
||||
"id", "channel_id", "platform", "models", "billing_mode", "input_price", "output_price",
|
||||
"cache_write_price", "cache_read_price", "image_input_price", "image_output_price",
|
||||
"cache_write_price", "cache_read_price", "fast_multiplier", "flex_multiplier", "image_input_price", "image_output_price",
|
||||
"per_request_price", "time_pricing", "created_at", "updated_at",
|
||||
}
|
||||
|
||||
@@ -33,7 +33,7 @@ func newChannelModelPricingTimePricingRepo(t *testing.T) (*channelRepository, sq
|
||||
func modelPricingTimePricingRow(timePricing any) *sqlmock.Rows {
|
||||
return sqlmock.NewRows(channelModelPricingTimePricingColumns).AddRow(
|
||||
int64(11), int64(7), "openai", `["gpt-5"]`, service.BillingModeToken,
|
||||
nil, nil, nil, nil, nil, nil, nil, timePricing,
|
||||
nil, nil, nil, nil, nil, nil, nil, nil, nil, timePricing,
|
||||
time.Date(2026, 8, 17, 0, 0, 0, 0, time.UTC), time.Date(2026, 8, 17, 1, 0, 0, 0, time.UTC),
|
||||
)
|
||||
}
|
||||
@@ -105,10 +105,10 @@ func TestChannelModelPricingTimePricingCreateAndUpdateRoundTrip(t *testing.T) {
|
||||
|
||||
t.Run("create writes JSON", func(t *testing.T) {
|
||||
repo, mock := newChannelModelPricingTimePricingRepo(t)
|
||||
mock.ExpectQuery(regexp.QuoteMeta("INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, image_input_price, image_output_price, per_request_price, time_pricing)")).
|
||||
mock.ExpectQuery(regexp.QuoteMeta("INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, fast_multiplier, flex_multiplier, image_input_price, image_output_price, per_request_price, time_pricing)")).
|
||||
WithArgs(
|
||||
int64(7), "openai", []byte(`["gpt-5"]`), service.BillingModeToken,
|
||||
nil, nil, nil, nil, nil, nil, nil, channelModelPricingTimePricingJSON,
|
||||
nil, nil, nil, nil, nil, nil, nil, nil, nil, channelModelPricingTimePricingJSON,
|
||||
).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"id", "created_at", "updated_at"}).AddRow(int64(11), time.Time{}, time.Time{}))
|
||||
|
||||
@@ -118,10 +118,10 @@ func TestChannelModelPricingTimePricingCreateAndUpdateRoundTrip(t *testing.T) {
|
||||
|
||||
t.Run("update writes JSON and entry ID", func(t *testing.T) {
|
||||
repo, mock := newChannelModelPricingTimePricingRepo(t)
|
||||
mock.ExpectExec(`(?s)UPDATE channel_model_pricing.*per_request_price = \$9, time_pricing = \$10, platform = \$11.*WHERE id = \$12`).
|
||||
mock.ExpectExec(`(?s)UPDATE channel_model_pricing.*per_request_price = \$11, time_pricing = \$12, platform = \$13.*WHERE id = \$14`).
|
||||
WithArgs(
|
||||
[]byte(`["gpt-5"]`), service.BillingModeToken,
|
||||
nil, nil, nil, nil, nil, nil, nil, channelModelPricingTimePricingJSON, "openai", int64(11),
|
||||
nil, nil, nil, nil, nil, nil, nil, nil, nil, channelModelPricingTimePricingJSON, "openai", int64(11),
|
||||
).
|
||||
WillReturnResult(sqlmock.NewResult(0, 1))
|
||||
|
||||
@@ -153,10 +153,10 @@ func TestChannelModelPricingTimePricingCreateAndUpdateWriteNullWhenDisabled(t *t
|
||||
|
||||
t.Run("create writes SQL NULL", func(t *testing.T) {
|
||||
repo, mock := newChannelModelPricingTimePricingRepo(t)
|
||||
mock.ExpectQuery(regexp.QuoteMeta("INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, image_input_price, image_output_price, per_request_price, time_pricing)")).
|
||||
mock.ExpectQuery(regexp.QuoteMeta("INSERT INTO channel_model_pricing (channel_id, platform, models, billing_mode, input_price, output_price, cache_write_price, cache_read_price, fast_multiplier, flex_multiplier, image_input_price, image_output_price, per_request_price, time_pricing)")).
|
||||
WithArgs(
|
||||
int64(7), "openai", []byte(`["gpt-5"]`), service.BillingModeToken,
|
||||
nil, nil, nil, nil, nil, nil, nil, nil,
|
||||
nil, nil, nil, nil, nil, nil, nil, nil, nil, nil,
|
||||
).
|
||||
WillReturnRows(sqlmock.NewRows([]string{"id", "created_at", "updated_at"}).AddRow(int64(11), time.Time{}, time.Time{}))
|
||||
|
||||
@@ -166,10 +166,10 @@ func TestChannelModelPricingTimePricingCreateAndUpdateWriteNullWhenDisabled(t *t
|
||||
|
||||
t.Run("update writes SQL NULL", func(t *testing.T) {
|
||||
repo, mock := newChannelModelPricingTimePricingRepo(t)
|
||||
mock.ExpectExec(`(?s)UPDATE channel_model_pricing.*per_request_price = \$9, time_pricing = \$10, platform = \$11.*WHERE id = \$12`).
|
||||
mock.ExpectExec(`(?s)UPDATE channel_model_pricing.*per_request_price = \$11, time_pricing = \$12, platform = \$13.*WHERE id = \$14`).
|
||||
WithArgs(
|
||||
[]byte(`["gpt-5"]`), service.BillingModeToken,
|
||||
nil, nil, nil, nil, nil, nil, nil, nil, "openai", int64(11),
|
||||
nil, nil, nil, nil, nil, nil, nil, nil, nil, nil, "openai", int64(11),
|
||||
).
|
||||
WillReturnResult(sqlmock.NewResult(0, 1))
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user