diff --git a/backend/internal/handler/admin/channel_handler.go b/backend/internal/handler/admin/channel_handler.go index ade8f0c95c..90c1a69b6d 100644 --- a/backend/internal/handler/admin/channel_handler.go +++ b/backend/internal/handler/admin/channel_handler.go @@ -57,17 +57,29 @@ type updateChannelRequest struct { } type channelModelPricingRequest struct { - Platform string `json:"platform" binding:"omitempty,max=50"` - Models []string `json:"models" binding:"required,min=1,max=100"` - BillingMode string `json:"billing_mode" binding:"omitempty,oneof=token per_request image"` - InputPrice *float64 `json:"input_price" binding:"omitempty,min=0"` - 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"` - 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"` - Intervals []pricingIntervalRequest `json:"intervals"` + Platform string `json:"platform" binding:"omitempty,max=50"` + Models []string `json:"models" binding:"required,min=1,max=100"` + BillingMode string `json:"billing_mode" binding:"omitempty,oneof=token per_request image"` + InputPrice *float64 `json:"input_price" binding:"omitempty,min=0"` + 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"` + 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"` + Intervals []pricingIntervalRequest `json:"intervals"` + TimePricing *channelTimePricingRequest `json:"time_pricing"` +} + +type channelTimePricingRequest struct { + Timezone string `json:"timezone"` + Periods []channelTimePricingPeriodRequest `json:"periods"` +} + +type channelTimePricingPeriodRequest struct { + StartTime string `json:"start_time"` + EndTime string `json:"end_time"` + Multiplier float64 `json:"multiplier"` } type pricingIntervalRequest struct { @@ -108,18 +120,30 @@ type channelResponse struct { } type channelModelPricingResponse struct { - ID int64 `json:"id"` - Platform string `json:"platform"` - Models []string `json:"models"` - BillingMode string `json:"billing_mode"` - InputPrice *float64 `json:"input_price"` - OutputPrice *float64 `json:"output_price"` - CacheWritePrice *float64 `json:"cache_write_price"` - CacheReadPrice *float64 `json:"cache_read_price"` - ImageInputPrice *float64 `json:"image_input_price"` - ImageOutputPrice *float64 `json:"image_output_price"` - PerRequestPrice *float64 `json:"per_request_price"` - Intervals []pricingIntervalResponse `json:"intervals"` + ID int64 `json:"id"` + Platform string `json:"platform"` + Models []string `json:"models"` + BillingMode string `json:"billing_mode"` + InputPrice *float64 `json:"input_price"` + OutputPrice *float64 `json:"output_price"` + CacheWritePrice *float64 `json:"cache_write_price"` + CacheReadPrice *float64 `json:"cache_read_price"` + ImageInputPrice *float64 `json:"image_input_price"` + ImageOutputPrice *float64 `json:"image_output_price"` + PerRequestPrice *float64 `json:"per_request_price"` + Intervals []pricingIntervalResponse `json:"intervals"` + TimePricing *channelTimePricingResponse `json:"time_pricing"` +} + +type channelTimePricingResponse struct { + Timezone string `json:"timezone"` + Periods []channelTimePricingPeriodResponse `json:"periods"` +} + +type channelTimePricingPeriodResponse struct { + StartTime string `json:"start_time"` + EndTime string `json:"end_time"` + Multiplier float64 `json:"multiplier"` } type pricingIntervalResponse struct { @@ -228,9 +252,25 @@ func pricingToResponse(p *service.ChannelModelPricing) channelModelPricingRespon ImageOutputPrice: p.ImageOutputPrice, PerRequestPrice: p.PerRequestPrice, Intervals: intervals, + TimePricing: timePricingToResponse(p.TimePricing), } } +func timePricingToResponse(value *service.ChannelTimePricing) *channelTimePricingResponse { + if value == nil { + return nil + } + periods := make([]channelTimePricingPeriodResponse, 0, len(value.Periods)) + for _, period := range value.Periods { + periods = append(periods, channelTimePricingPeriodResponse{ + StartTime: period.StartTime, + EndTime: period.EndTime, + Multiplier: period.Multiplier, + }) + } + return &channelTimePricingResponse{Timezone: value.Timezone, Periods: periods} +} + func intervalToResponse(iv service.PricingInterval) pricingIntervalResponse { return pricingIntervalResponse{ ID: iv.ID, @@ -280,11 +320,27 @@ func pricingRequestToService(reqs []channelModelPricingRequest) []service.Channe ImageOutputPrice: r.ImageOutputPrice, PerRequestPrice: r.PerRequestPrice, Intervals: intervals, + TimePricing: timePricingRequestToService(r.TimePricing), }) } return result } +func timePricingRequestToService(value *channelTimePricingRequest) *service.ChannelTimePricing { + if value == nil { + return nil + } + periods := make([]service.ChannelTimePricingPeriod, 0, len(value.Periods)) + for _, period := range value.Periods { + periods = append(periods, service.ChannelTimePricingPeriod{ + StartTime: period.StartTime, + EndTime: period.EndTime, + Multiplier: period.Multiplier, + }) + } + return &service.ChannelTimePricing{Timezone: value.Timezone, Periods: periods} +} + func accountStatsPricingRuleRequestToService(r accountStatsPricingRuleRequest) service.AccountStatsPricingRule { return service.AccountStatsPricingRule{ Name: r.Name, diff --git a/backend/internal/handler/admin/channel_handler_test.go b/backend/internal/handler/admin/channel_handler_test.go index d05a1a6a3b..7f9c050fa5 100644 --- a/backend/internal/handler/admin/channel_handler_test.go +++ b/backend/internal/handler/admin/channel_handler_test.go @@ -421,6 +421,49 @@ func TestPricingRequestToService_NilPriceFields(t *testing.T) { require.Nil(t, r.PerRequestPrice) } +func TestPricingRequestToService_TimePricing(t *testing.T) { + req := channelModelPricingRequest{ + Models: []string{"gpt-5"}, + BillingMode: "token", + TimePricing: &channelTimePricingRequest{ + Timezone: "Asia/Shanghai", + Periods: []channelTimePricingPeriodRequest{{ + StartTime: "09:00", EndTime: "12:00", Multiplier: 2, + }}, + }, + } + + got := pricingRequestToService([]channelModelPricingRequest{req}) + 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"}}}) + require.Nil(t, got[0].TimePricing) +} + +func TestPricingToResponse_TimePricing(t *testing.T) { + got := pricingToResponse(&service.ChannelModelPricing{ + BillingMode: service.BillingModeToken, + TimePricing: &service.ChannelTimePricing{ + Timezone: "Asia/Shanghai", + Periods: []service.ChannelTimePricingPeriod{{ + StartTime: "14:00", EndTime: "18:00", Multiplier: 1.25, + }}, + }, + }) + + require.NotNil(t, got.TimePricing) + require.Equal(t, "Asia/Shanghai", got.TimePricing.Timezone) + require.Equal(t, 1.25, got.TimePricing.Periods[0].Multiplier) +} + +func TestPricingToResponse_TimePricingNil(t *testing.T) { + got := pricingToResponse(&service.ChannelModelPricing{}) + require.Nil(t, got.TimePricing) +} + // --------------------------------------------------------------------------- // 3. SyncPricingModels handler // --------------------------------------------------------------------------- diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 28515d98f0..c9b034c141 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -1420,12 +1420,10 @@ const ( // 由 BeforeTurn 在每个 turn 开始时冻结,AfterTurn 的用量提交读取它;turn 在 // 连接内串行推进,互斥锁只为跨用量提交 goroutine 的读取安全。 // -// 零值语义(重要):ws_v2 passthrough ingress 只实现了 AfterTurn,没有任何 -// turn 起始回调,BeforeTurn 永远不会被调用。此时本值保持零,RecordUsage 经 -// openAIUsagePricingAt 回退到记录时刻——与引入分组利润控制前的基线一致。 -// 绝不能用建连时刻初始化:那会把透传连接的所有 turn 钉死在建连时的高峰因子, -// 客户端只要峰前一分钟建连并保活,整条连接就能全程按谷价结算,正是利润控制 -// 想堵的漏洞。透传 ingress 目前不做 turn 级利润复核,只有建连时的准入门。 +// ws_v2 passthrough ingress 没有 BeforeTurn,因此本值会保持零;AfterTurn 必须 +// 以 TurnStarted 已记录的所属 turn 开始时刻为回退,而不是用建连或记录时刻。 +// 这样每个 passthrough turn 都按自己的开始时刻计价,但不改变其仅在建连时执行 +// 准入门、没有 turn 级利润复核的既有行为。 type openAIWSTurnPricing struct { mu sync.Mutex at time.Time @@ -1437,10 +1435,13 @@ func (p *openAIWSTurnPricing) freeze(at time.Time) { p.mu.Unlock() } -func (p *openAIWSTurnPricing) current() time.Time { +func (p *openAIWSTurnPricing) currentOr(fallback time.Time) time.Time { p.mu.Lock() defer p.mu.Unlock() - return p.at + if !p.at.IsZero() { + return p.at + } + return fallback } // recordOpenAIProfitVeto 记录 OpenAI 侧选号循环的一次利润门终检否决:把账号 @@ -1703,6 +1704,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "missing first response.create message") return } + firstTurnStartedAt := time.Now() if msgType != coderws.MessageText && msgType != coderws.MessageBinary { closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "unsupported websocket message type") return @@ -2031,19 +2033,37 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { maxReasoningEffort, reasoningEffortMappings, _ := openAIReasoningEffortPolicyForRequest(c, apiKey) var requestPayloadHash string + var turnStartsMu sync.Mutex + turnStarts := make(map[int]time.Time, 4) + recordTurnStart := func(turn int, startedAt time.Time) { + if turn <= 0 || startedAt.IsZero() { + return + } + turnStartsMu.Lock() + turnStarts[turn] = startedAt + turnStartsMu.Unlock() + } + getTurnStart := func(turn int) time.Time { + turnStartsMu.Lock() + startedAt := turnStarts[turn] + delete(turnStarts, turn) + turnStartsMu.Unlock() + return startedAt + } // Passthrough rejects overlapping response.create frames, so one immutable // turn-tagged slot preserves the exact mapping used for the in-flight request. var turnChannelMapping atomic.Pointer[openAIWSTurnChannelMappingSnapshot] turnChannelMapping.Store(&openAIWSTurnChannelMappingSnapshot{turn: 1, mapping: channelMappingWS}) - // turn 级定价:BeforeTurn 重新冻结 pricingAt 并按最新门复核当前账号, - // AfterTurn 的计费读取所属 turn 的时刻。零值起步的语义见 - // openAIWSTurnPricing 的注释——绝不能用建连时刻初始化。 + // turn 级定价:BeforeTurn 重新冻结 pricingAt 并按最新门复核当前账号; + // passthrough 没有 BeforeTurn 时,AfterTurn 回退到 TurnStarted 的所属 turn 时刻。 var turnPricing openAIWSTurnPricing hooks := &service.OpenAIWSIngressHooks{ ClientLifecycleContext: clientLifecycleCtx, InitialRequestModel: reqModel, + InitialTurnStartedAt: firstTurnStartedAt, MaxReasoningEffort: maxReasoningEffort, ReasoningEffortMappings: reasoningEffortMappings, + TurnStarted: recordTurnStart, BeforeRequest: func(turn int, payload []byte, originalModel string) error { c.Set(securityAuditWSTurnContextKey, turn) if turn == 1 { @@ -2129,6 +2149,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { return nil }, AfterTurn: func(turn int, result *service.OpenAIForwardResult, turnErr error) { + turnStart := getTurnStart(turn) // F1: cyber 标记按 turn 生命周期清理——defer 保证任意早返回路径都执行; // CyberBlocked 必须在 submit 前同步预捕获(task 闭包由 worker 池异步执行, // 届时 defer 已清除标记)。 @@ -2196,7 +2217,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result) quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey) sessionID := service.ExtractClientSessionID(c) - turnRecordPricingAt := turnPricing.current() + turnRecordPricingAt := turnPricing.currentOr(turnStart) cyberBlocked := service.GetOpsCyberPolicy(c) != nil h.submitOpenAIUsageRecordTask(ctx, result, func(taskCtx context.Context) { if err := h.gatewayService.RecordUsage(taskCtx, &service.OpenAIRecordUsageInput{ diff --git a/backend/internal/handler/openai_ws_turn_pricing_test.go b/backend/internal/handler/openai_ws_turn_pricing_test.go index b9bb204fd6..3a59f323ba 100644 --- a/backend/internal/handler/openai_ws_turn_pricing_test.go +++ b/backend/internal/handler/openai_ws_turn_pricing_test.go @@ -7,16 +7,20 @@ import ( "github.com/stretchr/testify/require" ) -// TestOpenAIWSTurnPricingZeroValue 钉死 WS turn 定价的零值语义: -// 没有 turn 起始回调的 ingress 模式(ws_v2 passthrough 只实现 AfterTurn) -// 必须让 pricingAt 保持零,由 RecordUsage 回退到记录时刻。 -// -// 反例(本 PR 引入的回归):用建连时刻初始化,会把透传连接的所有 turn 钉死在 -// 建连时的高峰因子——客户端峰前一分钟建连并保活,整条连接就按谷价结算。 -func TestOpenAIWSTurnPricingZeroValue(t *testing.T) { - var p openAIWSTurnPricing - require.True(t, p.current().IsZero(), - "未经 turn 起始回调冻结时必须保持零值,交由 RecordUsage 回退记录时刻") +func TestOpenAIWSTurnPricingCurrentOr(t *testing.T) { + fallback := time.Date(2024, time.January, 2, 2, 0, 0, 0, time.UTC) + + t.Run("frozen time takes precedence", func(t *testing.T) { + frozen := fallback.Add(time.Minute) + var p openAIWSTurnPricing + p.freeze(frozen) + require.Equal(t, frozen, p.currentOr(fallback)) + }) + + t.Run("zero value falls back to turn start", func(t *testing.T) { + var p openAIWSTurnPricing + require.Equal(t, fallback, p.currentOr(fallback)) + }) } // TestOpenAIWSTurnPricingFreezePerTurn 钉死每个 turn 的 BeforeTurn 都会覆盖 @@ -27,8 +31,8 @@ func TestOpenAIWSTurnPricingFreezePerTurn(t *testing.T) { turn2 := time.Now() p.freeze(turn1) - require.Equal(t, turn1, p.current()) + require.Equal(t, turn1, p.currentOr(time.Time{})) p.freeze(turn2) - require.Equal(t, turn2, p.current(), "后续 turn 必须使用自己的定价时刻") + require.Equal(t, turn2, p.currentOr(time.Time{}), "后续 turn 必须使用自己的定价时刻") } diff --git a/backend/internal/repository/channel_repo_pricing.go b/backend/internal/repository/channel_repo_pricing.go index 995621139e..b7a9fd6b02 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, created_at, updated_at + `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 FROM channel_model_pricing WHERE channel_id = $1 ORDER BY id`, channelID, ) if err != nil { @@ -51,16 +51,20 @@ func (r *channelRepository) UpdateModelPricing(ctx context.Context, pricing *ser if err != nil { return fmt.Errorf("marshal models: %w", err) } + timePricingJSON, err := marshalChannelTimePricing(pricing.TimePricing) + if err != nil { + return err + } billingMode := pricing.BillingMode if billingMode == "" { billingMode = service.BillingModeToken } 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, platform = $10, updated_at = NOW() - WHERE id = $11`, + 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`, modelsJSON, billingMode, pricing.InputPrice, pricing.OutputPrice, pricing.CacheWritePrice, pricing.CacheReadPrice, - pricing.ImageInputPrice, pricing.ImageOutputPrice, pricing.PerRequestPrice, pricing.Platform, pricing.ID, + pricing.ImageInputPrice, pricing.ImageOutputPrice, pricing.PerRequestPrice, timePricingJSON, pricing.Platform, pricing.ID, ) if err != nil { return fmt.Errorf("update model pricing: %w", err) @@ -91,7 +95,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, created_at, updated_at + `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 FROM channel_model_pricing WHERE channel_id = ANY($1) ORDER BY channel_id, id`, pq.Array(channelIDs), ) @@ -169,16 +173,22 @@ func scanModelPricingRows(rows *sql.Rows) ([]service.ChannelModelPricing, []int6 for rows.Next() { var p service.ChannelModelPricing var modelsJSON []byte + var timePricingJSON []byte if err := rows.Scan( &p.ID, &p.ChannelID, &p.Platform, &modelsJSON, &p.BillingMode, &p.InputPrice, &p.OutputPrice, &p.CacheWritePrice, &p.CacheReadPrice, - &p.ImageInputPrice, &p.ImageOutputPrice, &p.PerRequestPrice, &p.CreatedAt, &p.UpdatedAt, + &p.ImageInputPrice, &p.ImageOutputPrice, &p.PerRequestPrice, &timePricingJSON, &p.CreatedAt, &p.UpdatedAt, ); err != nil { return nil, nil, fmt.Errorf("scan model pricing: %w", err) } if err := json.Unmarshal(modelsJSON, &p.Models); err != nil { p.Models = []string{} } + timePricing, err := unmarshalChannelTimePricing(timePricingJSON) + if err != nil { + return nil, nil, err + } + p.TimePricing = timePricing pricingIDs = append(pricingIDs, p.ID) result = append(result, p) } @@ -220,6 +230,10 @@ func createModelPricingExec(ctx context.Context, exec dbExec, pricing *service.C if err != nil { return fmt.Errorf("marshal models: %w", err) } + timePricingJSON, err := marshalChannelTimePricing(pricing.TimePricing) + if err != nil { + return err + } billingMode := pricing.BillingMode if billingMode == "" { billingMode = service.BillingModeToken @@ -229,11 +243,11 @@ 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) - VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11) 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, 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`, pricing.ChannelID, platform, modelsJSON, billingMode, pricing.InputPrice, pricing.OutputPrice, pricing.CacheWritePrice, pricing.CacheReadPrice, - pricing.ImageInputPrice, pricing.ImageOutputPrice, pricing.PerRequestPrice, + 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) @@ -249,6 +263,28 @@ func createModelPricingExec(ctx context.Context, exec dbExec, pricing *service.C return nil } +func marshalChannelTimePricing(config *service.ChannelTimePricing) (any, error) { + if config == nil || len(config.Periods) == 0 { + return nil, nil + } + data, err := json.Marshal(config) + if err != nil { + return nil, fmt.Errorf("marshal time pricing: %w", err) + } + return string(data), nil +} + +func unmarshalChannelTimePricing(data []byte) (*service.ChannelTimePricing, error) { + if len(data) == 0 { + return nil, nil + } + var config service.ChannelTimePricing + if err := json.Unmarshal(data, &config); err != nil { + return nil, fmt.Errorf("unmarshal time pricing: %w", err) + } + return &config, nil +} + func createIntervalExec(ctx context.Context, exec dbExec, iv *service.PricingInterval) error { return exec.QueryRowContext(ctx, `INSERT INTO channel_pricing_intervals diff --git a/backend/internal/repository/channel_repo_pricing_time_test.go b/backend/internal/repository/channel_repo_pricing_time_test.go new file mode 100644 index 0000000000..d7af1886e7 --- /dev/null +++ b/backend/internal/repository/channel_repo_pricing_time_test.go @@ -0,0 +1,181 @@ +//go:build unit + +package repository + +import ( + "context" + "regexp" + "strings" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + "github.com/Wei-Shaw/sub2api/internal/service" + "github.com/stretchr/testify/require" +) + +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", + "per_request_price", "time_pricing", "created_at", "updated_at", +} + +const channelModelPricingTimePricingJSON = `{"timezone":"Asia/Shanghai","periods":[{"start_time":"09:00","end_time":"12:00","multiplier":2}]}` + +func newChannelModelPricingTimePricingRepo(t *testing.T) (*channelRepository, sqlmock.Sqlmock) { + t.Helper() + db, mock, err := sqlmock.New() + require.NoError(t, err) + t.Cleanup(func() { _ = db.Close() }) + return &channelRepository{db: db}, mock +} + +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, + time.Date(2026, 8, 17, 0, 0, 0, 0, time.UTC), time.Date(2026, 8, 17, 1, 0, 0, 0, time.UTC), + ) +} + +func expectEmptyModelPricingIntervals(mock sqlmock.Sqlmock) { + mock.ExpectQuery(`SELECT id, pricing_id, min_tokens, max_tokens, tier_label`). + WithArgs(sqlmock.AnyArg()). + WillReturnRows(sqlmock.NewRows([]string{"id"})) +} + +func TestChannelModelPricingTimePricingListRoundTrip(t *testing.T) { + repo, mock := newChannelModelPricingTimePricingRepo(t) + mock.ExpectQuery(`(?s)SELECT .*per_request_price, time_pricing, created_at, updated_at.*FROM channel_model_pricing.*channel_id = \$1`). + WithArgs(int64(7)). + WillReturnRows(modelPricingTimePricingRow(channelModelPricingTimePricingJSON)) + expectEmptyModelPricingIntervals(mock) + + pricing, err := repo.ListModelPricing(context.Background(), 7) + require.NoError(t, err) + require.Len(t, pricing, 1) + require.NotNil(t, pricing[0].TimePricing) + require.Equal(t, "Asia/Shanghai", pricing[0].TimePricing.Timezone) + require.Len(t, pricing[0].TimePricing.Periods, 1) + require.Equal(t, 2.0, pricing[0].TimePricing.Periods[0].Multiplier) + require.NoError(t, mock.ExpectationsWereMet()) +} + +func TestChannelModelPricingTimePricingListNullAndMalformed(t *testing.T) { + t.Run("SQL NULL maps to nil", func(t *testing.T) { + repo, mock := newChannelModelPricingTimePricingRepo(t) + mock.ExpectQuery(`(?s)SELECT .*per_request_price, time_pricing, created_at, updated_at.*FROM channel_model_pricing.*channel_id = \$1`). + WithArgs(int64(7)). + WillReturnRows(modelPricingTimePricingRow(nil)) + expectEmptyModelPricingIntervals(mock) + + pricing, err := repo.ListModelPricing(context.Background(), 7) + require.NoError(t, err) + require.Len(t, pricing, 1) + require.Nil(t, pricing[0].TimePricing) + require.NoError(t, mock.ExpectationsWereMet()) + }) + + t.Run("malformed JSON returns repository error", func(t *testing.T) { + repo, mock := newChannelModelPricingTimePricingRepo(t) + mock.ExpectQuery(`(?s)SELECT .*per_request_price, time_pricing, created_at, updated_at.*FROM channel_model_pricing.*channel_id = \$1`). + WithArgs(int64(7)). + WillReturnRows(modelPricingTimePricingRow(`{"timezone":`)) + + _, err := repo.ListModelPricing(context.Background(), 7) + require.Error(t, err) + require.True(t, strings.Contains(err.Error(), "unmarshal time pricing"), "unexpected error: %v", err) + require.NoError(t, mock.ExpectationsWereMet()) + }) +} + +func TestChannelModelPricingTimePricingCreateAndUpdateRoundTrip(t *testing.T) { + pricing := &service.ChannelModelPricing{ + ID: 11, + ChannelID: 7, + Platform: "openai", + Models: []string{"gpt-5"}, + TimePricing: &service.ChannelTimePricing{ + Timezone: "Asia/Shanghai", + Periods: []service.ChannelTimePricingPeriod{{ + StartTime: "09:00", EndTime: "12:00", Multiplier: 2, + }}, + }, + } + + 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)")). + WithArgs( + int64(7), "openai", []byte(`["gpt-5"]`), service.BillingModeToken, + nil, nil, nil, nil, nil, nil, nil, channelModelPricingTimePricingJSON, + ). + WillReturnRows(sqlmock.NewRows([]string{"id", "created_at", "updated_at"}).AddRow(int64(11), time.Time{}, time.Time{})) + + require.NoError(t, repo.CreateModelPricing(context.Background(), pricing)) + require.NoError(t, mock.ExpectationsWereMet()) + }) + + 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`). + WithArgs( + []byte(`["gpt-5"]`), service.BillingModeToken, + nil, nil, nil, nil, nil, nil, nil, channelModelPricingTimePricingJSON, "openai", int64(11), + ). + WillReturnResult(sqlmock.NewResult(0, 1)) + + require.NoError(t, repo.UpdateModelPricing(context.Background(), pricing)) + require.NoError(t, mock.ExpectationsWereMet()) + }) +} + +func TestChannelModelPricingTimePricingCreateAndUpdateWriteNullWhenDisabled(t *testing.T) { + tests := []struct { + name string + timePricing *service.ChannelTimePricing + }{ + {name: "nil", timePricing: nil}, + {name: "empty periods", timePricing: &service.ChannelTimePricing{Timezone: "Asia/Shanghai"}}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + newPricing := func() *service.ChannelModelPricing { + return &service.ChannelModelPricing{ + ID: 11, + ChannelID: 7, + Platform: "openai", + Models: []string{"gpt-5"}, + TimePricing: tt.timePricing, + } + } + + 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)")). + WithArgs( + int64(7), "openai", []byte(`["gpt-5"]`), service.BillingModeToken, + nil, nil, nil, nil, nil, nil, nil, nil, + ). + WillReturnRows(sqlmock.NewRows([]string{"id", "created_at", "updated_at"}).AddRow(int64(11), time.Time{}, time.Time{})) + + require.NoError(t, repo.CreateModelPricing(context.Background(), newPricing())) + require.NoError(t, mock.ExpectationsWereMet()) + }) + + 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`). + WithArgs( + []byte(`["gpt-5"]`), service.BillingModeToken, + nil, nil, nil, nil, nil, nil, nil, nil, "openai", int64(11), + ). + WillReturnResult(sqlmock.NewResult(0, 1)) + + require.NoError(t, repo.UpdateModelPricing(context.Background(), newPricing())) + require.NoError(t, mock.ExpectationsWereMet()) + }) + }) + } +} diff --git a/backend/internal/service/admin_group.go b/backend/internal/service/admin_group.go index ac6259509d..f8f9acb4e7 100644 --- a/backend/internal/service/admin_group.go +++ b/backend/internal/service/admin_group.go @@ -987,6 +987,12 @@ func normalizeGroupModelPricing(platform string, pricing []ChannelModelPricing) out[i] = pricing[i].Clone() out[i].ID = 0 out[i].ChannelID = 0 + if out[i].TimePricing != nil && len(out[i].TimePricing.Periods) > 0 { + return nil, infraerrors.BadRequest( + "GROUP_MODEL_TIME_PRICING_UNSUPPORTED", + "group model pricing does not support time pricing", + ) + } if strings.TrimSpace(out[i].Platform) == "" { out[i].Platform = platform } diff --git a/backend/internal/service/admin_service_group_test.go b/backend/internal/service/admin_service_group_test.go index c0b4faef52..64bfdaeb02 100644 --- a/backend/internal/service/admin_service_group_test.go +++ b/backend/internal/service/admin_service_group_test.go @@ -4,8 +4,10 @@ package service import ( "context" + "net/http" "testing" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" "github.com/stretchr/testify/require" ) @@ -156,6 +158,61 @@ func (s *groupRepoStubForAdmin) UpdateSortOrders(_ context.Context, _ []GroupSor return nil } +func TestAdminService_CreateGroup_RejectsTimePricing(t *testing.T) { + repo := &groupRepoStubForAdmin{createID: 51} + svc := &adminServiceImpl{groupRepo: repo} + + _, err := svc.CreateGroup(context.Background(), &CreateGroupInput{ + Name: "time-pricing-group", + Platform: PlatformOpenAI, + RateMultiplier: 1, + ModelPricing: []ChannelModelPricing{{ + Platform: PlatformOpenAI, + Models: []string{"gpt-5"}, + BillingMode: BillingModeToken, + TimePricing: validTimePricingForTest(), + }}, + }) + + require.Error(t, err) + appErr := infraerrors.FromError(err) + require.Equal(t, int32(http.StatusBadRequest), appErr.Code) + require.Equal(t, "GROUP_MODEL_TIME_PRICING_UNSUPPORTED", appErr.Reason) + require.Nil(t, repo.created) +} + +func TestAdminService_UpdateGroup_RejectsTimePricing(t *testing.T) { + existing := &Group{ID: 1, Name: "existing", Platform: PlatformOpenAI, Status: StatusActive} + repo := &groupRepoStubForAdmin{getByID: existing} + svc := &adminServiceImpl{groupRepo: repo} + pricing := []ChannelModelPricing{{ + Platform: PlatformOpenAI, + Models: []string{"gpt-5"}, + BillingMode: BillingModeToken, + TimePricing: validTimePricingForTest(), + }} + + _, err := svc.UpdateGroup(context.Background(), existing.ID, &UpdateGroupInput{ModelPricing: &pricing}) + + require.Error(t, err) + appErr := infraerrors.FromError(err) + require.Equal(t, int32(http.StatusBadRequest), appErr.Code) + require.Equal(t, "GROUP_MODEL_TIME_PRICING_UNSUPPORTED", appErr.Reason) + require.Nil(t, repo.updated) +} + +func TestNormalizeGroupModelPricing_NormalizesEmptyTimePricing(t *testing.T) { + pricing, err := normalizeGroupModelPricing(PlatformOpenAI, []ChannelModelPricing{{ + Models: []string{"gpt-5"}, + BillingMode: BillingModeToken, + TimePricing: &ChannelTimePricing{Timezone: "Asia/Shanghai"}, + }}) + + require.NoError(t, err) + require.Len(t, pricing, 1) + require.Nil(t, pricing[0].TimePricing) +} + type compositeRouteRepoStubForAdmin struct { routes []CompositeModelRoute created *CompositeModelRoute diff --git a/backend/internal/service/billing_service.go b/backend/internal/service/billing_service.go index 0356f698a0..6ead8da021 100644 --- a/backend/internal/service/billing_service.go +++ b/backend/internal/service/billing_service.go @@ -167,6 +167,27 @@ type CostBreakdown struct { LongContextBillingApplied bool } +func applyCostBreakdownMultiplier(cost *CostBreakdown, multiplier float64) { + if cost == nil || multiplier == 1 { + return + } + cost.InputCost *= multiplier + cost.ImageInputCost *= multiplier + cost.OutputCost *= multiplier + cost.ImageOutputCost *= multiplier + cost.CacheCreationCost *= multiplier + cost.CacheReadCost *= multiplier + cost.TotalCost *= multiplier + cost.ActualCost *= multiplier +} + +func resolvedChannelTimeMultiplier(resolved *ResolvedPricing, at time.Time) float64 { + if resolved == nil || resolved.Source != PricingSourceChannel || resolved.channelPricing == nil { + return 1 + } + return resolved.channelPricing.TimePricing.MultiplierAt(at) +} + // ErrModelPricingUnavailable indicates that none of the configured pricing // sources can price the requested model. var ErrModelPricingUnavailable = errors.New("pricing not found") @@ -1030,6 +1051,7 @@ type CostInput struct { UsageUnits float64 // 音频等连续计量单位(分钟/小时/百万字符) SizeTier string // 按次/图片模式的层级标签("1K","2K","4K","HD" 等) RateMultiplier float64 + PricingAt time.Time // 渠道分时定价使用的计费时刻 ServiceTier string // "priority","flex","" 等 Resolver *ModelPricingResolver // 定价解析器 Resolved *ResolvedPricing // 可选:预解析的定价结果(避免重复 Resolve 调用) @@ -1104,7 +1126,9 @@ func (s *BillingService) calculateTokenCost(resolved *ResolvedPricing, input Cos applyLongCtx = applyLongCtx && *input.LongContextBillingEnabled } - return s.computeTokenBreakdown(pricing, input.Tokens, input.RateMultiplier, input.ServiceTier, applyLongCtx), nil + breakdown := s.computeTokenBreakdown(pricing, input.Tokens, input.RateMultiplier, input.ServiceTier, applyLongCtx) + applyCostBreakdownMultiplier(breakdown, resolvedChannelTimeMultiplier(resolved, input.PricingAt)) + return breakdown, nil } // computeTokenBreakdown 是 token 计费的核心逻辑,由 calculateTokenCost 和 calculateCostInternal 共用。 diff --git a/backend/internal/service/billing_service_unified_test.go b/backend/internal/service/billing_service_unified_test.go index eabbab3dfc..e61a899551 100644 --- a/backend/internal/service/billing_service_unified_test.go +++ b/backend/internal/service/billing_service_unified_test.go @@ -169,6 +169,154 @@ func TestCalculateCostUnified_ImageMode(t *testing.T) { require.Equal(t, string(BillingModeImage), cost.BillingMode) } +func channelTimeResolvedForTest(base *ModelPricing, intervals []PricingInterval) *ResolvedPricing { + return &ResolvedPricing{ + Mode: BillingModeToken, + BasePricing: base, + Intervals: intervals, + Source: PricingSourceChannel, + channelPricing: &ChannelModelPricing{ + BillingMode: BillingModeToken, + TimePricing: &ChannelTimePricing{ + Timezone: "Asia/Shanghai", + Periods: []ChannelTimePricingPeriod{{ + StartTime: "09:00", + EndTime: "12:00", + Multiplier: 2, + }}, + }, + }, + longContextPricingEnabled: true, + } +} + +func TestCalculateCostUnified_ChannelTimePricingScalesBaseAndActualCost(t *testing.T) { + billing := NewBillingService(&config.Config{}, nil) + resolved := channelTimeResolvedForTest(&ModelPricing{InputPricePerToken: 0.001}, nil) + + cost, err := billing.CalculateCostUnified(CostInput{ + Ctx: context.Background(), + Model: "model", + Tokens: UsageTokens{InputTokens: 1000}, + RateMultiplier: 0.8, + Resolver: &ModelPricingResolver{}, + Resolved: resolved, + PricingAt: time.Date(2026, 8, 17, 1, 0, 0, 0, time.UTC), + }) + require.NoError(t, err) + require.InDelta(t, 2.0, cost.InputCost, 1e-12) + require.InDelta(t, 2.0, cost.TotalCost, 1e-12) + require.InDelta(t, 1.6, cost.ActualCost, 1e-12) +} + +func TestCalculateCostUnified_ChannelTimePricingScalesMatchingInterval(t *testing.T) { + intervalInputPrice := 0.003 + resolved := channelTimeResolvedForTest( + &ModelPricing{InputPricePerToken: 0.001}, + []PricingInterval{{MinTokens: 0, InputPrice: &intervalInputPrice}}, + ) + billing := NewBillingService(&config.Config{}, nil) + + cost, err := billing.CalculateCostUnified(CostInput{ + Ctx: context.Background(), + Model: "model", + Tokens: UsageTokens{InputTokens: 1000}, + Resolver: &ModelPricingResolver{}, + Resolved: resolved, + PricingAt: time.Date(2026, 8, 17, 1, 0, 0, 0, time.UTC), + }) + require.NoError(t, err) + require.InDelta(t, 6.0, cost.InputCost, 1e-12) + require.InDelta(t, 6.0, cost.TotalCost, 1e-12) +} + +func TestCalculateCostUnified_ChannelTimePricingScalesBaseOnUnmatchedInterval(t *testing.T) { + intervalInputPrice := 0.003 + resolved := channelTimeResolvedForTest( + &ModelPricing{InputPricePerToken: 0.001}, + []PricingInterval{{MinTokens: 2000, InputPrice: &intervalInputPrice}}, + ) + billing := NewBillingService(&config.Config{}, nil) + + cost, err := billing.CalculateCostUnified(CostInput{ + Ctx: context.Background(), + Model: "model", + Tokens: UsageTokens{InputTokens: 1000}, + Resolver: &ModelPricingResolver{}, + Resolved: resolved, + PricingAt: time.Date(2026, 8, 17, 1, 0, 0, 0, time.UTC), + }) + require.NoError(t, err) + require.InDelta(t, 2.0, cost.InputCost, 1e-12) + require.InDelta(t, 2.0, cost.TotalCost, 1e-12) +} + +func TestCalculateCostUnified_ChannelTimePricingDoesNotApplyToGroupPricing(t *testing.T) { + resolved := channelTimeResolvedForTest(&ModelPricing{InputPricePerToken: 0.001}, nil) + resolved.Source = PricingSourceGroup + billing := NewBillingService(&config.Config{}, nil) + + cost, err := billing.CalculateCostUnified(CostInput{ + Ctx: context.Background(), + Model: "model", + Tokens: UsageTokens{InputTokens: 1000}, + Resolver: &ModelPricingResolver{}, + Resolved: resolved, + PricingAt: time.Date(2026, 8, 17, 1, 0, 0, 0, time.UTC), + }) + require.NoError(t, err) + require.InDelta(t, 1.0, cost.TotalCost, 1e-12) +} + +func TestCalculateCostUnified_ChannelTimePricingDoesNotApplyOutsideMatchingTime(t *testing.T) { + resolved := channelTimeResolvedForTest(&ModelPricing{InputPricePerToken: 0.001}, nil) + billing := NewBillingService(&config.Config{}, nil) + + for _, pricingAt := range []time.Time{ + time.Time{}, + time.Date(2026, 8, 17, 5, 0, 0, 0, time.UTC), + } { + cost, err := billing.CalculateCostUnified(CostInput{ + Ctx: context.Background(), + Model: "model", + Tokens: UsageTokens{InputTokens: 1000}, + Resolver: &ModelPricingResolver{}, + Resolved: resolved, + PricingAt: pricingAt, + }) + require.NoError(t, err) + require.InDelta(t, 1.0, cost.TotalCost, 1e-12) + } +} + +func TestApplyCostBreakdownMultiplierScalesAllMonetaryFields(t *testing.T) { + cost := &CostBreakdown{ + InputCost: 1, + ImageInputCost: 2, + OutputCost: 3, + ImageOutputCost: 4, + CacheCreationCost: 5, + CacheReadCost: 6, + TotalCost: 21, + ActualCost: 42, + BillingMode: string(BillingModeToken), + LongContextBillingApplied: true, + } + + applyCostBreakdownMultiplier(cost, 1.5) + + require.InDelta(t, 1.5, cost.InputCost, 1e-12) + require.InDelta(t, 3.0, cost.ImageInputCost, 1e-12) + require.InDelta(t, 4.5, cost.OutputCost, 1e-12) + require.InDelta(t, 6.0, cost.ImageOutputCost, 1e-12) + require.InDelta(t, 7.5, cost.CacheCreationCost, 1e-12) + require.InDelta(t, 9.0, cost.CacheReadCost, 1e-12) + require.InDelta(t, 31.5, cost.TotalCost, 1e-12) + require.InDelta(t, 63.0, cost.ActualCost, 1e-12) + require.Equal(t, string(BillingModeToken), cost.BillingMode) + require.True(t, cost.LongContextBillingApplied) +} + // TestCalculateCostUnified_RateMultiplierZeroProducesZero 锁定新行为: // 保存时强制 > 0;若 0 仍泄漏到计费层,按 0 计费(而非历史上的 1.0)。 func TestCalculateCostUnified_RateMultiplierZeroProducesZero(t *testing.T) { diff --git a/backend/internal/service/channel.go b/backend/internal/service/channel.go index 4d5523b0a3..5345693eb0 100644 --- a/backend/internal/service/channel.go +++ b/backend/internal/service/channel.go @@ -87,21 +87,35 @@ type AccountStatsPricingRule struct { // ChannelModelPricing 渠道模型定价条目 type ChannelModelPricing struct { - ID int64 `json:"id,omitempty"` - ChannelID int64 `json:"channel_id,omitempty"` - Platform string `json:"platform"` // 所属平台(anthropic/openai/gemini/...) - Models []string `json:"models"` - BillingMode BillingMode `json:"billing_mode"` - InputPrice *float64 `json:"input_price"` - OutputPrice *float64 `json:"output_price"` - CacheWritePrice *float64 `json:"cache_write_price"` - CacheReadPrice *float64 `json:"cache_read_price"` - ImageInputPrice *float64 `json:"image_input_price"` - ImageOutputPrice *float64 `json:"image_output_price"` - PerRequestPrice *float64 `json:"per_request_price"` - Intervals []PricingInterval `json:"intervals"` - CreatedAt time.Time `json:"created_at,omitempty"` - UpdatedAt time.Time `json:"updated_at,omitempty"` + ID int64 `json:"id,omitempty"` + ChannelID int64 `json:"channel_id,omitempty"` + Platform string `json:"platform"` // 所属平台(anthropic/openai/gemini/...) + Models []string `json:"models"` + BillingMode BillingMode `json:"billing_mode"` + InputPrice *float64 `json:"input_price"` + OutputPrice *float64 `json:"output_price"` + CacheWritePrice *float64 `json:"cache_write_price"` + CacheReadPrice *float64 `json:"cache_read_price"` + ImageInputPrice *float64 `json:"image_input_price"` + ImageOutputPrice *float64 `json:"image_output_price"` + PerRequestPrice *float64 `json:"per_request_price"` + Intervals []PricingInterval `json:"intervals"` + TimePricing *ChannelTimePricing `json:"time_pricing,omitempty"` + CreatedAt time.Time `json:"created_at,omitempty"` + UpdatedAt time.Time `json:"updated_at,omitempty"` +} + +// ChannelTimePricing 渠道模型定价的分时倍率配置。 +type ChannelTimePricing struct { + Timezone string `json:"timezone"` + Periods []ChannelTimePricingPeriod `json:"periods"` +} + +// ChannelTimePricingPeriod 是秒级的左闭右开分时倍率区间,并兼容历史 HH:mm 数据。 +type ChannelTimePricingPeriod struct { + StartTime string `json:"start_time"` + EndTime string `json:"end_time"` + Multiplier float64 `json:"multiplier"` } // PricingInterval 定价区间(token 区间 / 按次分层 / 图片分辨率分层) @@ -195,6 +209,12 @@ func (p ChannelModelPricing) Clone() ChannelModelPricing { cp.Intervals = make([]PricingInterval, len(p.Intervals)) copy(cp.Intervals, p.Intervals) } + if p.TimePricing != nil { + cp.TimePricing = &ChannelTimePricing{Timezone: p.TimePricing.Timezone} + if p.TimePricing.Periods != nil { + cp.TimePricing.Periods = append([]ChannelTimePricingPeriod(nil), p.TimePricing.Periods...) + } + } return cp } diff --git a/backend/internal/service/channel_service.go b/backend/internal/service/channel_service.go index 3a8c5556ad..a2c4bf6ebf 100644 --- a/backend/internal/service/channel_service.go +++ b/backend/internal/service/channel_service.go @@ -646,7 +646,47 @@ func validatePricingEntries(pricing []ChannelModelPricing) error { if err := validatePricingIntervals(pricing); err != nil { return err } - return validatePricingBillingMode(pricing) + if err := validatePricingBillingMode(pricing); err != nil { + return err + } + return validatePricingTimePricing(pricing) +} + +func validatePricingTimePricing(pricing []ChannelModelPricing) error { + for i := range pricing { + config := pricing[i].TimePricing + if config == nil { + continue + } + if len(config.Periods) == 0 { + pricing[i].TimePricing = nil + continue + } + mode := pricing[i].BillingMode + if mode != "" && mode != BillingModeToken { + return infraerrors.BadRequest("TIME_PRICING_UNSUPPORTED_MODE", "time pricing only supports token billing mode") + } + if err := validateChannelTimePricing(config); err != nil { + return infraerrors.BadRequest("INVALID_TIME_PRICING", fmt.Sprintf( + "invalid time pricing for platform '%s' models %v: %v", pricing[i].Platform, pricing[i].Models, err)) + } + } + return nil +} + +func validateAccountStatsPricingRules(rules []AccountStatsPricingRule) error { + for i := range rules { + for _, pricing := range rules[i].Pricing { + if pricing.TimePricing != nil && len(pricing.TimePricing.Periods) > 0 { + return fmt.Errorf("account stats pricing rule #%d: %w", i+1, + infraerrors.BadRequest("ACCOUNT_STATS_TIME_PRICING_UNSUPPORTED", "account stats pricing does not support time pricing")) + } + } + if err := validatePricingEntries(rules[i].Pricing); err != nil { + return fmt.Errorf("account stats pricing rule #%d: %w", i+1, err) + } + } + return nil } // validatePricingBillingMode 校验计费模式配置:按次/图片模式必须配价格或区间,所有价格字段不能为负,区间至少有一个价格字段。 @@ -755,10 +795,8 @@ func (s *ChannelService) Create(ctx context.Context, input *CreateChannelInput) if err := validateChannelConfig(channel.ModelPricing, channel.ModelMapping); err != nil { return nil, err } - for i, rule := range channel.AccountStatsPricingRules { - if err := validatePricingEntries(rule.Pricing); err != nil { - return nil, fmt.Errorf("account stats pricing rule #%d: %w", i+1, err) - } + if err := validateAccountStatsPricingRules(channel.AccountStatsPricingRules); err != nil { + return nil, err } if err := s.repo.Create(ctx, channel); err != nil { @@ -799,10 +837,8 @@ func (s *ChannelService) Update(ctx context.Context, id int64, input *UpdateChan if err := validateChannelConfig(channel.ModelPricing, channel.ModelMapping); err != nil { return nil, err } - for i, rule := range channel.AccountStatsPricingRules { - if err := validatePricingEntries(rule.Pricing); err != nil { - return nil, fmt.Errorf("account stats pricing rule #%d: %w", i+1, err) - } + if err := validateAccountStatsPricingRules(channel.AccountStatsPricingRules); err != nil { + return nil, err } oldGroupIDs := s.getOldGroupIDs(ctx, id) diff --git a/backend/internal/service/channel_service_test.go b/backend/internal/service/channel_service_test.go index f56f61e917..9a81913b04 100644 --- a/backend/internal/service/channel_service_test.go +++ b/backend/internal/service/channel_service_test.go @@ -5,8 +5,10 @@ package service import ( "context" "errors" + "net/http" "testing" + infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/pagination" "github.com/stretchr/testify/require" "github.com/tidwall/gjson" @@ -2460,6 +2462,66 @@ func TestValidatePricingBillingMode(t *testing.T) { } } +func validTimePricingForTest() *ChannelTimePricing { + return &ChannelTimePricing{Timezone: "Asia/Shanghai", Periods: []ChannelTimePricingPeriod{ + {StartTime: "09:00", EndTime: "12:00", Multiplier: 2}, + }} +} + +func TestValidatePricingTimePricing(t *testing.T) { + token := []ChannelModelPricing{{BillingMode: BillingModeToken, TimePricing: validTimePricingForTest()}} + require.NoError(t, validatePricingTimePricing(token)) + + implicitToken := []ChannelModelPricing{{TimePricing: validTimePricingForTest()}} + require.NoError(t, validatePricingTimePricing(implicitToken)) + + image := []ChannelModelPricing{{BillingMode: BillingModeImage, TimePricing: validTimePricingForTest()}} + modeErr := infraerrors.FromError(validatePricingTimePricing(image)) + require.Equal(t, int32(http.StatusBadRequest), modeErr.Code) + require.Equal(t, "TIME_PRICING_UNSUPPORTED_MODE", modeErr.Reason) + + invalid := []ChannelModelPricing{{ + Platform: PlatformOpenAI, + Models: []string{"gpt-5"}, + BillingMode: BillingModeToken, + TimePricing: &ChannelTimePricing{Timezone: "UTC+8", Periods: validTimePricingForTest().Periods}, + }} + invalidErr := infraerrors.FromError(validatePricingTimePricing(invalid)) + require.Equal(t, int32(http.StatusBadRequest), invalidErr.Code) + require.Equal(t, "INVALID_TIME_PRICING", invalidErr.Reason) + require.Contains(t, invalidErr.Message, "platform 'openai'") + require.Contains(t, invalidErr.Message, "models [gpt-5]") + + invalidMultiplier := []ChannelModelPricing{{ + Platform: PlatformOpenAI, + Models: []string{"gpt-5"}, + BillingMode: BillingModeToken, + TimePricing: &ChannelTimePricing{Timezone: "Asia/Shanghai", Periods: []ChannelTimePricingPeriod{{ + StartTime: "09:00", EndTime: "12:00", Multiplier: 1e-12, + }}}, + }} + invalidMultiplierRawErr := validatePricingTimePricing(invalidMultiplier) + require.Error(t, invalidMultiplierRawErr) + invalidMultiplierErr := infraerrors.FromError(invalidMultiplierRawErr) + require.Equal(t, int32(http.StatusBadRequest), invalidMultiplierErr.Code) + require.Equal(t, "INVALID_TIME_PRICING", invalidMultiplierErr.Reason) + + empty := []ChannelModelPricing{{BillingMode: BillingModeToken, TimePricing: &ChannelTimePricing{Timezone: "Asia/Shanghai"}}} + require.NoError(t, validatePricingTimePricing(empty)) + require.Nil(t, empty[0].TimePricing) +} + +func TestValidateAccountStatsPricingRulesRejectsTimePricing(t *testing.T) { + rules := []AccountStatsPricingRule{{Pricing: []ChannelModelPricing{{ + BillingMode: BillingModeToken, + TimePricing: validTimePricingForTest(), + }}}} + + appErr := infraerrors.FromError(validateAccountStatsPricingRules(rules)) + require.Equal(t, int32(http.StatusBadRequest), appErr.Code) + require.Equal(t, "ACCOUNT_STATS_TIME_PRICING_UNSUPPORTED", appErr.Reason) +} + // --------------------------------------------------------------------------- // 12. Antigravity wildcard mapping isolation // --------------------------------------------------------------------------- diff --git a/backend/internal/service/channel_test.go b/backend/internal/service/channel_test.go index 2f371f8a1c..19e45a02db 100644 --- a/backend/internal/service/channel_test.go +++ b/backend/internal/service/channel_test.go @@ -197,6 +197,14 @@ func TestChannelModelPricingClone(t *testing.T) { Intervals: []PricingInterval{ {MinTokens: 0, TierLabel: "tier1"}, }, + TimePricing: &ChannelTimePricing{ + Timezone: "Asia/Shanghai", + Periods: []ChannelTimePricingPeriod{{ + StartTime: "09:00", + EndTime: "12:00", + Multiplier: 2, + }}, + }, } cloned := original.Clone() @@ -207,6 +215,13 @@ func TestChannelModelPricingClone(t *testing.T) { cloned.Intervals[0].TierLabel = "hacked" require.Equal(t, "tier1", original.Intervals[0].TierLabel) + + cloned.TimePricing.Timezone = "America/New_York" + cloned.TimePricing.Periods[0].StartTime = "10:00" + cloned.TimePricing.Periods[0].Multiplier = 3 + require.Equal(t, "Asia/Shanghai", original.TimePricing.Timezone) + require.Equal(t, "09:00", original.TimePricing.Periods[0].StartTime) + require.Equal(t, 2.0, original.TimePricing.Periods[0].Multiplier) } // --- BillingMode.IsValid --- @@ -513,7 +528,6 @@ func TestSupportedModels_WildcardExpandedFromPricing(t *testing.T) { } } - func TestSupportedModels_MissingPricingKeepsNilPricing(t *testing.T) { ch := &Channel{ ModelMapping: map[string]map[string]string{ diff --git a/backend/internal/service/custom_channel_time_pricing.go b/backend/internal/service/custom_channel_time_pricing.go new file mode 100644 index 0000000000..759eac987e --- /dev/null +++ b/backend/internal/service/custom_channel_time_pricing.go @@ -0,0 +1,140 @@ +package service + +import ( + "fmt" + "math" + "sort" + "strings" + "sync" + "time" +) + +var channelTimePricingLocations sync.Map + +type parsedChannelTimePeriod struct { + start int + end int + multiplier float64 +} + +// validateChannelTimePricing 校验分时倍率配置。nil 或空 periods 表示未启用。 +func validateChannelTimePricing(config *ChannelTimePricing) error { + if config == nil || len(config.Periods) == 0 { + return nil + } + if _, err := loadChannelTimePricingLocation(config.Timezone); err != nil { + return fmt.Errorf("timezone: %w", err) + } + _, err := parseChannelTimePeriods(config.Periods) + return err +} + +// MultiplierAt 返回 at 对应的分时倍率。无配置或脏配置均安全降级为 1。 +func (config *ChannelTimePricing) MultiplierAt(at time.Time) float64 { + if config == nil || len(config.Periods) == 0 || at.IsZero() { + return 1.0 + } + if err := validateChannelTimePricing(config); err != nil { + return 1.0 + } + location, err := loadChannelTimePricingLocation(config.Timezone) + if err != nil { + return 1.0 + } + periods, err := parseChannelTimePeriods(config.Periods) + if err != nil { + return 1.0 + } + + local := at.In(location) + second := local.Hour()*60*60 + local.Minute()*60 + local.Second() + for _, period := range periods { + if second >= period.start && second < period.end { + return period.multiplier + } + } + return 1.0 +} + +func loadChannelTimePricingLocation(name string) (*time.Location, error) { + if strings.TrimSpace(name) == "" { + return nil, fmt.Errorf("timezone is required") + } + if name == "Local" { + return nil, fmt.Errorf("local is not a supported timezone") + } + if cached, ok := channelTimePricingLocations.Load(name); ok { + location, valid := cached.(*time.Location) + if valid && location != nil { + return location, nil + } + channelTimePricingLocations.Delete(name) + } + location, err := time.LoadLocation(name) + if err != nil { + return nil, err + } + actual, _ := channelTimePricingLocations.LoadOrStore(name, location) + actualLocation, ok := actual.(*time.Location) + if !ok || actualLocation == nil { + return nil, fmt.Errorf("invalid cached timezone %q", name) + } + return actualLocation, nil +} + +func parseChannelTime(value string, end bool) (int, error) { + if end && (value == "00:00" || value == "00:00:00") { + return 24 * 60 * 60, nil + } + layout := "15:04:05" + if len(value) == len("15:04") { + layout = "15:04" + } + parsed, err := time.Parse(layout, value) + if err != nil || parsed.Format(layout) != value { + return 0, fmt.Errorf("time %q must use HH:mm or HH:mm:ss format", value) + } + return parsed.Hour()*60*60 + parsed.Minute()*60 + parsed.Second(), nil +} + +func parseChannelTimePeriods(periods []ChannelTimePricingPeriod) ([]parsedChannelTimePeriod, error) { + parsed := make([]parsedChannelTimePeriod, 0, len(periods)) + for _, period := range periods { + if math.IsNaN(period.Multiplier) || math.IsInf(period.Multiplier, 0) || period.Multiplier <= 0 { + return nil, fmt.Errorf("multiplier must be finite and greater than 0") + } + if period.Multiplier < 0.01 { + return nil, fmt.Errorf("multiplier must be at least 0.01") + } + scaled := period.Multiplier * 100 + if math.IsNaN(scaled) || math.IsInf(scaled, 0) { + return nil, fmt.Errorf("multiplier must remain finite when scaled") + } + if math.Abs(scaled-math.Round(scaled)) > 1e-9 { + return nil, fmt.Errorf("multiplier must have at most two decimal places") + } + + start, err := parseChannelTime(period.StartTime, false) + if err != nil { + return nil, err + } + end, err := parseChannelTime(period.EndTime, true) + if err != nil { + return nil, err + } + if period.StartTime == period.EndTime || start >= end { + return nil, fmt.Errorf("start time must be before end time") + } + parsed = append(parsed, parsedChannelTimePeriod{start: start, end: end, multiplier: period.Multiplier}) + } + + sort.Slice(parsed, func(i, j int) bool { + return parsed[i].start < parsed[j].start + }) + for i := 1; i < len(parsed); i++ { + if parsed[i].start < parsed[i-1].end { + return nil, fmt.Errorf("time pricing periods overlap") + } + } + return parsed, nil +} diff --git a/backend/internal/service/custom_channel_time_pricing_test.go b/backend/internal/service/custom_channel_time_pricing_test.go new file mode 100644 index 0000000000..f1038af0d3 --- /dev/null +++ b/backend/internal/service/custom_channel_time_pricing_test.go @@ -0,0 +1,166 @@ +//go:build unit + +package service + +import ( + "math" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +func timeConfig(periods ...ChannelTimePricingPeriod) *ChannelTimePricing { + return &ChannelTimePricing{Timezone: "Asia/Shanghai", Periods: periods} +} + +func onePeriod() []ChannelTimePricingPeriod { + return []ChannelTimePricingPeriod{{StartTime: "09:00", EndTime: "12:00", Multiplier: 2}} +} + +func TestValidateChannelTimePricing(t *testing.T) { + tests := []struct { + name string + config *ChannelTimePricing + wantErr string + }{ + {name: "nil disabled", config: nil}, + {name: "empty disabled", config: &ChannelTimePricing{Timezone: "Asia/Shanghai"}}, + {name: "adjacent", config: timeConfig( + ChannelTimePricingPeriod{StartTime: "09:00", EndTime: "12:00", Multiplier: 2}, + ChannelTimePricingPeriod{StartTime: "12:00", EndTime: "14:00", Multiplier: 1.5})}, + {name: "midnight split", config: timeConfig( + ChannelTimePricingPeriod{StartTime: "22:00", EndTime: "00:00", Multiplier: 2}, + ChannelTimePricingPeriod{StartTime: "00:00", EndTime: "02:00", Multiplier: 2})}, + {name: "second precision", config: timeConfig( + ChannelTimePricingPeriod{StartTime: "09:00:00", EndTime: "12:00:00", Multiplier: 2}, + ChannelTimePricingPeriod{StartTime: "14:00:00", EndTime: "18:00:00", Multiplier: 2})}, + {name: "second precision overlap", config: timeConfig( + ChannelTimePricingPeriod{StartTime: "09:00:00", EndTime: "12:00:00", Multiplier: 2}, + ChannelTimePricingPeriod{StartTime: "11:59:59", EndTime: "14:00:00", Multiplier: 2}), wantErr: "overlap"}, + {name: "empty timezone", config: &ChannelTimePricing{Periods: onePeriod()}, wantErr: "timezone"}, + {name: "whitespace timezone", config: &ChannelTimePricing{Timezone: " ", Periods: onePeriod()}, wantErr: "timezone"}, + {name: "timezone", config: &ChannelTimePricing{Timezone: "UTC+8", Periods: onePeriod()}, wantErr: "timezone"}, + {name: "format", config: timeConfig(ChannelTimePricingPeriod{StartTime: "9:00", EndTime: "12:00", Multiplier: 2}), wantErr: "HH:mm"}, + {name: "equal midnight", config: timeConfig(ChannelTimePricingPeriod{StartTime: "00:00", EndTime: "00:00", Multiplier: 2}), wantErr: "before"}, + {name: "cross midnight", config: timeConfig(ChannelTimePricingPeriod{StartTime: "22:00", EndTime: "02:00", Multiplier: 2}), wantErr: "before"}, + {name: "overlap", config: timeConfig( + ChannelTimePricingPeriod{StartTime: "09:00", EndTime: "12:00", Multiplier: 2}, + ChannelTimePricingPeriod{StartTime: "11:59", EndTime: "14:00", Multiplier: 2}), wantErr: "overlap"}, + {name: "zero", config: timeConfig(ChannelTimePricingPeriod{StartTime: "09:00", EndTime: "12:00", Multiplier: 0}), wantErr: "greater than 0"}, + {name: "minimum positive", config: timeConfig(ChannelTimePricingPeriod{StartTime: "09:00", EndTime: "12:00", Multiplier: 0.01})}, + {name: "tiny positive", config: timeConfig(ChannelTimePricingPeriod{StartTime: "09:00", EndTime: "12:00", Multiplier: 1e-12}), wantErr: "at least 0.01"}, + {name: "below minimum", config: timeConfig(ChannelTimePricingPeriod{StartTime: "09:00", EndTime: "12:00", Multiplier: 0.001}), wantErr: "at least 0.01"}, + {name: "three decimals", config: timeConfig(ChannelTimePricingPeriod{StartTime: "09:00", EndTime: "12:00", Multiplier: 1.001}), wantErr: "decimal"}, + {name: "scaled overflow", config: timeConfig(ChannelTimePricingPeriod{StartTime: "09:00", EndTime: "12:00", Multiplier: math.MaxFloat64}), wantErr: "finite"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := validateChannelTimePricing(tt.config) + if tt.wantErr == "" { + require.NoError(t, err) + return + } + require.Error(t, err) + require.True(t, strings.Contains(err.Error(), tt.wantErr), "error %q does not contain %q", err, tt.wantErr) + }) + } +} + +func TestChannelTimePricingMultiplierAt(t *testing.T) { + config := timeConfig(ChannelTimePricingPeriod{StartTime: "09:00", EndTime: "12:00", Multiplier: 2}) + tests := []struct { + name string + at time.Time + want float64 + }{ + {name: "Shanghai 08:59", at: time.Date(2026, 6, 29, 0, 59, 0, 0, time.UTC), want: 1}, + {name: "Shanghai 09:00", at: time.Date(2026, 6, 29, 1, 0, 0, 0, time.UTC), want: 2}, + {name: "Shanghai 11:59", at: time.Date(2026, 6, 29, 3, 59, 0, 0, time.UTC), want: 2}, + {name: "Shanghai 12:00", at: time.Date(2026, 6, 29, 4, 0, 0, 0, time.UTC), want: 1}, + {name: "Shanghai 14:00", at: time.Date(2026, 6, 29, 6, 0, 0, 0, time.UTC), want: 1}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, config.MultiplierAt(tt.at)) + }) + } + + newYork := &ChannelTimePricing{Timezone: "America/New_York", Periods: onePeriod()} + at := time.Date(2026, 6, 29, 14, 0, 0, 0, time.UTC) + require.Equal(t, 1.0, config.MultiplierAt(at)) + require.Equal(t, 2.0, newYork.MultiplierAt(at)) +} + +func TestChannelTimePricingMultiplierAtSecondPrecision(t *testing.T) { + config := timeConfig(ChannelTimePricingPeriod{StartTime: "09:00:30", EndTime: "09:00:45", Multiplier: 2}) + shanghai, err := time.LoadLocation("Asia/Shanghai") + require.NoError(t, err) + + tests := []struct { + name string + at time.Time + want float64 + }{ + {name: "before", at: time.Date(2026, 6, 29, 9, 0, 29, 0, shanghai), want: 1}, + {name: "start", at: time.Date(2026, 6, 29, 9, 0, 30, 0, shanghai), want: 2}, + {name: "last matching second", at: time.Date(2026, 6, 29, 9, 0, 44, 999_999_999, shanghai), want: 2}, + {name: "end", at: time.Date(2026, 6, 29, 9, 0, 45, 0, shanghai), want: 1}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, config.MultiplierAt(tt.at)) + }) + } +} + +func TestChannelTimePricingMultiplierAtMidnightSplit(t *testing.T) { + config := timeConfig( + ChannelTimePricingPeriod{StartTime: "22:00", EndTime: "00:00", Multiplier: 2}, + ChannelTimePricingPeriod{StartTime: "00:00", EndTime: "02:00", Multiplier: 3}, + ) + shanghai, err := time.LoadLocation("Asia/Shanghai") + require.NoError(t, err) + tests := []struct { + name string + at time.Time + want float64 + }{ + {name: "23:59", at: time.Date(2026, 6, 29, 23, 59, 0, 0, shanghai), want: 2}, + {name: "next day 00:00", at: time.Date(2026, 6, 30, 0, 0, 0, 0, shanghai), want: 3}, + {name: "02:00", at: time.Date(2026, 6, 30, 2, 0, 0, 0, shanghai), want: 1}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + require.Equal(t, tt.want, config.MultiplierAt(tt.at)) + }) + } +} + +func TestChannelTimePricingMultiplierAtDegradesForInvalidConfigurations(t *testing.T) { + var nilConfig *ChannelTimePricing + zeroTime := time.Time{} + validAt := time.Date(2026, 6, 29, 1, 0, 0, 0, time.UTC) + + require.Equal(t, 1.0, nilConfig.MultiplierAt(validAt)) + require.Equal(t, 1.0, timeConfig().MultiplierAt(validAt)) + require.Equal(t, 1.0, timeConfig(ChannelTimePricingPeriod{StartTime: "09:00", EndTime: "12:00", Multiplier: 2}).MultiplierAt(zeroTime)) + require.Equal(t, 1.0, (&ChannelTimePricing{Periods: onePeriod()}).MultiplierAt(validAt)) + require.Equal(t, 1.0, (&ChannelTimePricing{Timezone: " ", Periods: onePeriod()}).MultiplierAt(validAt)) + require.Equal(t, 1.0, (&ChannelTimePricing{Timezone: "UTC+8", Periods: onePeriod()}).MultiplierAt(validAt)) + require.Equal(t, 1.0, timeConfig(ChannelTimePricingPeriod{StartTime: "22:00", EndTime: "02:00", Multiplier: 2}).MultiplierAt(validAt)) +} + +func TestChannelTimePricingRejectsLocalTimezone(t *testing.T) { + config := &ChannelTimePricing{Timezone: "Local", Periods: onePeriod()} + + err := validateChannelTimePricing(config) + require.Error(t, err) + require.Contains(t, err.Error(), "timezone") + require.Equal(t, 1.0, config.MultiplierAt(time.Date(2026, 6, 29, 1, 0, 0, 0, time.UTC))) +} diff --git a/backend/internal/service/gateway_record_usage_test.go b/backend/internal/service/gateway_record_usage_test.go index 517d4723cb..5361644745 100644 --- a/backend/internal/service/gateway_record_usage_test.go +++ b/backend/internal/service/gateway_record_usage_test.go @@ -345,6 +345,38 @@ func TestGatewayServiceRecordUsage_PeakRateAffectsTokenModeImageOutputTokens(t * require.InDelta(t, expectedActual, userRepo.lastAmount, 1e-12) } +func TestGatewayServiceRecordUsage_TimePricingUsesPricingAt(t *testing.T) { + groupID := int64(904) + requestStart := time.Date(2024, time.January, 2, 2, 0, 0, 0, time.UTC) // 上海 10:00 + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + userRepo := &openAIRecordUsageUserRepoStub{} + svc := newGatewayRecordUsageServiceForTest(usageRepo, userRepo, &openAIRecordUsageSubRepoStub{}) + svc.resolver = newOpenAITokenImageChannelPricingResolverWithTimeForTest(t, groupID, "gpt-5.1", &ChannelTimePricing{ + Timezone: "Asia/Shanghai", + Periods: []ChannelTimePricingPeriod{{StartTime: "09:00", EndTime: "12:00", Multiplier: 2}}, + }) + + err := svc.RecordUsage(context.Background(), &RecordUsageInput{ + Result: &ForwardResult{ + RequestID: "gateway_time_pricing_request_start", + Model: "gpt-5.1", + Usage: ClaudeUsage{InputTokens: 1000, OutputTokens: 500}, + }, + APIKey: &APIKey{ID: 804, GroupID: i64p(groupID), Group: &Group{ + ID: groupID, RateMultiplier: 0.8, SubscriptionType: SubscriptionTypeSubscription, + }}, + User: &User{ID: 604}, + Account: &Account{ID: 704}, + PricingAt: requestStart, + }) + + require.NoError(t, err) + require.NotNil(t, usageRepo.lastLog) + baseCost := 1000*3e-6 + 500*15e-6 + require.InDelta(t, baseCost*2, usageRepo.lastLog.TotalCost, 1e-12) + require.InDelta(t, baseCost*2*0.8, usageRepo.lastLog.ActualCost, 1e-12) + require.InDelta(t, 0.8, usageRepo.lastLog.RateMultiplier, 1e-12) +} func TestGatewayServiceRecordUsage_UsesExplicitPricingAtForPeakRate(t *testing.T) { for _, platform := range []string{PlatformAnthropic, PlatformGemini, PlatformGrok, PlatformAntigravity} { t.Run(platform, func(t *testing.T) { diff --git a/backend/internal/service/gateway_usage_billing.go b/backend/internal/service/gateway_usage_billing.go index d28a1be19a..c9671328be 100644 --- a/backend/internal/service/gateway_usage_billing.go +++ b/backend/internal/service/gateway_usage_billing.go @@ -838,7 +838,7 @@ func (s *GatewayService) recordUsageCore(ctx context.Context, input *recordUsage } // 计算费用 - cost := s.calculateRecordUsageCost(ctx, result, apiKey, billingModel, multiplier, imageMultiplier, opts) + cost := s.calculateRecordUsageCost(ctx, result, apiKey, billingModel, multiplier, imageMultiplier, pricingAt, opts) // response_model:按上游成功响应自报的模型计费(渠道显式开启才生效)。 // 采纳条件见 responseModelBillingDeclaration + hasIdentifiedResponseModelPricing // + responseModelBillingAdoptable。任一条件不满足都静默回落基线,即开启本模式前的 @@ -850,7 +850,7 @@ func (s *GatewayService) recordUsageCore(ctx context.Context, input *recordUsage result.ImageCount > 0 || result.AudioUsage != nil || result.SearchCount > 0, ); responseModel != "" && !strings.EqualFold(responseModel, strings.TrimSpace(billingModel)) { if identified, responseChannelPriced := s.hasIdentifiedResponseModelPricing(ctx, responseModel, apiKey); identified { - responseCost := s.calculateRecordUsageCost(ctx, result, apiKey, responseModel, multiplier, imageMultiplier, opts) + responseCost := s.calculateRecordUsageCost(ctx, result, apiKey, responseModel, multiplier, imageMultiplier, pricingAt, opts) baselineChannelPriced := s.resolveChannelPricing(ctx, billingModel, apiKey) != nil if responseModelBillingAdoptable(cost, responseCost, baselineChannelPriced, responseChannelPriced) { // billingModel 到此为止只是定价查表的入参,后续流程只消费 cost, @@ -939,12 +939,13 @@ func (s *GatewayService) calculateRecordUsageCost( billingModel string, multiplier float64, imageMultiplier float64, + pricingAt time.Time, opts *recordUsageOpts, ) *CostBreakdown { // 图片生成:渠道定价为 token 计费时走 token 路径,否则走图片计费 if result.ImageCount > 0 { if resolved := s.resolveChannelPricing(ctx, billingModel, apiKey); resolved != nil && resolved.Mode == BillingModeToken { - return s.calculateTokenCost(ctx, result, apiKey, billingModel, multiplier, opts) + return s.calculateTokenCost(ctx, result, apiKey, billingModel, multiplier, pricingAt, opts) } return s.calculateImageCost(ctx, result, apiKey, billingModel, imageMultiplier) } @@ -968,7 +969,7 @@ func (s *GatewayService) calculateRecordUsageCost( } // Token 计费;SearchCount 为叠加 surcharge(不替代 token)。 - tokenCost := s.calculateTokenCost(ctx, result, apiKey, billingModel, multiplier, opts) + tokenCost := s.calculateTokenCost(ctx, result, apiKey, billingModel, multiplier, pricingAt, opts) if result.SearchCount > 0 { price := groupSearchPricePer1kFromAPIKey(apiKey) if price != nil && *price == 0 { @@ -1127,6 +1128,7 @@ func (s *GatewayService) calculateTokenCost( apiKey *APIKey, billingModel string, multiplier float64, + pricingAt time.Time, opts *recordUsageOpts, ) *CostBreakdown { tokens := UsageTokens{ @@ -1154,6 +1156,7 @@ func (s *GatewayService) calculateTokenCost( Tokens: tokens, RequestCount: 1, RateMultiplier: multiplier, + PricingAt: pricingAt, Resolver: s.resolver, Resolved: resolved, }) @@ -1164,7 +1167,7 @@ func (s *GatewayService) calculateTokenCost( gid := apiKey.Group.ID cost, err = s.billingService.CalculateCostUnified(CostInput{ Ctx: ctx, Model: billingModel, GroupID: &gid, Group: apiKey.Group, - Tokens: tokens, RequestCount: 1, RateMultiplier: multiplier, Resolver: s.resolver, + Tokens: tokens, RequestCount: 1, RateMultiplier: multiplier, PricingAt: pricingAt, Resolver: s.resolver, }) } else { cost, err = s.billingService.CalculateCost(billingModel, tokens, multiplier) diff --git a/backend/internal/service/openai_alpha_search_billing_test.go b/backend/internal/service/openai_alpha_search_billing_test.go index 99cacc88f7..b8c6236d9d 100644 --- a/backend/internal/service/openai_alpha_search_billing_test.go +++ b/backend/internal/service/openai_alpha_search_billing_test.go @@ -5,6 +5,7 @@ package service import ( "context" "testing" + "time" "github.com/Wei-Shaw/sub2api/internal/config" "github.com/stretchr/testify/require" @@ -50,7 +51,7 @@ func TestCalculateOpenAIRecordUsageCostWebSearchPerCall(t *testing.T) { // 即使 token 倍率(含高峰,3.0)更高也不采用。 apiKey := &APIKey{ID: 1, GroupID: &groupID, Group: &Group{ID: groupID, Platform: PlatformOpenAI}} result := &OpenAIForwardResult{Model: "gpt-5.6-sol", UpstreamModel: "gpt-5.6-sol", WebSearchCalls: 1} - cost, err := svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 3.0, 1.0, 1.0, 2.0, UsageTokens{}, "", boolPtr(false)) + cost, err := svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 3.0, 1.0, 1.0, 2.0, UsageTokens{}, "", boolPtr(false), time.Time{}) require.NoError(t, err) require.Equal(t, string(BillingModePerRequest), cost.BillingMode) require.InDelta(t, 0.01, cost.TotalCost, 1e-12) @@ -58,7 +59,7 @@ func TestCalculateOpenAIRecordUsageCostWebSearchPerCall(t *testing.T) { // 分组配置单价 0.005 apiKey.Group.WebSearchPricePerCall = float64Ptr(0.005) - cost, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{}, "", boolPtr(false)) + cost, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{}, "", boolPtr(false), time.Time{}) require.NoError(t, err) require.InDelta(t, 0.005, cost.TotalCost, 1e-12) require.InDelta(t, 0.005, cost.ActualCost, 1e-12) @@ -66,7 +67,7 @@ func TestCalculateOpenAIRecordUsageCostWebSearchPerCall(t *testing.T) { // WebSearchCalls = 0 时不得走按次分支(无定价数据会返回 pricing 错误, // 证明回落到了 token 路径而不是被按次分支吞掉)。 result.WebSearchCalls = 0 - _, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{InputTokens: 10}, "", boolPtr(false)) + _, err = svc.calculateOpenAIRecordUsageCost(context.Background(), result, apiKey, []string{"gpt-5.6-sol"}, 1.0, 1.0, 1.0, 1.0, UsageTokens{InputTokens: 10}, "", boolPtr(false), time.Time{}) require.Error(t, err) } diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go index e59b3b205a..21e5cd4096 100644 --- a/backend/internal/service/openai_gateway_record_usage_test.go +++ b/backend/internal/service/openai_gateway_record_usage_test.go @@ -504,6 +504,71 @@ func TestOpenAIGatewayServiceRecordUsage_PeakRateAffectsTokenModeImageOutputToke require.InDelta(t, expectedActual, userRepo.lastAmount, 1e-12) } +func TestOpenAIGatewayServiceRecordUsage_TimePricingUsesPricingAt(t *testing.T) { + groupID := int64(16) + requestStart := time.Date(2024, time.January, 2, 2, 0, 0, 0, time.UTC) // 上海 10:00 + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + userRepo := &openAIRecordUsageUserRepoStub{} + svc := newOpenAIRecordUsageServiceForTest(usageRepo, userRepo, &openAIRecordUsageSubRepoStub{}, nil) + svc.resolver = newOpenAITokenImageChannelPricingResolverWithTimeForTest(t, groupID, "gpt-5.1", &ChannelTimePricing{ + Timezone: "Asia/Shanghai", + Periods: []ChannelTimePricingPeriod{{StartTime: "09:00", EndTime: "12:00", Multiplier: 2}}, + }) + + err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{ + RequestID: "openai_time_pricing_request_start", + Model: "gpt-5.1", + Usage: OpenAIUsage{InputTokens: 1000, OutputTokens: 500}, + }, + APIKey: &APIKey{ID: 1006, GroupID: i64p(groupID), Group: &Group{ + ID: groupID, RateMultiplier: 0.8, SubscriptionType: SubscriptionTypeSubscription, + }}, + User: &User{ID: 2006}, + Account: &Account{ID: 3006}, + PricingAt: requestStart, + }) + + require.NoError(t, err) + require.NotNil(t, usageRepo.lastLog) + baseCost := 1000*3e-6 + 500*15e-6 + require.InDelta(t, baseCost*2, usageRepo.lastLog.TotalCost, 1e-12) + require.InDelta(t, baseCost*2*0.8, usageRepo.lastLog.ActualCost, 1e-12) + require.InDelta(t, 0.8, usageRepo.lastLog.RateMultiplier, 1e-12) +} + +func TestOpenAIGatewayServiceRecordUsage_TimePricingUsesExplicitPricingAt(t *testing.T) { + groupID := int64(17) + pricingAt := time.Date(2024, time.January, 2, 0, 0, 0, 0, time.UTC) // 上海 08:00 + usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} + userRepo := &openAIRecordUsageUserRepoStub{} + svc := newOpenAIRecordUsageServiceForTest(usageRepo, userRepo, &openAIRecordUsageSubRepoStub{}, nil) + svc.resolver = newOpenAITokenImageChannelPricingResolverWithTimeForTest(t, groupID, "gpt-5.1", &ChannelTimePricing{ + Timezone: "Asia/Shanghai", + Periods: []ChannelTimePricingPeriod{{StartTime: "09:00", EndTime: "12:00", Multiplier: 2}}, + }) + + err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{ + RequestID: "openai_time_pricing_explicit", + Model: "gpt-5.1", + Usage: OpenAIUsage{InputTokens: 1000, OutputTokens: 500}, + }, + APIKey: &APIKey{ID: 1007, GroupID: i64p(groupID), Group: &Group{ + ID: groupID, RateMultiplier: 0.8, SubscriptionType: SubscriptionTypeSubscription, + }}, + User: &User{ID: 2007}, + Account: &Account{ID: 3007}, + PricingAt: pricingAt, + }) + + require.NoError(t, err) + require.NotNil(t, usageRepo.lastLog) + baseCost := 1000*3e-6 + 500*15e-6 + require.InDelta(t, baseCost, usageRepo.lastLog.TotalCost, 1e-12) + require.InDelta(t, baseCost*0.8, usageRepo.lastLog.ActualCost, 1e-12) + require.InDelta(t, 0.8, usageRepo.lastLog.RateMultiplier, 1e-12) +} func TestOpenAIGatewayServiceRecordUsage_IncludesEndpointMetadata(t *testing.T) { usageRepo := &openAIRecordUsageLogRepoStub{inserted: true} userRepo := &openAIRecordUsageUserRepoStub{} @@ -2640,6 +2705,20 @@ func newOpenAITokenImageChannelPricingResolverForTest(t *testing.T, groupID int6 return NewModelPricingResolver(cs, NewBillingService(&config.Config{}, nil)) } +func newOpenAITokenImageChannelPricingResolverWithTimeForTest( + t *testing.T, + groupID int64, + model string, + timePricing *ChannelTimePricing, +) *ModelPricingResolver { + t.Helper() + resolver := newOpenAITokenImageChannelPricingResolverForTest(t, groupID, model) + cached, ok := resolver.channelService.cache.Load().(*channelCache) + require.True(t, ok) + cached.pricingByGroupModel[channelModelKey{groupID: groupID, model: model}].TimePricing = timePricing + return resolver +} + type openAIMediaPriceGroupRepoStub struct { GroupRepository group *Group @@ -2668,6 +2747,7 @@ func TestGatewayServiceCalculateRecordUsageCost_ChannelImageBillingUsesImageCoun "gemini-image", 0.15, 1.0, + time.Time{}, nil, ) @@ -2707,6 +2787,7 @@ func TestGatewayServiceCalculateRecordUsageCost_ChannelImageBillingUsesSizeTier( "gemini-image", 1.0, 1.0, + time.Time{}, nil, ) @@ -2739,6 +2820,7 @@ func TestGatewayServiceCalculateRecordUsageCost_GroupImagePriceOverridesChannelI "gemini-image", 1.0, 1.0, + time.Time{}, nil, ) @@ -2802,6 +2884,7 @@ func TestGatewayServiceCalculateRecordUsageCost_ChannelImageBillingNormalizesMis "gemini-image", 1.0, 1.0, + time.Time{}, nil, ) diff --git a/backend/internal/service/openai_gateway_search_surcharge_test.go b/backend/internal/service/openai_gateway_search_surcharge_test.go index 2633a5b930..d2d069ff5c 100644 --- a/backend/internal/service/openai_gateway_search_surcharge_test.go +++ b/backend/internal/service/openai_gateway_search_surcharge_test.go @@ -5,6 +5,7 @@ package service import ( "context" "testing" + "time" "github.com/stretchr/testify/require" ) @@ -37,6 +38,7 @@ func TestCalculateOpenAIRecordUsageCost_SearchIsAdditiveToTokens(t *testing.T) { UsageTokens{InputTokens: 1000, OutputTokens: 500}, "", boolPtr(false), + time.Time{}, ) require.NoError(t, err) require.NotNil(t, cost) @@ -67,6 +69,7 @@ func TestCalculateOpenAIRecordUsageCost_SearchOnlyWhenNoTokenPricing(t *testing. UsageTokens{}, "", boolPtr(false), + time.Time{}, ) require.NoError(t, err) require.NotNil(t, cost) @@ -112,6 +115,7 @@ func TestCalculateOpenAIRecordUsageCost_TokenPricingErrorNotSwallowedBySearch(t UsageTokens{InputTokens: 1000, OutputTokens: 500}, "", boolPtr(false), + time.Time{}, ) require.Error(t, err) require.Nil(t, cost) diff --git a/backend/internal/service/openai_gateway_usage.go b/backend/internal/service/openai_gateway_usage.go index c55bcdbeac..a8b3c0fc0f 100644 --- a/backend/internal/service/openai_gateway_usage.go +++ b/backend/internal/service/openai_gateway_usage.go @@ -177,7 +177,8 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec // 变价);未装配 PricingAt 的路径回退记录时刻,保持既有行为。不并入上面的 // Resolve,以免污染 user:group 倍率缓存。 baseMultiplier := multiplier - multiplier, imageMultiplier := computePeakAwareMultipliers(apiKey, baseMultiplier, openAIUsagePricingAt(input)) + pricingAt := openAIUsagePricingAt(input) + multiplier, imageMultiplier := computePeakAwareMultipliers(apiKey, baseMultiplier, pricingAt) videoMultiplier := resolveVideoRateMultiplier(apiKey, baseMultiplier) var cost *CostBreakdown @@ -224,6 +225,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec tokens, serviceTier, longContextBillingGate, + pricingAt, ) if err != nil { if !isUsagePricingUnavailableError(err) { @@ -256,7 +258,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec responseModels := usageBillingModelCandidates(responseModel) responseCost, responseErr := s.calculateOpenAIRecordUsageCost( ctx, result, apiKey, responseModels, multiplier, imageMultiplier, - videoMultiplier, baseMultiplier, tokens, serviceTier, longContextBillingGate, + videoMultiplier, baseMultiplier, tokens, serviceTier, longContextBillingGate, pricingAt, ) // 基线定价源以 baselineBillingModel 为准:它正是 calculateOpenAIRecordUsageCost // 内部做渠道定价判断时使用的模型,且"首候选有渠道价"必然意味着首候选就是实际 @@ -505,6 +507,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost( tokens UsageTokens, serviceTier string, longContextBillingGate *bool, + pricingAt time.Time, ) (*CostBreakdown, error) { billingModel := firstUsageBillingModel(billingModels) if result != nil && result.WebSearchCalls > 0 { @@ -554,6 +557,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost( apiKey, candidate, multiplier, + pricingAt, tokens, serviceTier, longContextBillingGate, @@ -643,6 +647,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageTokenCost( apiKey *APIKey, billingModel string, multiplier float64, + pricingAt time.Time, tokens UsageTokens, serviceTier string, longContextBillingGate *bool, @@ -651,7 +656,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageTokenCost( gid := apiKey.Group.ID return s.billingService.CalculateCostUnified(CostInput{ Ctx: ctx, Model: billingModel, GroupID: &gid, Group: apiKey.Group, - Tokens: tokens, RequestCount: 1, RateMultiplier: multiplier, + Tokens: tokens, RequestCount: 1, RateMultiplier: multiplier, PricingAt: pricingAt, ServiceTier: serviceTier, Resolver: s.resolver, LongContextBillingEnabled: longContextBillingGate, }) diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go index 51898060ac..1928c1d6e4 100644 --- a/backend/internal/service/openai_ws_forwarder.go +++ b/backend/internal/service/openai_ws_forwarder.go @@ -215,10 +215,13 @@ type OpenAIWSIngressHooks struct { // before channel or account mapping. Ingress modes preserve it for usage // attribution while MapRequestModel determines the upstream model. InitialRequestModel string + // InitialTurnStartedAt freezes when the first response.create was accepted. + InitialTurnStartedAt time.Time // MaxReasoningEffort limits explicit reasoning effort values for this WS session. MaxReasoningEffort string // ReasoningEffortMappings rewrites explicit effort values for this WS session. ReasoningEffortMappings []ReasoningEffortMapping + TurnStarted func(turn int, startedAt time.Time) BeforeTurn func(turn int) error BeforeRequest func(turn int, payload []byte, originalModel string) error // MapRequestModel resolves the current turn's client model to the model diff --git a/backend/internal/service/openai_ws_forwarder_ingress.go b/backend/internal/service/openai_ws_forwarder_ingress.go index ce7e8f0386..bdd59f5e1c 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress.go +++ b/backend/internal/service/openai_ws_forwarder_ingress.go @@ -92,11 +92,9 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( if wsDecision.Transport != OpenAIUpstreamTransportResponsesWebsocketV2 { return fmt.Errorf("websocket ingress requires ws_v2 transport, got=%s", wsDecision.Transport) } - // 注意:透传 relay 只回调 hooks.AfterTurn,没有 turn 起始回调, - // 因此下面这条路径永远不会触发 hooks.BeforeTurn——分组利润控制的 - // turn 级复核与 turn 级 pricingAt 冻结都不覆盖透传 ingress, - // 只有建连时的准入门生效。handler 侧据此把 turn 定价留作零值, - // 由 RecordUsage 回退到记录时刻(见 openAIWSTurnPricing 注释)。 + // 透传 relay 通过 TurnStarted 记录每个 turn 的开始时刻,但不触发 + // BeforeTurn;因此仍只有建连时的利润准入门,没有 turn 级复核。 + // handler 计费在 turn 定价未冻结时回退到对应的 turn 开始时刻。 return s.proxyResponsesWebSocketV2Passthrough( ctx, c, diff --git a/backend/internal/service/openai_ws_passthrough_turn_pricing_test.go b/backend/internal/service/openai_ws_passthrough_turn_pricing_test.go index bb6e5ebde6..16be2bd4c6 100644 --- a/backend/internal/service/openai_ws_passthrough_turn_pricing_test.go +++ b/backend/internal/service/openai_ws_passthrough_turn_pricing_test.go @@ -60,18 +60,13 @@ func startPassthroughHookRecordingServer( return server, serverErr } -// TestPassthroughIngressNeverCallsBeforeTurn 钉死 ws_v2 透传 ingress 与 handler -// 侧 turn 定价的耦合:透传 relay 只回调 AfterTurn,没有任何 turn 起始回调, -// 因此 hooks.BeforeTurn 永远不会触发。 +// TestPassthroughIngressReportsTurnStartedBeforeAfterTurnWithoutBeforeTurn 钉死 +// ws_v2 透传 ingress 与 handler 侧 turn 定价的耦合:透传 relay 不触发 +// BeforeTurn,但会在每个 AfterTurn 前通过 TurnStarted 报告同一 turn 的开始时刻。 // -// handler 依赖这一点:openAIWSTurnPricing 零值起步,透传连接的每个 turn 都拿 -// 不到冻结的 pricingAt,RecordUsage 回退到记录时刻——与引入分组利润控制前的 -// 基线一致。若把 turn 定价初始化成建连时刻,透传连接的所有 turn 就会被钉死在 -// 建连时的高峰因子,客户端峰前建连保活即可全程按谷价结算。 -// -// 若本断言因为透传补齐了 turn 起始回调而失败:这是好事,请同步复核 -// openAIWSTurnPricing 的零值语义与透传路径的 turn 级利润复核。 -func TestPassthroughIngressNeverCallsBeforeTurn(t *testing.T) { +// handler 的 recordTurnStart 保存该时刻,AfterTurn 再用 currentOr(turnStart) +// 作为计费 PricingAt;不触发 BeforeTurn 也意味着透传仍没有 turn 级利润复核。 +func TestPassthroughIngressReportsTurnStartedBeforeAfterTurnWithoutBeforeTurn(t *testing.T) { gin.SetMode(gin.TestMode) controlCtx, cancelControl := context.WithCancelCause(context.Background()) defer cancelControl(context.Canceled) @@ -81,17 +76,29 @@ func TestPassthroughIngressNeverCallsBeforeTurn(t *testing.T) { var hooksMu sync.Mutex beforeTurnCalls := 0 - afterTurnCalls := 0 + expectedTurnStartedAt := time.Date(2026, time.August, 17, 9, 59, 59, 0, time.UTC) + type hookEvent struct { + name string + turn int + startedAt time.Time + } + var hookEvents []hookEvent hooks := &OpenAIWSIngressHooks{ + InitialTurnStartedAt: expectedTurnStartedAt, + TurnStarted: func(turn int, startedAt time.Time) { + hooksMu.Lock() + hookEvents = append(hookEvents, hookEvent{name: "TurnStarted", turn: turn, startedAt: startedAt}) + hooksMu.Unlock() + }, BeforeTurn: func(int) error { hooksMu.Lock() beforeTurnCalls++ hooksMu.Unlock() return nil }, - AfterTurn: func(int, *OpenAIForwardResult, error) { + AfterTurn: func(turn int, _ *OpenAIForwardResult, _ error) { hooksMu.Lock() - afterTurnCalls++ + hookEvents = append(hookEvents, hookEvent{name: "AfterTurn", turn: turn}) hooksMu.Unlock() }, } @@ -120,9 +127,109 @@ func TestPassthroughIngressNeverCallsBeforeTurn(t *testing.T) { } hooksMu.Lock() - gotBefore, gotAfter := beforeTurnCalls, afterTurnCalls + gotBefore := beforeTurnCalls + gotEvents := append([]hookEvent(nil), hookEvents...) hooksMu.Unlock() - require.Zero(t, gotBefore, "透传 ingress 没有 turn 起始回调,BeforeTurn 不应被调用") - require.Positive(t, gotAfter, "透传 ingress 仍应回调 AfterTurn 提交用量") + require.Zero(t, gotBefore, "透传 ingress 不应调用 BeforeTurn") + require.GreaterOrEqual(t, len(gotEvents), 2, "透传 ingress 应报告 TurnStarted 和 AfterTurn") + require.Equal(t, "TurnStarted", gotEvents[0].name) + require.Equal(t, expectedTurnStartedAt, gotEvents[0].startedAt, "TurnStarted 必须携带入口冻结的首轮开始时刻") + require.Equal(t, "AfterTurn", gotEvents[1].name) + require.Equal(t, gotEvents[0].turn, gotEvents[1].turn, "TurnStarted 后应提交同一 turn 的 AfterTurn") +} + +func TestPassthroughIngressFreezesSubsequentTurnBeforeRequestPolicy(t *testing.T) { + testPassthroughIngressFreezesSubsequentTurnBeforeRequestPolicy(t, coderws.MessageText) +} + +func TestPassthroughIngressFreezesBinarySubsequentTurnBeforeRequestPolicy(t *testing.T) { + testPassthroughIngressFreezesSubsequentTurnBeforeRequestPolicy(t, coderws.MessageBinary) +} + +func testPassthroughIngressFreezesSubsequentTurnBeforeRequestPolicy(t *testing.T, secondMessageType coderws.MessageType) { + t.Helper() + gin.SetMode(gin.TestMode) + controlCtx, cancelControl := context.WithCancelCause(context.Background()) + defer cancelControl(context.Canceled) + + upstream := newStagedPassthroughConn() + upstream.Send(`{"type":"response.completed","response":{"id":"resp_first","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`) + + type turnStart struct { + turn int + startedAt time.Time + } + turnStarts := make(chan turnStart, 2) + beforeRequestEntered := make(chan time.Time, 1) + releaseBeforeRequest := make(chan struct{}) + hooks := &OpenAIWSIngressHooks{ + InitialTurnStartedAt: time.Now(), + TurnStarted: func(turn int, startedAt time.Time) { + turnStarts <- turnStart{turn: turn, startedAt: startedAt} + }, + BeforeRequest: func(turn int, _ []byte, _ string) error { + if turn == 2 { + beforeRequestEntered <- time.Now() + <-releaseBeforeRequest + } + return nil + }, + } + + server, serverErr := startPassthroughHookRecordingServer( + t, + controlCtx, + newPassthroughLifecycleService(passthroughLifecycleConfig(), upstream), + passthroughLifecycleAccount(), + hooks, + ) + defer server.Close() + clientConn := dialPassthroughLifecycleClient(t, server) + defer func() { _ = clientConn.CloseNow() }() + + require.Equal(t, "response.create", gjson.GetBytes(requirePassthroughUpstreamWrite(t, upstream, 3*time.Second), "type").String()) + firstCompleted, err := readPassthroughLifecycleFrame(t, clientConn, 3*time.Second) + require.NoError(t, err) + require.Equal(t, "resp_first", gjson.GetBytes(firstCompleted, "response.id").String()) + select { + case first := <-turnStarts: + require.Equal(t, 1, first.turn) + case <-time.After(time.Second): + t.Fatal("first turn start was not reported") + } + + writeCtx, cancelWrite := context.WithTimeout(context.Background(), 3*time.Second) + err = clientConn.Write(writeCtx, secondMessageType, []byte(`{"type":"response.create","model":"gpt-5.1","previous_response_id":"resp_first"}`)) + cancelWrite() + require.NoError(t, err) + + var policyEnteredAt time.Time + select { + case policyEnteredAt = <-beforeRequestEntered: + case <-time.After(time.Second): + t.Fatal("second turn did not enter BeforeRequest") + } + close(releaseBeforeRequest) + require.Equal(t, "response.create", gjson.GetBytes(requirePassthroughUpstreamWrite(t, upstream, 3*time.Second), "type").String()) + upstream.Send(`{"type":"response.completed","response":{"id":"resp_second","model":"gpt-5.1","usage":{"input_tokens":1,"output_tokens":1}}}`) + secondCompleted, err := readPassthroughLifecycleFrame(t, clientConn, 3*time.Second) + require.NoError(t, err) + require.Equal(t, "resp_second", gjson.GetBytes(secondCompleted, "response.id").String()) + + select { + case second := <-turnStarts: + require.Equal(t, 2, second.turn) + require.False(t, second.startedAt.After(policyEnteredAt), "第二轮开始时刻必须在 BeforeRequest 策略执行前冻结") + case <-time.After(time.Second): + t.Fatal("second turn start was not reported") + } + + _ = clientConn.CloseNow() + cancelControl(context.Canceled) + select { + case <-serverErr: + case <-time.After(3 * time.Second): + t.Fatal("passthrough ingress did not exit") + } } diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay.go b/backend/internal/service/openai_ws_v2/passthrough_relay.go index 626faf891a..abc94e8f1c 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay.go @@ -50,6 +50,7 @@ type RelayTurnResult struct { Usage Usage RequestID string TerminalEventType string + StartedAt time.Time Duration time.Duration FirstTokenMs *int } @@ -65,6 +66,8 @@ type RelayOptions struct { WriteTimeout time.Duration IdleTimeout time.Duration UpstreamDrainTimeout time.Duration + FirstTurnStartedAt time.Time + TakeNextTurnStartedAt func() time.Time FirstMessageType coderws.MessageType FirstMessageSent bool StartClientAfterFirstDownstream bool @@ -93,6 +96,7 @@ type relayState struct { usage Usage requestModelMu sync.RWMutex requestModel string + pendingTurnStart atomic.Pointer[time.Time] lastResponseID string lastResponseModel string responseConflict bool @@ -114,6 +118,7 @@ type observedUpstreamEvent struct { eventType string responseID string usage Usage + startedAt time.Time responseModel string responseConflict bool duration time.Duration @@ -161,6 +166,13 @@ func Relay( } startAt := nowFn() state := &relayState{requestModel: result.RequestModel} + if isClientResponseCreateFrame(firstMessageType, firstClientMessage) { + firstTurnStartedAt := options.FirstTurnStartedAt + if firstTurnStartedAt.IsZero() { + firstTurnStartedAt = startAt + } + state.setPendingTurnStartedAt(firstTurnStartedAt) + } onTrace := options.OnTrace relayCtx, relayCancel := context.WithCancel(ctx) @@ -178,8 +190,16 @@ func Relay( return upstreamConn.WriteFrame(writeCtx, msgType, payload) } writeClientFrameUpstream := func(msgType coderws.MessageType, payload []byte) error { - if msgType == coderws.MessageText && strings.TrimSpace(gjson.GetBytes(payload, "type").String()) == "response.create" { + if isClientResponseCreateFrame(msgType, payload) { state.setRequestModel(strings.TrimSpace(gjson.GetBytes(payload, "model").String())) + turnStartedAt := time.Time{} + if options.TakeNextTurnStartedAt != nil { + turnStartedAt = options.TakeNextTurnStartedAt() + } + if turnStartedAt.IsZero() { + turnStartedAt = nowFn() + } + state.setPendingTurnStartedAt(turnStartedAt) } return writeUpstream(msgType, payload) } @@ -412,6 +432,13 @@ func Relay( return result, nil } +func isClientResponseCreateFrame(msgType coderws.MessageType, payload []byte) bool { + if msgType != coderws.MessageText && msgType != coderws.MessageBinary { + return false + } + return strings.TrimSpace(gjson.GetBytes(payload, "type").String()) == "response.create" +} + func runClientToUpstream( ctx context.Context, clientConn FrameConn, @@ -724,6 +751,7 @@ func observeUpstreamMessage( if duration < 0 { duration = 0 } + observed.startedAt = turnTiming.startAt observed.duration = duration observed.firstToken = openAIWSRelayCloneIntPtr(turnTiming.firstTokenMs) } @@ -754,6 +782,7 @@ func emitTurnComplete( Usage: observed.usage, RequestID: responseID, TerminalEventType: observed.eventType, + StartedAt: observed.startedAt, Duration: observed.duration, FirstTokenMs: openAIWSRelayCloneIntPtr(observed.firstToken), }) @@ -815,7 +844,11 @@ func openAIWSRelayGetOrInitTurnTiming(state *relayState, responseID string, now } timing, ok := state.turnTimingByID[responseID] if !ok || timing == nil || timing.startAt.IsZero() { - timing = &relayTurnTiming{startAt: now} + startAt := state.consumePendingTurnStartedAt() + if startAt.IsZero() { + startAt = now + } + timing = &relayTurnTiming{startAt: startAt} state.turnTimingByID[responseID] = timing state.activeTurn = timing return timing @@ -823,6 +856,25 @@ func openAIWSRelayGetOrInitTurnTiming(state *relayState, responseID string, now return timing } +func (s *relayState) setPendingTurnStartedAt(startedAt time.Time) { + if s == nil || startedAt.IsZero() { + return + } + startedAtCopy := startedAt + s.pendingTurnStart.Store(&startedAtCopy) +} + +func (s *relayState) consumePendingTurnStartedAt() time.Time { + if s == nil { + return time.Time{} + } + startedAt := s.pendingTurnStart.Swap(nil) + if startedAt == nil { + return time.Time{} + } + return *startedAt +} + func openAIWSRelayDeleteTurnTiming(state *relayState, responseID string) (relayTurnTiming, bool) { if state == nil || state.turnTimingByID == nil { return relayTurnTiming{}, false diff --git a/backend/internal/service/openai_ws_v2/passthrough_relay_test.go b/backend/internal/service/openai_ws_v2/passthrough_relay_test.go index c41e7d293b..12e5506698 100644 --- a/backend/internal/service/openai_ws_v2/passthrough_relay_test.go +++ b/backend/internal/service/openai_ws_v2/passthrough_relay_test.go @@ -565,6 +565,143 @@ func TestRelay_OnTurnComplete_ProvidesTurnMetrics(t *testing.T) { require.Greater(t, result.Duration.Milliseconds(), int64(0)) } +func TestRelay_OnTurnComplete_UsesResponseCreateTimeAcrossPricingBoundary(t *testing.T) { + t.Parallel() + + clientConn := newPassthroughTestFrameConn(nil, false) + upstreamConn := newPassthroughTestFrameConn([]passthroughTestFrame{ + { + msgType: coderws.MessageText, + payload: []byte(`{"type":"response.completed","response":{"id":"resp_boundary","usage":{"input_tokens":1,"output_tokens":1}}}`), + }, + }, true) + + responseCreateAt := time.Date(2026, time.August, 17, 9, 59, 59, 0, time.UTC) + upstreamResponseAt := responseCreateAt.Add(time.Second) + var nowCalls atomic.Int64 + nowFn := func() time.Time { + if nowCalls.Add(1) == 1 { + return responseCreateAt + } + return upstreamResponseAt + } + + var turn RelayTurnResult + _, relayExit := Relay( + context.Background(), + clientConn, + upstreamConn, + []byte(`{"type":"response.create","model":"gpt-5.3-codex","input":[]}`), + RelayOptions{ + Now: nowFn, + OnTurnComplete: func(current RelayTurnResult) { turn = current }, + }, + ) + + require.Nil(t, relayExit) + require.Equal(t, responseCreateAt, turn.StartedAt) +} + +func TestRelay_OnTurnComplete_UsesExplicitFirstTurnStartedAt(t *testing.T) { + t.Parallel() + + clientConn := newPassthroughTestFrameConn(nil, false) + upstreamConn := newPassthroughTestFrameConn([]passthroughTestFrame{ + { + msgType: coderws.MessageText, + payload: []byte(`{"type":"response.completed","response":{"id":"resp_initial_boundary","usage":{"input_tokens":1,"output_tokens":1}}}`), + }, + }, true) + + responseCreateAt := time.Date(2026, time.August, 17, 9, 59, 59, 0, time.UTC) + relayStartedAt := responseCreateAt.Add(time.Second) + var turn RelayTurnResult + _, relayExit := Relay( + context.Background(), + clientConn, + upstreamConn, + []byte(`{"type":"response.create","model":"gpt-5.3-codex","input":[]}`), + RelayOptions{ + FirstTurnStartedAt: responseCreateAt, + Now: func() time.Time { return relayStartedAt }, + OnTurnComplete: func(current RelayTurnResult) { turn = current }, + }, + ) + + require.Nil(t, relayExit) + require.Equal(t, responseCreateAt, turn.StartedAt) +} + +func TestRelay_OnTurnComplete_UsesSubsequentResponseCreateTimeAcrossPricingBoundary(t *testing.T) { + t.Parallel() + + clientConn := newPassthroughTestFrameConn(nil, false) + upstreamConn := newPassthroughTestFrameConn(nil, false) + firstTurnAt := time.Date(2026, time.August, 17, 9, 0, 0, 0, time.UTC) + secondTurnAt := time.Date(2026, time.August, 17, 9, 59, 59, 0, time.UTC) + secondResponseAt := secondTurnAt.Add(time.Second) + var clock atomic.Int64 + clock.Store(firstTurnAt.UnixNano()) + + turns := make(chan RelayTurnResult, 2) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + done := make(chan struct{}) + go func() { + defer close(done) + _, _ = Relay( + ctx, + clientConn, + upstreamConn, + []byte(`{"type":"response.create","model":"gpt-5.3-codex","input":[]}`), + RelayOptions{ + Now: func() time.Time { return time.Unix(0, clock.Load()).UTC() }, + OnTurnComplete: func(current RelayTurnResult) { + turns <- current + }, + }, + ) + }() + + require.Eventually(t, func() bool { return len(upstreamConn.Writes()) == 1 }, time.Second, time.Millisecond) + upstreamConn.readCh <- passthroughTestFrame{ + msgType: coderws.MessageText, + payload: []byte(`{"type":"response.completed","response":{"id":"resp_first","usage":{"input_tokens":1,"output_tokens":1}}}`), + } + select { + case <-turns: + case <-time.After(time.Second): + t.Fatal("first turn did not complete") + } + + clock.Store(secondTurnAt.UnixNano()) + clientConn.readCh <- passthroughTestFrame{ + msgType: coderws.MessageText, + payload: []byte(`{"type":"response.create","model":"gpt-5.3-codex","input":[]}`), + } + require.Eventually(t, func() bool { return len(upstreamConn.Writes()) == 2 }, time.Second, time.Millisecond) + clock.Store(secondResponseAt.UnixNano()) + upstreamConn.readCh <- passthroughTestFrame{ + msgType: coderws.MessageText, + payload: []byte(`{"type":"response.completed","response":{"id":"resp_second","usage":{"input_tokens":1,"output_tokens":1}}}`), + } + + var secondTurn RelayTurnResult + select { + case secondTurn = <-turns: + case <-time.After(time.Second): + t.Fatal("second turn did not complete") + } + require.Equal(t, secondTurnAt, secondTurn.StartedAt) + + cancel() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("relay did not stop after cancellation") + } +} + func TestRelay_BinaryFramePassthrough(t *testing.T) { t.Parallel() diff --git a/backend/internal/service/openai_ws_v2_passthrough_adapter.go b/backend/internal/service/openai_ws_v2_passthrough_adapter.go index 775e26980c..df094a3367 100644 --- a/backend/internal/service/openai_ws_v2_passthrough_adapter.go +++ b/backend/internal/service/openai_ws_v2_passthrough_adapter.go @@ -915,6 +915,7 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( completedTurns := atomic.Int32{} turnLifecycle := newOpenAIWSPassthroughTurnLifecycle(true) + var acceptedTurnStartedAt atomic.Pointer[time.Time] clientFrameConn := &openAIWSClientFrameConn{ conn: clientConn, controlCtx: ctx, @@ -936,13 +937,15 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( // capturedSessionModel 的读写都发生在该 goroutine 内,因此无需 // 加锁/原子化。 filter: func(msgType coderws.MessageType, payload []byte) (out []byte, blocked *OpenAIFastBlockedError, filterErr error) { - if msgType != coderws.MessageText { + if msgType != coderws.MessageText && msgType != coderws.MessageBinary { return payload, nil, nil } eventType := strings.TrimSpace(gjson.GetBytes(payload, "type").String()) isResponseCreate := eventType == "response.create" + responseCreateAt := time.Time{} acceptedTurn := false if isResponseCreate { + responseCreateAt = time.Now() if !turnLifecycle.beginResponseCreate(clientFrameConn.markTurnStarted) { err := errors.New("overlapping response.create is not supported") return payload, nil, NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, err.Error(), err) @@ -1035,6 +1038,8 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( // service_tier 时按 default 处理,billing 应如实反映。 if policyErr == nil && blocked == nil && isResponseCreate { usageMeta.updateFromResponseCreate(out, model, requestModelForThisFrame) + responseCreateAtCopy := responseCreateAt + acceptedTurnStartedAt.Store(&responseCreateAtCopy) acceptedTurn = true } return out, blocked, policyErr @@ -1072,7 +1077,7 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( if readErr != nil { return msgType, payload, readErr } - if msgType == coderws.MessageText && strings.TrimSpace(gjson.GetBytes(payload, "type").String()) == "response.create" { + if (msgType == coderws.MessageText || msgType == coderws.MessageBinary) && strings.TrimSpace(gjson.GetBytes(payload, "type").String()) == "response.create" { return msgType, payload, nil } if writeErr := upstreamFrameConn.WriteFrame(readCtx, msgType, payload); writeErr != nil { @@ -1081,13 +1086,25 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( } } + firstTurnStartedAt := time.Time{} + if hooks != nil { + firstTurnStartedAt = hooks.InitialTurnStartedAt + } relayResult, relayExit := openaiwsv2.RunEntry(openaiwsv2.EntryInput{ Ctx: ctx, ClientConn: policyClientConn, UpstreamConn: relayUpstreamFrameConn, FirstClientMessage: firstClientMessage, Options: openaiwsv2.RelayOptions{ - WriteTimeout: s.openAIWSWriteTimeout(), + WriteTimeout: s.openAIWSWriteTimeout(), + FirstTurnStartedAt: firstTurnStartedAt, + TakeNextTurnStartedAt: func() time.Time { + startedAt := acceptedTurnStartedAt.Swap(nil) + if startedAt == nil { + return time.Time{} + } + return *startedAt + }, // Passthrough idle is enforced only after a completed turn by // clientFrameConn. The relay-wide activity watchdog would also // terminate a healthy active upstream turn. @@ -1105,6 +1122,9 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( }, OnTurnComplete: func(turn openaiwsv2.RelayTurnResult) { turnNo := int(completedTurns.Add(1)) + if hooks != nil && hooks.TurnStarted != nil && !turn.StartedAt.IsZero() { + hooks.TurnStarted(turnNo, turn.StartedAt) + } turnRequestModel, turnUpstreamModel := usageMeta.turnModels(turn.RequestModel) turnResult := &OpenAIForwardResult{ RequestID: turn.RequestID, @@ -1265,6 +1285,9 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( ) // 正常路径按 terminal 事件逐 turn 已回调;仅在零 turn 场景兜底回调一次。 if turnCount == 0 && hooks != nil && hooks.AfterTurn != nil { + if hooks.TurnStarted != nil { + hooks.TurnStarted(1, time.Now().Add(-result.Duration)) + } hooks.AfterTurn(1, result, nil) } return nil @@ -1331,6 +1354,9 @@ func (s *OpenAIGatewayService) proxyResponsesWebSocketV2Passthrough( relayExit.WroteDownstream, ) if hooks != nil && hooks.AfterTurn != nil { + if hooks.TurnStarted != nil { + hooks.TurnStarted(turnCount+1, time.Now().Add(-result.Duration)) + } hooks.AfterTurn(turnCount+1, nil, turnErr) } return turnErr diff --git a/backend/migrations/225_channel_model_time_pricing.sql b/backend/migrations/225_channel_model_time_pricing.sql new file mode 100644 index 0000000000..dbfb84cde4 --- /dev/null +++ b/backend/migrations/225_channel_model_time_pricing.sql @@ -0,0 +1,5 @@ +ALTER TABLE channel_model_pricing + ADD COLUMN IF NOT EXISTS time_pricing JSONB NULL; + +COMMENT ON COLUMN channel_model_pricing.time_pricing IS + 'Optional IANA timezone and recurring daily multiplier periods for channel token pricing'; diff --git a/frontend/src/api/admin/channels.ts b/frontend/src/api/admin/channels.ts index fdbeadf57a..6556417ced 100644 --- a/frontend/src/api/admin/channels.ts +++ b/frontend/src/api/admin/channels.ts @@ -21,6 +21,17 @@ export interface PricingInterval { sort_order: number } +export interface ChannelTimePricingPeriod { + start_time: string + end_time: string + multiplier: number +} + +export interface ChannelTimePricing { + timezone: string + periods: ChannelTimePricingPeriod[] +} + export interface ChannelModelPricing { id?: number platform: string @@ -34,6 +45,7 @@ export interface ChannelModelPricing { image_output_price: number | null per_request_price: number | null intervals: PricingInterval[] + time_pricing: ChannelTimePricing | null } export interface AccountStatsPricingRule { diff --git a/frontend/src/components/admin/channel/PricingEntryCard.vue b/frontend/src/components/admin/channel/PricingEntryCard.vue index d09e19feb2..f2a0d6502b 100644 --- a/frontend/src/components/admin/channel/PricingEntryCard.vue +++ b/frontend/src/components/admin/channel/PricingEntryCard.vue @@ -87,7 +87,12 @@ @@ -156,6 +161,12 @@ /> + + @@ -238,6 +249,7 @@ import Select from '@/components/common/Select.vue' import Icon from '@/components/icons/Icon.vue' import IntervalRow from './IntervalRow.vue' import ModelTagInput from './ModelTagInput.vue' +import TimePricingSection from './TimePricingSection.vue' import type { PricingFormEntry, IntervalFormEntry } from './types' import { perTokenToMTok, getPlatformTagClass } from './types' import type { BillingMode } from '@/api/admin/channels' @@ -249,8 +261,10 @@ const props = withDefaults(defineProps<{ entry: PricingFormEntry platform?: string hideTokenIntervals?: boolean + enableTimePricing?: boolean }>(), { hideTokenIntervals: false, + enableTimePricing: false, }) const emit = defineEmits<{ diff --git a/frontend/src/components/admin/channel/TimePricingSection.vue b/frontend/src/components/admin/channel/TimePricingSection.vue new file mode 100644 index 0000000000..b2e31ca46d --- /dev/null +++ b/frontend/src/components/admin/channel/TimePricingSection.vue @@ -0,0 +1,161 @@ + + + + + + {{ t('admin.channels.form.timePricing') }} + + + {{ t('admin.channels.form.timezone') }} + + + + + + {{ t('admin.channels.form.addTimePeriod') }} + + + + + + + + {{ t('admin.channels.form.startTime') }} + + + + + + {{ t('admin.channels.form.endTime') }} + + + + + + {{ t('admin.channels.form.multiplier') }} + + + + + + + + + + + + diff --git a/frontend/src/components/admin/channel/__tests__/PricingEntryCard.timePricing.spec.ts b/frontend/src/components/admin/channel/__tests__/PricingEntryCard.timePricing.spec.ts new file mode 100644 index 0000000000..45982256eb --- /dev/null +++ b/frontend/src/components/admin/channel/__tests__/PricingEntryCard.timePricing.spec.ts @@ -0,0 +1,71 @@ +import { shallowMount } from '@vue/test-utils' +import { describe, expect, it, vi } from 'vitest' +import PricingEntryCard from '../PricingEntryCard.vue' +import type { PricingFormEntry } from '../types' + +vi.mock('vue-i18n', async importOriginal => ({ + ...await importOriginal(), + useI18n: () => ({ t: (key: string) => key }), +})) + +function createEntry(billingMode: PricingFormEntry['billing_mode'] = 'token'): PricingFormEntry { + return { + models: [], + billing_mode: billingMode, + input_price: null, + output_price: null, + cache_write_price: null, + cache_read_price: null, + image_input_price: null, + image_output_price: null, + per_request_price: null, + intervals: [], + time_pricing: { + timezone: 'Asia/Shanghai', + periods: [{ start_time: '09:00', end_time: '12:00', multiplier: '2.00' }], + }, + } +} + +describe('PricingEntryCard time pricing visibility', () => { + it('is hidden by default', () => { + const wrapper = shallowMount(PricingEntryCard, { + props: { entry: createEntry() }, + }) + + expect(wrapper.findComponent({ name: 'TimePricingSection' }).exists()).toBe(false) + }) + + it('is shown for token pricing when explicitly enabled', () => { + const wrapper = shallowMount(PricingEntryCard, { + props: { entry: createEntry(), enableTimePricing: true }, + }) + + expect(wrapper.findComponent({ name: 'TimePricingSection' }).exists()).toBe(true) + }) + + it('is hidden for non-token pricing even when explicitly enabled', () => { + const wrapper = shallowMount(PricingEntryCard, { + props: { entry: createEntry('per_request'), enableTimePricing: true }, + }) + + expect(wrapper.findComponent({ name: 'TimePricingSection' }).exists()).toBe(false) + }) + + it('clears time periods when changing billing mode', () => { + const entry = createEntry() + const wrapper = shallowMount(PricingEntryCard, { + props: { entry, enableTimePricing: true }, + }) + + wrapper.findComponent({ name: 'Select' }).vm.$emit('update:modelValue', 'image') + + expect(wrapper.emitted('update')?.[0]?.[0]).toEqual({ + ...entry, + billing_mode: 'image', + intervals: [], + time_pricing: { timezone: 'Asia/Shanghai', periods: [] }, + }) + expect(entry.time_pricing.periods).toHaveLength(1) + }) +}) diff --git a/frontend/src/components/admin/channel/__tests__/TimePricingSection.spec.ts b/frontend/src/components/admin/channel/__tests__/TimePricingSection.spec.ts new file mode 100644 index 0000000000..5ce5348a27 --- /dev/null +++ b/frontend/src/components/admin/channel/__tests__/TimePricingSection.spec.ts @@ -0,0 +1,183 @@ +import { mount } from '@vue/test-utils' +import { beforeEach, describe, expect, it, vi } from 'vitest' +import TimePricingSection from '../TimePricingSection.vue' +import { createDefaultTimePricingForm } from '../types' + +vi.mock('vue-i18n', async importOriginal => ({ + ...await importOriginal(), + useI18n: () => ({ t: (key: string) => key }), +})) + +const SelectStub = { + name: 'Select', + props: ['modelValue', 'options'], + emits: ['update:modelValue'], + template: '', +} + +describe('TimePricingSection', () => { + beforeEach(() => { + vi.clearAllMocks() + }) + + it('adds a neutral period immutably', async () => { + const value = createDefaultTimePricingForm() + const wrapper = mount(TimePricingSection, { + props: { modelValue: value }, + global: { stubs: { Select: SelectStub, Icon: true } }, + }) + + await wrapper.get('[data-testid="add-time-period"]').trigger('click') + + expect(wrapper.emitted('update:modelValue')?.[0]?.[0]).toEqual({ + timezone: 'Asia/Shanghai', + periods: [{ start_time: '', end_time: '', multiplier: '1.00' }], + }) + expect(value.periods).toHaveLength(0) + }) + + it('removes without mutating props', async () => { + const value = { + timezone: 'Asia/Shanghai', + periods: [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: '2.00' }], + } + const wrapper = mount(TimePricingSection, { + props: { modelValue: value }, + global: { stubs: { Select: SelectStub, Icon: true } }, + }) + + const removeButton = wrapper.get('[data-testid="remove-time-period-0"]') + expect(removeButton.attributes('title')).toBe('admin.channels.form.removeTimePeriod') + expect(removeButton.attributes('aria-label')).toBe('admin.channels.form.removeTimePeriod') + await removeButton.trigger('click') + + expect(wrapper.emitted('update:modelValue')?.[0]?.[0]).toEqual({ + timezone: 'Asia/Shanghai', + periods: [], + }) + expect(value.periods).toHaveLength(1) + }) + + it('uses locale-independent 24-hour second inputs', () => { + const value = { + timezone: 'Asia/Shanghai', + periods: [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: '2.00' }], + } + const wrapper = mount(TimePricingSection, { + props: { modelValue: value }, + global: { stubs: { Select: SelectStub, Icon: true } }, + }) + + const timeInputs = wrapper.findAll('input[inputmode="numeric"]') + expect(timeInputs).toHaveLength(2) + expect(timeInputs.every(input => input.attributes('type') === 'text')).toBe(true) + expect(timeInputs.every(input => input.attributes('maxlength') === '8')).toBe(true) + expect(timeInputs.every(input => input.attributes('placeholder') === 'HH:mm:ss')).toBe(true) + }) + + it('updates the timezone and normalizes full-width colons immutably', async () => { + const value = { + timezone: 'Asia/Shanghai', + periods: [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: '2.00' }], + } + const wrapper = mount(TimePricingSection, { + props: { modelValue: value }, + global: { stubs: { Select: SelectStub, Icon: true } }, + }) + + wrapper.getComponent(SelectStub).vm.$emit('update:modelValue', 'Europe/London') + await wrapper.get('input').setValue('10:30:15') + + expect(wrapper.emitted('update:modelValue')?.[0]?.[0]).toEqual({ + ...value, + timezone: 'Europe/London', + }) + expect(wrapper.emitted('update:modelValue')?.[1]?.[0]).toEqual({ + ...value, + periods: [{ ...value.periods[0], start_time: '10:30:15' }], + }) + expect(value).toEqual({ + timezone: 'Asia/Shanghai', + periods: [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: '2.00' }], + }) + }) + + it('normalizes 24:00:00 to 00:00:00', async () => { + const value = { + timezone: 'Asia/Shanghai', + periods: [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: '2.00' }], + } + const wrapper = mount(TimePricingSection, { + props: { modelValue: value }, + global: { stubs: { Select: SelectStub, Icon: true } }, + }) + + await wrapper.get('input').setValue('24:00:00') + + expect(wrapper.emitted('update:modelValue')?.[0]?.[0]).toEqual({ + ...value, + periods: [{ ...value.periods[0], start_time: '00:00:00' }], + }) + }) + + it('formats only a valid positive multiplier on blur', async () => { + const value = { + timezone: 'Asia/Shanghai', + periods: [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: '2' }], + } + const wrapper = mount(TimePricingSection, { + props: { modelValue: value }, + global: { stubs: { Select: SelectStub, Icon: true } }, + }) + + const multiplier = wrapper.get('input[type="number"]') + expect(multiplier.attributes('min')).toBe('0.01') + expect(multiplier.attributes('step')).toBe('0.01') + await multiplier.trigger('blur') + + expect(wrapper.emitted('update:modelValue')?.[0]?.[0]).toEqual({ + ...value, + periods: [{ ...value.periods[0], multiplier: '2.00' }], + }) + + }) + + it.each(['1.234', '1e2', '.5'])('keeps invalid multiplier %s unchanged on blur', async invalid => { + const value = { + timezone: 'Asia/Shanghai', + periods: [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: invalid }], + } + const wrapper = mount(TimePricingSection, { + props: { modelValue: value }, + global: { stubs: { Select: SelectStub, Icon: true } }, + }) + + await wrapper.get('input[type="number"]').trigger('blur') + + expect(wrapper.emitted('update:modelValue')).toBeUndefined() + expect(value.periods[0].multiplier).toBe(invalid) + }) + + it('uses unique input ids across component instances', () => { + const value = { + timezone: 'Asia/Shanghai', + periods: [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: '2.00' }], + } + const wrapper = mount({ + components: { TimePricingSection }, + data: () => ({ first: value, second: value }), + template: ` + + + + + `, + }, { + global: { stubs: { Select: SelectStub, Icon: true } }, + }) + + const ids = wrapper.findAll('input').map(input => input.attributes('id')) + expect(ids).toHaveLength(6) + expect(new Set(ids).size).toBe(ids.length) + }) +}) diff --git a/frontend/src/components/admin/channel/__tests__/types.spec.ts b/frontend/src/components/admin/channel/__tests__/types.spec.ts index 04758012f1..b1ebef252a 100644 --- a/frontend/src/components/admin/channel/__tests__/types.spec.ts +++ b/frontend/src/components/admin/channel/__tests__/types.spec.ts @@ -1,5 +1,14 @@ import { describe, expect, it } from 'vitest' -import { validateIntervals, type IntervalFormEntry } from '../types' +import { + apiTimePricingToForm, + createDefaultTimePricingForm, + formTimePricingToAPI, + validateIntervals, + validateTimePricing, + type IntervalFormEntry, + type TimePricingFormEntry, + type TimePricingPeriodFormEntry, +} from '../types' function makeInterval(over: Partial): IntervalFormEntry { return { @@ -81,3 +90,67 @@ describe('validateIntervals', () => { }) }) }) + +describe('time pricing', () => { + it('uses a disabled Shanghai default', () => { + const form = createDefaultTimePricingForm() + expect(form).toEqual({ timezone: 'Asia/Shanghai', periods: [] }) + expect(formTimePricingToAPI(form)).toBeNull() + }) + + it('round-trips and formats multiplier', () => { + const form = apiTimePricingToForm({ + timezone: 'Asia/Shanghai', + periods: [{ start_time: '09:00', end_time: '12:00', multiplier: 2 }], + }) + expect(form.periods[0]).toEqual({ + start_time: '09:00:00', + end_time: '12:00:00', + multiplier: '2.00', + }) + expect(formTimePricingToAPI(form)).toEqual({ + timezone: 'Asia/Shanghai', + periods: [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: 2 }], + }) + }) + + it.each([ + ['separated', [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: '2.00' }, { start_time: '14:00:00', end_time: '18:00:00', multiplier: '2.00' }], null], + ['adjacent', [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: '2.00' }, { start_time: '12:00:00', end_time: '14:00:00', multiplier: '1.50' }], null], + ['midnight split', [{ start_time: '22:00:00', end_time: '00:00:00', multiplier: '2.00' }, { start_time: '00:00:00', end_time: '02:00:00', multiplier: '2.00' }], null], + ['overlap by one second', [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: '2.00' }, { start_time: '11:59:59', end_time: '14:00:00', multiplier: '2.00' }], 'overlap'], + ['cross midnight', [{ start_time: '22:00:00', end_time: '02:00:00', multiplier: '2.00' }], 'range'], + ['equal midnight', [{ start_time: '00:00:00', end_time: '00:00:00', multiplier: '2.00' }], 'range'], + ['missing seconds', [{ start_time: '09:00', end_time: '12:00', multiplier: '2.00' }], 'format'], + ['zero', [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: '0.00' }], 'multiplier'], + ['three decimals', [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: '1.001' }], 'multiplier'], + ])('%s', (_name, periods, errorKey) => { + const result = validateTimePricing({ + timezone: 'Asia/Shanghai', + periods: periods as TimePricingPeriodFormEntry[], + }, t) + if (errorKey === null) expect(result).toBeNull() + else expect(result).toContain(String(errorKey)) + }) + + it('rejects non-IANA timezone', () => { + expect(validateTimePricing({ + timezone: 'UTC+8', + periods: [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: '2.00' }], + }, t)).toContain('timezone') + }) + + it.each([ + ['missing', undefined], + ['blank', ' '], + ])('rejects a %s timezone without throwing during conversion', (_name, timezone) => { + const form = { + timezone, + periods: [{ start_time: '09:00:00', end_time: '12:00:00', multiplier: '2.00' }], + } as unknown as TimePricingFormEntry + + expect(validateTimePricing(form, t)).toContain('timezone') + expect(() => formTimePricingToAPI(form)).not.toThrow() + expect(formTimePricingToAPI(form)?.timezone).toBe('') + }) +}) diff --git a/frontend/src/components/admin/channel/types.ts b/frontend/src/components/admin/channel/types.ts index 5c43b7d349..8809eff072 100644 --- a/frontend/src/components/admin/channel/types.ts +++ b/frontend/src/components/admin/channel/types.ts @@ -1,4 +1,4 @@ -import type { BillingMode, PricingInterval } from '@/api/admin/channels' +import type { BillingMode, ChannelTimePricing, PricingInterval } from '@/api/admin/channels' type TranslateFn = (key: string, params?: Record) => string @@ -25,6 +25,134 @@ export interface PricingFormEntry { image_output_price: number | string | null per_request_price: number | string | null intervals: IntervalFormEntry[] + time_pricing: TimePricingFormEntry +} + +export interface TimePricingPeriodFormEntry { + start_time: string + end_time: string + multiplier: number | string +} + +export interface TimePricingFormEntry { + timezone: string + periods: TimePricingPeriodFormEntry[] +} + +export const DEFAULT_TIME_PRICING_TIMEZONE = 'Asia/Shanghai' + +const CLOCK_TIME = /^(?:[01]\d|2[0-3]):[0-5]\d:[0-5]\d$/ +const LEGACY_CLOCK_TIME = /^(?:[01]\d|2[0-3]):[0-5]\d$/ +const TWO_DECIMAL_MULTIPLIER = /^\d+(?:\.\d{1,2})?$/ + +export function isValidTimePricingMultiplier(value: number | string): boolean { + const multiplier = String(value) + const numericValue = Number(multiplier) + return TWO_DECIMAL_MULTIPLIER.test(multiplier) && + Number.isFinite(numericValue) && numericValue > 0 +} + +export const COMMON_TIMEZONES = [ + 'UTC', 'Asia/Shanghai', 'Asia/Tokyo', 'Asia/Seoul', 'Asia/Singapore', 'Asia/Kolkata', + 'Australia/Sydney', 'Europe/London', 'Europe/Paris', 'Europe/Berlin', + 'America/New_York', 'America/Chicago', 'America/Denver', 'America/Los_Angeles', + 'America/Toronto', 'America/Sao_Paulo', 'Pacific/Auckland', 'Pacific/Honolulu', +] + +export function createDefaultTimePricingForm(): TimePricingFormEntry { + return { timezone: DEFAULT_TIME_PRICING_TIMEZONE, periods: [] } +} + +export function apiTimePricingToForm(value: ChannelTimePricing | null | undefined): TimePricingFormEntry { + if (!value) return createDefaultTimePricingForm() + return { + timezone: value.timezone || DEFAULT_TIME_PRICING_TIMEZONE, + periods: (value.periods || []).map(period => ({ + start_time: LEGACY_CLOCK_TIME.test(period.start_time) ? `${period.start_time}:00` : period.start_time, + end_time: LEGACY_CLOCK_TIME.test(period.end_time) ? `${period.end_time}:00` : period.end_time, + multiplier: Number(period.multiplier).toFixed(2), + })), + } +} + +export function formTimePricingToAPI(value: TimePricingFormEntry | null | undefined): ChannelTimePricing | null { + if (!value?.periods?.length) return null + const timezone = typeof value.timezone === 'string' ? value.timezone.trim() : '' + return { + timezone, + periods: value.periods.map(period => ({ + start_time: period.start_time, + end_time: period.end_time, + multiplier: Number(period.multiplier), + })), + } +} + +function timeToSeconds(time: string, isEnd: boolean): number { + if (isEnd && time === '00:00:00') return 24 * 60 * 60 + const [hours, minutes, seconds] = time.split(':').map(Number) + return hours * 60 * 60 + minutes * 60 + seconds +} + +function timePricingValidationMessage(t: TranslateFn, key: string): string { + return t(`admin.channels.timePricingValidation.${key}`) +} + +export function validateTimePricing(value: TimePricingFormEntry, t: TranslateFn): string | null { + if (!value?.periods?.length) return null + + if (typeof value.timezone !== 'string' || value.timezone.trim() === '') { + return timePricingValidationMessage(t, 'timezone') + } + const timezone = value.timezone.trim() + + try { + new Intl.DateTimeFormat('en-US', { timeZone: timezone }) + } catch { + return timePricingValidationMessage(t, 'timezone') + } + + const periods = [] as { start: number, end: number }[] + for (const period of value.periods) { + if (!CLOCK_TIME.test(period.start_time) || !CLOCK_TIME.test(period.end_time)) { + return timePricingValidationMessage(t, 'format') + } + if (period.start_time === period.end_time) { + return timePricingValidationMessage(t, 'range') + } + + const start = timeToSeconds(period.start_time, false) + const end = timeToSeconds(period.end_time, true) + if (start >= end) return timePricingValidationMessage(t, 'range') + + if (!isValidTimePricingMultiplier(period.multiplier)) { + return timePricingValidationMessage(t, 'multiplier') + } + periods.push({ start, end }) + } + + const sorted = [...periods].sort((a, b) => a.start - b.start) + for (let i = 1; i < sorted.length; i++) { + if (sorted[i].start < sorted[i - 1].end) { + return timePricingValidationMessage(t, 'overlap') + } + } + return null +} + +export function formatTimezoneOffset(timezone: string, at = new Date()): string { + try { + const part = new Intl.DateTimeFormat('en-US', { + timeZone: timezone, + timeZoneName: 'shortOffset', + }).formatToParts(at).find(item => item.type === 'timeZoneName')?.value + if (!part || part === 'GMT') return 'UTC+00:00' + const match = /^GMT([+-])(\d{1,2})(?::(\d{2}))?$/.exec(part) + if (!match) return '' + return `UTC${match[1]}${match[2].padStart(2, '0')}:${match[3] || '00'}` + } catch { + return '' + } } // 价格转换:后端存 per-token,前端显示 per-MTok ($/1M tokens) diff --git a/frontend/src/i18n/locales/en/admin/channels.ts b/frontend/src/i18n/locales/en/admin/channels.ts index d68ea5060d..bed17bc030 100644 --- a/frontend/src/i18n/locales/en/admin/channels.ts +++ b/frontend/src/i18n/locales/en/admin/channels.ts @@ -80,6 +80,13 @@ export default { perRequestPrice: 'per-request price' } }, + timePricingValidation: { + timezone: 'Select a valid IANA time zone', + format: 'Start and end times must use HH:mm:ss format', + range: 'Start time must be earlier than end time; split ranges across midnight', + overlap: 'Time periods must not overlap', + multiplier: 'Multiplier must be greater than 0 with at most two decimal places' + }, deleteConfirm: 'Are you sure you want to delete channel "{name}"? This cannot be undone.', columns: { name: 'Name', @@ -122,6 +129,13 @@ export default { imageOutputPrice: 'Image Output Price', pricePlaceholder: 'Default', intervals: 'Context Intervals (optional)', + timePricing: 'Time-based pricing (optional)', + timezone: 'Time zone', + addTimePeriod: 'Add period', + startTime: 'Start time', + endTime: 'End time', + multiplier: 'Multiplier', + removeTimePeriod: 'Remove period', minTokens: 'Min', maxTokens: 'Max', inclusive: '(inclusive)', diff --git a/frontend/src/i18n/locales/zh/admin/channels.ts b/frontend/src/i18n/locales/zh/admin/channels.ts index 808e78fa88..12cb22d565 100644 --- a/frontend/src/i18n/locales/zh/admin/channels.ts +++ b/frontend/src/i18n/locales/zh/admin/channels.ts @@ -80,6 +80,13 @@ export default { perRequestPrice: '单次价格' } }, + timePricingValidation: { + timezone: '请选择有效的 IANA 时区', + format: '开始时间和结束时间必须使用 HH:mm:ss 格式', + range: '开始时间必须早于结束时间;跨午夜请拆分为两个时间段', + overlap: '时间段不能重叠', + multiplier: '倍率必须大于 0,且最多保留两位小数' + }, deleteConfirm: '确定要删除渠道「{name}」吗?此操作不可撤销。', columns: { name: '名称', @@ -122,6 +129,13 @@ export default { imageOutputPrice: '图片输出价格', pricePlaceholder: '默认', intervals: '上下文区间定价(可选)', + timePricing: '时间段定价(可选)', + timezone: '时区', + addTimePeriod: '添加时间段', + startTime: '开始时间', + endTime: '结束时间', + multiplier: '倍率', + removeTimePeriod: '删除时间段', minTokens: '最小', maxTokens: '最大', inclusive: '(含)', diff --git a/frontend/src/views/admin/ChannelsView.vue b/frontend/src/views/admin/ChannelsView.vue index 9dceb5362a..d431cb1293 100644 --- a/frontend/src/views/admin/ChannelsView.vue +++ b/frontend/src/views/admin/ChannelsView.vue @@ -447,6 +447,7 @@ :key="idx" :entry="entry" :platform="section.platform" + enable-time-pricing @update="updatePricingEntry(sIdx, idx, $event)" @remove="removePricingEntry(sIdx, idx)" /> @@ -632,7 +633,7 @@ import { extractApiErrorMessage } from '@/utils/apiError' import { adminAPI } from '@/api/admin' import type { Channel, ChannelModelPricing, CreateChannelRequest, UpdateChannelRequest, AccountStatsPricingRule } from '@/api/admin/channels' import type { PricingFormEntry } from '@/components/admin/channel/types' -import { mTokToPerToken, perTokenToMTok, apiIntervalsToForm, formIntervalsToAPI, findModelConflict, validateIntervals } from '@/components/admin/channel/types' +import { apiIntervalsToForm, apiTimePricingToForm, createDefaultTimePricingForm, findModelConflict, formIntervalsToAPI, formTimePricingToAPI, mTokToPerToken, perTokenToMTok, validateIntervals, validateTimePricing } from '@/components/admin/channel/types' import type { AdminGroup, GroupPlatform } from '@/types' import type { Column } from '@/components/common/types' import { platformTextClass, platformBadgeLightClass } from '@/utils/platformColors' @@ -859,7 +860,8 @@ function addPricingEntry(sectionIdx: number) { image_input_price: null, image_output_price: null, per_request_price: null, - intervals: [] + intervals: [], + time_pricing: createDefaultTimePricingForm() }) } @@ -892,7 +894,8 @@ async function syncLatestModels(sectionIdx: number) { image_input_price: null, image_output_price: null, per_request_price: null, - intervals: [] + intervals: [], + time_pricing: createDefaultTimePricingForm() }) appStore.showSuccess(t('admin.channels.form.syncModelsSuccess', { count: newModels.length })) } catch (error) { @@ -957,7 +960,8 @@ function addRulePricingEntry(sectionIdx: number, ruleIndex: number) { image_input_price: null, image_output_price: null, per_request_price: null, - intervals: [] + intervals: [], + time_pricing: createDefaultTimePricingForm() }) } @@ -1073,7 +1077,8 @@ function accountStatsRulesToAPI(): AccountStatsPricingRule[] { image_input_price: mTokToPerToken(p.image_input_price), image_output_price: mTokToPerToken(p.image_output_price), per_request_price: p.per_request_price != null && p.per_request_price !== '' ? Number(p.per_request_price) : null, - intervals: formIntervalsToAPI(p.intervals || []) + intervals: formIntervalsToAPI(p.intervals || []), + time_pricing: null })) }) } @@ -1114,7 +1119,8 @@ function formToAPI(): { group_ids: number[], model_pricing: ChannelModelPricing[ image_input_price: mTokToPerToken(entry.image_input_price), image_output_price: mTokToPerToken(entry.image_output_price), per_request_price: entry.per_request_price != null && entry.per_request_price !== '' ? Number(entry.per_request_price) : null, - intervals: formIntervalsToAPI(entry.intervals || []) + intervals: formIntervalsToAPI(entry.intervals || []), + time_pricing: formTimePricingToAPI(entry.time_pricing) }) } } @@ -1211,7 +1217,8 @@ function apiToForm(channel: Channel): PlatformSection[] { image_input_price: perTokenToMTok(p.image_input_price), image_output_price: perTokenToMTok(p.image_output_price), per_request_price: p.per_request_price, - intervals: apiIntervalsToForm(p.intervals || []) + intervals: apiIntervalsToForm(p.intervals || []), + time_pricing: apiTimePricingToForm(p.time_pricing) } as PricingFormEntry)) // Read web_search_emulation from features_config @@ -1400,7 +1407,8 @@ function distributeRulesToPlatforms(apiRules: AccountStatsPricingRule[]) { image_input_price: perTokenToMTok(p.image_input_price), image_output_price: perTokenToMTok(p.image_output_price), per_request_price: p.per_request_price, - intervals: apiIntervalsToForm(p.intervals || []) + intervals: apiIntervalsToForm(p.intervals || []), + time_pricing: createDefaultTimePricingForm() } as PricingFormEntry)) } section.account_stats_pricing_rules.push(formRule) @@ -1523,6 +1531,20 @@ async function handleSubmit() { } } + // 校验时间段定价,并切换到对应平台便于修正 + for (const section of form.platforms.filter(s => s.enabled)) { + for (const entry of section.model_pricing) { + const timePricingError = validateTimePricing(entry.time_pricing, t) + if (timePricingError) { + const platformLabel = t('admin.groups.platforms.' + section.platform, section.platform) + const modelLabel = entry.models.join(', ') || t('admin.channels.form.unnamed') + appStore.showError(`${platformLabel} - ${modelLabel}: ${timePricingError}`) + activeTab.value = section.platform + return + } + } + } + const { group_ids, model_pricing, model_mapping, features_config } = formToAPI() submitting.value = true diff --git a/frontend/src/views/admin/GroupsView.vue b/frontend/src/views/admin/GroupsView.vue index 904b8926d4..9a427e7c64 100644 --- a/frontend/src/views/admin/GroupsView.vue +++ b/frontend/src/views/admin/GroupsView.vue @@ -4429,6 +4429,7 @@ import PricingEntryCard from "@/components/admin/channel/PricingEntryCard.vue"; import type { PricingFormEntry } from "@/components/admin/channel/types"; import { apiIntervalsToForm, + createDefaultTimePricingForm, formIntervalsToAPI, mTokToPerToken, perTokenToMTok, @@ -4499,6 +4500,7 @@ const emptyGroupPricing = (): PricingFormEntry => ({ image_output_price: null, per_request_price: null, intervals: [], + time_pricing: createDefaultTimePricingForm(), }); const addGroupPricing = (entries: PricingFormEntry[]) => @@ -4518,6 +4520,7 @@ const groupPricingFromAPI = ( image_output_price: perTokenToMTok(entry.image_output_price), per_request_price: entry.per_request_price, intervals: apiIntervalsToForm(entry.intervals || []), + time_pricing: createDefaultTimePricingForm(), })); const groupPricingToAPI = ( @@ -4541,6 +4544,7 @@ const groupPricingToAPI = ( entry.billing_mode === "token" ? [] : formIntervalsToAPI(entry.intervals || []), + time_pricing: null, })); const { t } = useI18n();