功能:支持渠道模型分时倍率定价

This commit is contained in:
lyen1688
2026-08-17 19:45:07 +08:00
parent e330c243a8
commit 9f24a55305
40 changed files with 2297 additions and 131 deletions
@@ -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
// ---------------------------------------------------------------------------
@@ -1420,12 +1420,10 @@ const (
// 由 BeforeTurn 在每个 turn 开始时冻结,AfterTurn 的用量提交读取它;turn 在
// 连接内串行推进,互斥锁只为跨用量提交 goroutine 的读取安全。
//
// 零值语义(重要):ws_v2 passthrough ingress 只实现了 AfterTurn,没有任何
// turn 起始回调,BeforeTurn 永远不会被调用。此时本值保持零,RecordUsage 经
// openAIUsagePricingAt 回退到记录时刻——与引入分组利润控制前的基线一致。
// 绝不能用建连时刻初始化:那会把透传连接的所有 turn 钉死在建连时的高峰因子,
// 客户端只要峰前一分钟建连并保活,整条连接就能全程按谷价结算,正是利润控制
// 想堵的漏洞。透传 ingress 目前不做 turn 级利润复核,只有建连时的准入门。
// ws_v2 passthrough ingress 没有 BeforeTurn,因此本值会保持零;AfterTurn 必须
// 以 TurnStarted 已记录的所属 turn 开始时刻为回退,而不是用建连或记录时刻。
// 这样每个 passthrough turn 都按自己的开始时刻计价,但不改变其仅在建连时执行
// 准入门、没有 turn 级利润复核的既有行为。
type openAIWSTurnPricing struct {
mu sync.Mutex
at time.Time
@@ -1437,10 +1435,13 @@ func (p *openAIWSTurnPricing) freeze(at time.Time) {
p.mu.Unlock()
}
func (p *openAIWSTurnPricing) current() time.Time {
func (p *openAIWSTurnPricing) currentOr(fallback time.Time) time.Time {
p.mu.Lock()
defer p.mu.Unlock()
return p.at
if !p.at.IsZero() {
return p.at
}
return fallback
}
// recordOpenAIProfitVeto 记录 OpenAI 侧选号循环的一次利润门终检否决:把账号
@@ -1703,6 +1704,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "missing first response.create message")
return
}
firstTurnStartedAt := time.Now()
if msgType != coderws.MessageText && msgType != coderws.MessageBinary {
closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, "unsupported websocket message type")
return
@@ -2031,19 +2033,37 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
maxReasoningEffort, reasoningEffortMappings, _ := openAIReasoningEffortPolicyForRequest(c, apiKey)
var requestPayloadHash string
var turnStartsMu sync.Mutex
turnStarts := make(map[int]time.Time, 4)
recordTurnStart := func(turn int, startedAt time.Time) {
if turn <= 0 || startedAt.IsZero() {
return
}
turnStartsMu.Lock()
turnStarts[turn] = startedAt
turnStartsMu.Unlock()
}
getTurnStart := func(turn int) time.Time {
turnStartsMu.Lock()
startedAt := turnStarts[turn]
delete(turnStarts, turn)
turnStartsMu.Unlock()
return startedAt
}
// Passthrough rejects overlapping response.create frames, so one immutable
// turn-tagged slot preserves the exact mapping used for the in-flight request.
var turnChannelMapping atomic.Pointer[openAIWSTurnChannelMappingSnapshot]
turnChannelMapping.Store(&openAIWSTurnChannelMappingSnapshot{turn: 1, mapping: channelMappingWS})
// turn 级定价:BeforeTurn 重新冻结 pricingAt 并按最新门复核当前账号,
// AfterTurn 的计费读取所属 turn 的时刻。零值起步的语义见
// openAIWSTurnPricing 的注释——绝不能用建连时刻初始化。
// turn 级定价:BeforeTurn 重新冻结 pricingAt 并按最新门复核当前账号;
// passthrough 没有 BeforeTurn 时,AfterTurn 回退到 TurnStarted 的所属 turn 时刻。
var turnPricing openAIWSTurnPricing
hooks := &service.OpenAIWSIngressHooks{
ClientLifecycleContext: clientLifecycleCtx,
InitialRequestModel: reqModel,
InitialTurnStartedAt: firstTurnStartedAt,
MaxReasoningEffort: maxReasoningEffort,
ReasoningEffortMappings: reasoningEffortMappings,
TurnStarted: recordTurnStart,
BeforeRequest: func(turn int, payload []byte, originalModel string) error {
c.Set(securityAuditWSTurnContextKey, turn)
if turn == 1 {
@@ -2129,6 +2149,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
return nil
},
AfterTurn: func(turn int, result *service.OpenAIForwardResult, turnErr error) {
turnStart := getTurnStart(turn)
// F1: cyber 标记按 turn 生命周期清理——defer 保证任意早返回路径都执行;
// CyberBlocked 必须在 submit 前同步预捕获(task 闭包由 worker 池异步执行,
// 届时 defer 已清除标记)。
@@ -2196,7 +2217,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
upstreamEndpoint := resolveOpenAIUpstreamEndpoint(c, account, result)
quotaPlatform := service.QuotaPlatform(c.Request.Context(), apiKey)
sessionID := service.ExtractClientSessionID(c)
turnRecordPricingAt := turnPricing.current()
turnRecordPricingAt := turnPricing.currentOr(turnStart)
cyberBlocked := service.GetOpsCyberPolicy(c) != nil
h.submitOpenAIUsageRecordTask(ctx, result, func(taskCtx context.Context) {
if err := h.gatewayService.RecordUsage(taskCtx, &service.OpenAIRecordUsageInput{
@@ -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())
})
})
}
}
+6
View File
@@ -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
+25 -1
View File
@@ -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) {
+35 -15
View File
@@ -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
}
+45 -9
View File
@@ -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
// ---------------------------------------------------------------------------
+15 -1
View File
@@ -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)
}
@@ -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
@@ -224,6 +225,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
tokens,
serviceTier,
longContextBillingGate,
pricingAt,
)
if err != nil {
if !isUsagePricingUnavailableError(err) {
@@ -256,7 +258,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
responseModels := usageBillingModelCandidates(responseModel)
responseCost, responseErr := s.calculateOpenAIRecordUsageCost(
ctx, result, apiKey, responseModels, multiplier, imageMultiplier,
videoMultiplier, baseMultiplier, tokens, serviceTier, longContextBillingGate,
videoMultiplier, baseMultiplier, tokens, serviceTier, longContextBillingGate, pricingAt,
)
// 基线定价源以 baselineBillingModel 为准:它正是 calculateOpenAIRecordUsageCost
// 内部做渠道定价判断时使用的模型,且"首候选有渠道价"必然意味着首候选就是实际
@@ -505,6 +507,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost(
tokens UsageTokens,
serviceTier string,
longContextBillingGate *bool,
pricingAt time.Time,
) (*CostBreakdown, error) {
billingModel := firstUsageBillingModel(billingModels)
if result != nil && result.WebSearchCalls > 0 {
@@ -554,6 +557,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageCost(
apiKey,
candidate,
multiplier,
pricingAt,
tokens,
serviceTier,
longContextBillingGate,
@@ -643,6 +647,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageTokenCost(
apiKey *APIKey,
billingModel string,
multiplier float64,
pricingAt time.Time,
tokens UsageTokens,
serviceTier string,
longContextBillingGate *bool,
@@ -651,7 +656,7 @@ func (s *OpenAIGatewayService) calculateOpenAIRecordUsageTokenCost(
gid := apiKey.Group.ID
return s.billingService.CalculateCostUnified(CostInput{
Ctx: ctx, Model: billingModel, GroupID: &gid, Group: apiKey.Group,
Tokens: tokens, RequestCount: 1, RateMultiplier: multiplier,
Tokens: tokens, RequestCount: 1, RateMultiplier: multiplier, PricingAt: pricingAt,
ServiceTier: serviceTier, Resolver: s.resolver,
LongContextBillingEnabled: longContextBillingGate,
})
@@ -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';
+12
View File
@@ -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('')
})
})
+129 -1
View File
@@ -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: '(含)',
+30 -8
View File
@@ -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
+4
View File
@@ -4429,6 +4429,7 @@ import PricingEntryCard from "@/components/admin/channel/PricingEntryCard.vue";
import type { PricingFormEntry } from "@/components/admin/channel/types";
import {
apiIntervalsToForm,
createDefaultTimePricingForm,
formIntervalsToAPI,
mTokToPerToken,
perTokenToMTok,
@@ -4499,6 +4500,7 @@ const emptyGroupPricing = (): PricingFormEntry => ({
image_output_price: null,
per_request_price: null,
intervals: [],
time_pricing: createDefaultTimePricingForm(),
});
const addGroupPricing = (entries: PricingFormEntry[]) =>
@@ -4518,6 +4520,7 @@ const groupPricingFromAPI = (
image_output_price: perTokenToMTok(entry.image_output_price),
per_request_price: entry.per_request_price,
intervals: apiIntervalsToForm(entry.intervals || []),
time_pricing: createDefaultTimePricingForm(),
}));
const groupPricingToAPI = (
@@ -4541,6 +4544,7 @@ const groupPricingToAPI = (
entry.billing_mode === "token"
? []
: formIntervalsToAPI(entry.intervals || []),
time_pricing: null,
}));
const { t } = useI18n();