Skip to content

Commit 4edcfe1

Browse files
fix(usage): preserve requested model in gateway billing paths
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent) Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
1 parent 9259dcb commit 4edcfe1

4 files changed

Lines changed: 79 additions & 11 deletions

File tree

backend/internal/service/gateway_record_usage_test.go

Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -162,6 +162,32 @@ func TestGatewayServiceRecordUsage_BillingFingerprintFallsBackToContextRequestID
162162
require.Equal(t, "local:req-local-123", billingRepo.lastCmd.RequestPayloadHash)
163163
}
164164

165+
func TestGatewayServiceRecordUsage_PreservesRequestedAndUpstreamModels(t *testing.T) {
166+
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
167+
svc := newGatewayRecordUsageServiceForTest(usageRepo, &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{})
168+
mappedModel := "claude-sonnet-4-20250514"
169+
170+
err := svc.RecordUsage(context.Background(), &RecordUsageInput{
171+
Result: &ForwardResult{
172+
RequestID: "gateway_models_split",
173+
Usage: ClaudeUsage{InputTokens: 10, OutputTokens: 6},
174+
Model: "claude-sonnet-4",
175+
UpstreamModel: mappedModel,
176+
Duration: time.Second,
177+
},
178+
APIKey: &APIKey{ID: 501, Quota: 100},
179+
User: &User{ID: 601},
180+
Account: &Account{ID: 701},
181+
})
182+
183+
require.NoError(t, err)
184+
require.NotNil(t, usageRepo.lastLog)
185+
require.Equal(t, "claude-sonnet-4", usageRepo.lastLog.Model)
186+
require.Equal(t, "claude-sonnet-4", usageRepo.lastLog.RequestedModel)
187+
require.NotNil(t, usageRepo.lastLog.UpstreamModel)
188+
require.Equal(t, mappedModel, *usageRepo.lastLog.UpstreamModel)
189+
}
190+
165191
func TestGatewayServiceRecordUsage_UsageLogWriteErrorDoesNotSkipBilling(t *testing.T) {
166192
usageRepo := &openAIRecordUsageLogRepoStub{inserted: false, err: MarkUsageLogCreateNotPersisted(context.Canceled)}
167193
userRepo := &openAIRecordUsageUserRepoStub{}

backend/internal/service/gateway_service.go

Lines changed: 15 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -482,10 +482,12 @@ type ClaudeUsage struct {
482482

483483
// ForwardResult 转发结果
484484
type ForwardResult struct {
485-
RequestID string
486-
Usage ClaudeUsage
487-
Model string
488-
UpstreamModel string // Actual upstream model after mapping (empty = no mapping)
485+
RequestID string
486+
Usage ClaudeUsage
487+
Model string
488+
// UpstreamModel is the actual upstream model after mapping.
489+
// Prefer empty when it is identical to Model; persistence normalizes equal values away as no-op mappings.
490+
UpstreamModel string
489491
Stream bool
490492
Duration time.Duration
491493
FirstTokenMs *int // 首字时间(流式请求)
@@ -7516,6 +7518,7 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu
75167518
}
75177519

75187520
var cost *CostBreakdown
7521+
billingModel := forwardResultBillingModel(result.Model, result.UpstreamModel)
75197522

75207523
// 根据请求类型选择计费方式
75217524
if result.MediaType == "image" || result.MediaType == "video" {
@@ -7531,7 +7534,7 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu
75317534
if result.MediaType == "image" {
75327535
cost = s.billingService.CalculateSoraImageCost(result.ImageSize, result.ImageCount, soraConfig, multiplier)
75337536
} else {
7534-
cost = s.billingService.CalculateSoraVideoCost(result.Model, soraConfig, multiplier)
7537+
cost = s.billingService.CalculateSoraVideoCost(billingModel, soraConfig, multiplier)
75357538
}
75367539
} else if result.MediaType == "prompt" {
75377540
cost = &CostBreakdown{}
@@ -7545,7 +7548,7 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu
75457548
Price4K: apiKey.Group.ImagePrice4K,
75467549
}
75477550
}
7548-
cost = s.billingService.CalculateImageCost(result.Model, result.ImageSize, result.ImageCount, groupConfig, multiplier)
7551+
cost = s.billingService.CalculateImageCost(billingModel, result.ImageSize, result.ImageCount, groupConfig, multiplier)
75497552
} else {
75507553
// Token 计费
75517554
tokens := UsageTokens{
@@ -7557,7 +7560,7 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu
75577560
CacheCreation1hTokens: result.Usage.CacheCreation1hTokens,
75587561
}
75597562
var err error
7560-
cost, err = s.billingService.CalculateCost(result.Model, tokens, multiplier)
7563+
cost, err = s.billingService.CalculateCost(billingModel, tokens, multiplier)
75617564
if err != nil {
75627565
logger.LegacyPrintf("service.gateway", "Calculate cost failed: %v", err)
75637566
cost = &CostBreakdown{ActualCost: 0}
@@ -7589,6 +7592,7 @@ func (s *GatewayService) RecordUsage(ctx context.Context, input *RecordUsageInpu
75897592
AccountID: account.ID,
75907593
RequestID: requestID,
75917594
Model: result.Model,
7595+
RequestedModel: result.Model,
75927596
UpstreamModel: optionalNonEqualStringPtr(result.UpstreamModel, result.Model),
75937597
ReasoningEffort: result.ReasoningEffort,
75947598
InboundEndpoint: optionalTrimmedStringPtr(input.InboundEndpoint),
@@ -7719,6 +7723,7 @@ func (s *GatewayService) RecordUsageWithLongContext(ctx context.Context, input *
77197723
}
77207724

77217725
var cost *CostBreakdown
7726+
billingModel := forwardResultBillingModel(result.Model, result.UpstreamModel)
77227727

77237728
// 根据请求类型选择计费方式
77247729
if result.ImageCount > 0 {
@@ -7731,7 +7736,7 @@ func (s *GatewayService) RecordUsageWithLongContext(ctx context.Context, input *
77317736
Price4K: apiKey.Group.ImagePrice4K,
77327737
}
77337738
}
7734-
cost = s.billingService.CalculateImageCost(result.Model, result.ImageSize, result.ImageCount, groupConfig, multiplier)
7739+
cost = s.billingService.CalculateImageCost(billingModel, result.ImageSize, result.ImageCount, groupConfig, multiplier)
77357740
} else {
77367741
// Token 计费(使用长上下文计费方法)
77377742
tokens := UsageTokens{
@@ -7743,7 +7748,7 @@ func (s *GatewayService) RecordUsageWithLongContext(ctx context.Context, input *
77437748
CacheCreation1hTokens: result.Usage.CacheCreation1hTokens,
77447749
}
77457750
var err error
7746-
cost, err = s.billingService.CalculateCostWithLongContext(result.Model, tokens, multiplier, input.LongContextThreshold, input.LongContextMultiplier)
7751+
cost, err = s.billingService.CalculateCostWithLongContext(billingModel, tokens, multiplier, input.LongContextThreshold, input.LongContextMultiplier)
77477752
if err != nil {
77487753
logger.LegacyPrintf("service.gateway", "Calculate cost failed: %v", err)
77497754
cost = &CostBreakdown{ActualCost: 0}
@@ -7771,6 +7776,7 @@ func (s *GatewayService) RecordUsageWithLongContext(ctx context.Context, input *
77717776
AccountID: account.ID,
77727777
RequestID: requestID,
77737778
Model: result.Model,
7779+
RequestedModel: result.Model,
77747780
UpstreamModel: optionalNonEqualStringPtr(result.UpstreamModel, result.Model),
77757781
ReasoningEffort: result.ReasoningEffort,
77767782
InboundEndpoint: optionalTrimmedStringPtr(input.InboundEndpoint),

backend/internal/service/openai_gateway_record_usage_test.go

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -879,6 +879,7 @@ func TestOpenAIGatewayServiceRecordUsage_UsesRequestedModelAndUpstreamModelMetad
879879
require.NoError(t, err)
880880
require.NotNil(t, usageRepo.lastLog)
881881
require.Equal(t, "gpt-5.1", usageRepo.lastLog.Model)
882+
require.Equal(t, "gpt-5.1", usageRepo.lastLog.RequestedModel)
882883
require.NotNil(t, usageRepo.lastLog.UpstreamModel)
883884
require.Equal(t, "gpt-5.1-codex", *usageRepo.lastLog.UpstreamModel)
884885
require.NotNil(t, usageRepo.lastLog.ServiceTier)
@@ -894,6 +895,40 @@ func TestOpenAIGatewayServiceRecordUsage_UsesRequestedModelAndUpstreamModelMetad
894895
require.Equal(t, 1, userRepo.deductCalls)
895896
}
896897

898+
func TestOpenAIGatewayServiceRecordUsage_BillsMappedRequestsUsingUpstreamModelFallback(t *testing.T) {
899+
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
900+
userRepo := &openAIRecordUsageUserRepoStub{}
901+
subRepo := &openAIRecordUsageSubRepoStub{}
902+
svc := newOpenAIRecordUsageServiceForTest(usageRepo, userRepo, subRepo, nil)
903+
usage := OpenAIUsage{InputTokens: 20, OutputTokens: 10}
904+
905+
expectedCost, err := svc.billingService.CalculateCost("gpt-5.1-codex", UsageTokens{
906+
InputTokens: 20,
907+
OutputTokens: 10,
908+
}, 1.1)
909+
require.NoError(t, err)
910+
911+
err = svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{
912+
Result: &OpenAIForwardResult{
913+
RequestID: "resp_upstream_model_billing_fallback",
914+
Model: "gpt-5.1",
915+
UpstreamModel: "gpt-5.1-codex",
916+
Usage: usage,
917+
Duration: time.Second,
918+
},
919+
APIKey: &APIKey{ID: 10},
920+
User: &User{ID: 20},
921+
Account: &Account{ID: 30},
922+
})
923+
924+
require.NoError(t, err)
925+
require.NotNil(t, usageRepo.lastLog)
926+
require.Equal(t, "gpt-5.1", usageRepo.lastLog.Model)
927+
require.Equal(t, expectedCost.ActualCost, usageRepo.lastLog.ActualCost)
928+
require.Equal(t, expectedCost.TotalCost, usageRepo.lastLog.TotalCost)
929+
require.Equal(t, expectedCost.ActualCost, userRepo.lastAmount)
930+
}
931+
897932
func TestOpenAIGatewayServiceRecordUsage_SubscriptionBillingSetsSubscriptionFields(t *testing.T) {
898933
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
899934
userRepo := &openAIRecordUsageUserRepoStub{}

backend/internal/service/openai_gateway_service.go

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -4110,9 +4110,9 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
41104110
multiplier = resolver.Resolve(ctx, user.ID, *apiKey.GroupID, apiKey.Group.RateMultiplier)
41114111
}
41124112

4113-
billingModel := result.Model
4113+
billingModel := forwardResultBillingModel(result.Model, result.UpstreamModel)
41144114
if result.BillingModel != "" {
4115-
billingModel = result.BillingModel
4115+
billingModel = strings.TrimSpace(result.BillingModel)
41164116
}
41174117
serviceTier := ""
41184118
if result.ServiceTier != nil {
@@ -4140,6 +4140,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
41404140
AccountID: account.ID,
41414141
RequestID: requestID,
41424142
Model: result.Model,
4143+
RequestedModel: result.Model,
41434144
UpstreamModel: optionalNonEqualStringPtr(result.UpstreamModel, result.Model),
41444145
ServiceTier: result.ServiceTier,
41454146
ReasoningEffort: result.ReasoningEffort,

0 commit comments

Comments
 (0)