mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 17:27:58 +08:00
Merge pull request #5737 from lyen1688/feat/channel-time-pricing
功能:支持渠道模型分时倍率定价(支持峰谷定价)
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
@@ -1446,12 +1446,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
|
||||
@@ -1463,10 +1461,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 侧选号循环的一次利润门终检否决:把账号
|
||||
@@ -1729,6 +1730,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
|
||||
@@ -2057,19 +2059,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 {
|
||||
@@ -2155,6 +2175,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 已清除标记)。
|
||||
@@ -2222,7 +2243,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{
|
||||
|
||||
@@ -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 必须使用自己的定价时刻")
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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())
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 共用。
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)))
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
@@ -64,7 +64,7 @@ func TestCalculateOpenAIRecordUsageCost_EmptyCandidatesIsPricingUnavailable(t *t
|
||||
|
||||
_, err := svc.calculateOpenAIRecordUsageCost(
|
||||
context.Background(), nil, apiKey, nil,
|
||||
1.0, 1.0, 1.0, 1.0, UsageTokens{InputTokens: 100}, "", nil,
|
||||
1.0, 1.0, 1.0, 1.0, UsageTokens{InputTokens: 100}, "", nil, time.Time{},
|
||||
)
|
||||
require.Error(t, err)
|
||||
require.True(t, isUsagePricingUnavailableError(err),
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
@@ -225,6 +226,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
|
||||
tokens,
|
||||
serviceTier,
|
||||
longContextBillingGate,
|
||||
pricingAt,
|
||||
)
|
||||
if err != nil {
|
||||
if !isUsagePricingUnavailableError(err) {
|
||||
@@ -257,7 +259,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
|
||||
responseModels := s.filterCNProviderBillingModelCandidates(ctx, account, apiKey, 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
|
||||
// 内部做渠道定价判断时使用的模型,且"首候选有渠道价"必然意味着首候选就是实际
|
||||
@@ -506,6 +508,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 {
|
||||
@@ -555,6 +558,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost(
|
||||
apiKey,
|
||||
candidate,
|
||||
multiplier,
|
||||
pricingAt,
|
||||
tokens,
|
||||
serviceTier,
|
||||
longContextBillingGate,
|
||||
@@ -647,6 +651,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageTokenCost(
|
||||
apiKey *APIKey,
|
||||
billingModel string,
|
||||
multiplier float64,
|
||||
pricingAt time.Time,
|
||||
tokens UsageTokens,
|
||||
serviceTier string,
|
||||
longContextBillingGate *bool,
|
||||
@@ -655,7 +660,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,
|
||||
})
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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';
|
||||
@@ -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 {
|
||||
|
||||
@@ -87,7 +87,12 @@
|
||||
</label>
|
||||
<Select
|
||||
:modelValue="entry.billing_mode"
|
||||
@update:modelValue="emit('update', { ...entry, billing_mode: $event as BillingMode, intervals: [] })"
|
||||
@update:modelValue="emit('update', {
|
||||
...entry,
|
||||
billing_mode: $event as BillingMode,
|
||||
intervals: [],
|
||||
time_pricing: { ...entry.time_pricing, periods: [] },
|
||||
})"
|
||||
:options="billingModeOptions"
|
||||
class="mt-1"
|
||||
/>
|
||||
@@ -156,6 +161,12 @@
|
||||
/>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<TimePricingSection
|
||||
v-if="enableTimePricing"
|
||||
:model-value="entry.time_pricing"
|
||||
@update:model-value="emit('update', { ...entry, time_pricing: $event })"
|
||||
/>
|
||||
</div>
|
||||
|
||||
<!-- Per-request mode -->
|
||||
@@ -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<{
|
||||
|
||||
@@ -0,0 +1,161 @@
|
||||
<template>
|
||||
<section class="mt-3 border-t border-gray-200 pt-3 dark:border-dark-600">
|
||||
<div class="flex flex-col gap-2 sm:flex-row sm:items-end sm:justify-between">
|
||||
<div class="min-w-0 flex-1 sm:max-w-sm">
|
||||
<label class="block text-xs font-medium text-gray-500 dark:text-gray-400">
|
||||
{{ t('admin.channels.form.timePricing') }}
|
||||
</label>
|
||||
<label class="mt-2 block text-xs text-gray-400">
|
||||
{{ t('admin.channels.form.timezone') }}
|
||||
</label>
|
||||
<Select
|
||||
:model-value="modelValue.timezone"
|
||||
:options="timezoneOptions"
|
||||
:aria-label="t('admin.channels.form.timezone')"
|
||||
searchable
|
||||
creatable
|
||||
class="mt-1 w-full"
|
||||
@update:model-value="updateTimezone"
|
||||
/>
|
||||
</div>
|
||||
<button
|
||||
type="button"
|
||||
class="self-start text-xs text-primary-600 hover:text-primary-700 sm:self-end sm:pb-2"
|
||||
data-testid="add-time-period"
|
||||
@click="addPeriod"
|
||||
>
|
||||
+ {{ t('admin.channels.form.addTimePeriod') }}
|
||||
</button>
|
||||
</div>
|
||||
|
||||
<div v-if="modelValue.periods.length > 0" class="mt-3 space-y-3">
|
||||
<div
|
||||
v-for="(period, index) in modelValue.periods"
|
||||
:key="index"
|
||||
class="grid grid-cols-1 gap-2 border-t border-gray-200 pt-3 sm:grid-cols-[minmax(0,1fr)_minmax(0,1fr)_minmax(0,1fr)_2rem] sm:items-end dark:border-dark-600"
|
||||
>
|
||||
<div class="min-w-0">
|
||||
<label :for="`${inputIdPrefix}-start-${index}`" class="block text-xs text-gray-400">
|
||||
{{ t('admin.channels.form.startTime') }}
|
||||
</label>
|
||||
<input
|
||||
:id="`${inputIdPrefix}-start-${index}`"
|
||||
:value="period.start_time"
|
||||
type="text"
|
||||
inputmode="numeric"
|
||||
maxlength="8"
|
||||
placeholder="HH:mm:ss"
|
||||
pattern="[0-9]{2}:[0-9]{2}:[0-9]{2}"
|
||||
autocomplete="off"
|
||||
class="input mt-1 w-full text-sm"
|
||||
@input="updatePeriod(index, 'start_time', normalizeClockTime(($event.target as HTMLInputElement).value))"
|
||||
/>
|
||||
</div>
|
||||
<div class="min-w-0">
|
||||
<label :for="`${inputIdPrefix}-end-${index}`" class="block text-xs text-gray-400">
|
||||
{{ t('admin.channels.form.endTime') }}
|
||||
</label>
|
||||
<input
|
||||
:id="`${inputIdPrefix}-end-${index}`"
|
||||
:value="period.end_time"
|
||||
type="text"
|
||||
inputmode="numeric"
|
||||
maxlength="8"
|
||||
placeholder="HH:mm:ss"
|
||||
pattern="[0-9]{2}:[0-9]{2}:[0-9]{2}"
|
||||
autocomplete="off"
|
||||
class="input mt-1 w-full text-sm"
|
||||
@input="updatePeriod(index, 'end_time', normalizeClockTime(($event.target as HTMLInputElement).value))"
|
||||
/>
|
||||
</div>
|
||||
<div class="min-w-0">
|
||||
<label :for="`${inputIdPrefix}-multiplier-${index}`" class="block text-xs text-gray-400">
|
||||
{{ t('admin.channels.form.multiplier') }}
|
||||
</label>
|
||||
<input
|
||||
:id="`${inputIdPrefix}-multiplier-${index}`"
|
||||
:value="period.multiplier"
|
||||
type="number"
|
||||
min="0.01"
|
||||
step="0.01"
|
||||
class="input mt-1 w-full text-sm"
|
||||
@input="updatePeriod(index, 'multiplier', ($event.target as HTMLInputElement).value)"
|
||||
@blur="formatMultiplier(index, ($event.target as HTMLInputElement).value)"
|
||||
/>
|
||||
</div>
|
||||
<button
|
||||
type="button"
|
||||
class="flex h-8 w-8 items-center justify-center rounded text-gray-400 hover:text-red-500"
|
||||
:title="t('admin.channels.form.removeTimePeriod')"
|
||||
:aria-label="t('admin.channels.form.removeTimePeriod')"
|
||||
:data-testid="`remove-time-period-${index}`"
|
||||
@click="removePeriod(index)"
|
||||
>
|
||||
<Icon name="trash" size="sm" />
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { getCurrentInstance } from 'vue'
|
||||
import { useI18n } from 'vue-i18n'
|
||||
import Select from '@/components/common/Select.vue'
|
||||
import Icon from '@/components/icons/Icon.vue'
|
||||
import {
|
||||
COMMON_TIMEZONES,
|
||||
formatTimezoneOffset,
|
||||
isValidTimePricingMultiplier,
|
||||
type TimePricingFormEntry,
|
||||
type TimePricingPeriodFormEntry,
|
||||
} from './types'
|
||||
|
||||
const { t } = useI18n()
|
||||
|
||||
const props = defineProps<{ modelValue: TimePricingFormEntry }>()
|
||||
const emit = defineEmits<{ 'update:modelValue': [value: TimePricingFormEntry] }>()
|
||||
const inputIdPrefix = `time-pricing-${getCurrentInstance()?.uid}`
|
||||
|
||||
const timezoneOptions = COMMON_TIMEZONES.map(value => {
|
||||
const offset = formatTimezoneOffset(value)
|
||||
return { value, label: offset ? `${value} (${offset})` : value }
|
||||
})
|
||||
|
||||
function updateTimezone(value: string | number | boolean | null) {
|
||||
emit('update:modelValue', { ...props.modelValue, timezone: String(value ?? '') })
|
||||
}
|
||||
|
||||
function normalizeClockTime(value: string): string {
|
||||
const normalized = value.replace(/:/g, ':')
|
||||
return normalized === '24:00:00' ? '00:00:00' : normalized
|
||||
}
|
||||
|
||||
function addPeriod() {
|
||||
emit('update:modelValue', {
|
||||
...props.modelValue,
|
||||
periods: [
|
||||
...props.modelValue.periods,
|
||||
{ start_time: '', end_time: '', multiplier: '1.00' },
|
||||
],
|
||||
})
|
||||
}
|
||||
|
||||
function updatePeriod(index: number, field: keyof TimePricingPeriodFormEntry, value: string) {
|
||||
const periods = props.modelValue.periods.map((period, current) =>
|
||||
current === index ? { ...period, [field]: value } : period)
|
||||
emit('update:modelValue', { ...props.modelValue, periods })
|
||||
}
|
||||
|
||||
function formatMultiplier(index: number, value: string) {
|
||||
if (!isValidTimePricingMultiplier(value)) return
|
||||
updatePeriod(index, 'multiplier', Number(value).toFixed(2))
|
||||
}
|
||||
|
||||
function removePeriod(index: number) {
|
||||
emit('update:modelValue', {
|
||||
...props.modelValue,
|
||||
periods: props.modelValue.periods.filter((_period, current) => current !== index),
|
||||
})
|
||||
}
|
||||
</script>
|
||||
@@ -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<typeof import('vue-i18n')>(),
|
||||
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)
|
||||
})
|
||||
})
|
||||
@@ -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<typeof import('vue-i18n')>(),
|
||||
useI18n: () => ({ t: (key: string) => key }),
|
||||
}))
|
||||
|
||||
const SelectStub = {
|
||||
name: 'Select',
|
||||
props: ['modelValue', 'options'],
|
||||
emits: ['update:modelValue'],
|
||||
template: '<button type="button" data-testid="timezone-select" />',
|
||||
}
|
||||
|
||||
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: `
|
||||
<div>
|
||||
<TimePricingSection v-model="first" />
|
||||
<TimePricingSection v-model="second" />
|
||||
</div>
|
||||
`,
|
||||
}, {
|
||||
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)
|
||||
})
|
||||
})
|
||||
@@ -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>): 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('')
|
||||
})
|
||||
})
|
||||
|
||||
@@ -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, unknown>) => 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)
|
||||
|
||||
@@ -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)',
|
||||
|
||||
@@ -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: '(含)',
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -4441,6 +4441,7 @@ import PricingEntryCard from "@/components/admin/channel/PricingEntryCard.vue";
|
||||
import type { PricingFormEntry } from "@/components/admin/channel/types";
|
||||
import {
|
||||
apiIntervalsToForm,
|
||||
createDefaultTimePricingForm,
|
||||
formIntervalsToAPI,
|
||||
mTokToPerToken,
|
||||
perTokenToMTok,
|
||||
@@ -4511,6 +4512,7 @@ const emptyGroupPricing = (): PricingFormEntry => ({
|
||||
image_output_price: null,
|
||||
per_request_price: null,
|
||||
intervals: [],
|
||||
time_pricing: createDefaultTimePricingForm(),
|
||||
});
|
||||
|
||||
const addGroupPricing = (entries: PricingFormEntry[]) =>
|
||||
@@ -4530,6 +4532,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 = (
|
||||
@@ -4553,6 +4556,7 @@ const groupPricingToAPI = (
|
||||
entry.billing_mode === "token"
|
||||
? []
|
||||
: formIntervalsToAPI(entry.intervals || []),
|
||||
time_pricing: null,
|
||||
}));
|
||||
|
||||
const { t } = useI18n();
|
||||
|
||||
Reference in New Issue
Block a user