fix(usage): preserve final upstream model

This commit is contained in:
visa2
2026-07-27 00:43:53 +08:00
parent 1f45c99de7
commit be65c713ff
5 changed files with 62 additions and 12 deletions
@@ -223,6 +223,35 @@ func TestGatewayServiceRecordUsage_PreservesChannelMappedUpstreamModel(t *testin
require.Equal(t, "gpt-5.6-terra", *usageRepo.lastLog.UpstreamModel)
}
func TestGatewayServiceRecordUsage_PreservesLoopedChannelAndAccountUpstreamModel(t *testing.T) {
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
svc := newGatewayRecordUsageServiceForTest(usageRepo, &openAIRecordUsageUserRepoStub{}, &openAIRecordUsageSubRepoStub{})
err := svc.RecordUsage(context.Background(), &RecordUsageInput{
Result: &ForwardResult{
RequestID: "gateway_looped_mapping_models",
Usage: ClaudeUsage{InputTokens: 10, OutputTokens: 6},
Model: "gpt-5.6-terra",
UpstreamModel: "gpt-5.6-sol",
Duration: time.Second,
},
APIKey: &APIKey{ID: 501, Quota: 100},
User: &User{ID: 601},
Account: &Account{ID: 701},
ChannelUsageFields: ChannelUsageFields{
OriginalModel: "gpt-5.6-sol",
ChannelMappedModel: "gpt-5.6-terra",
},
})
require.NoError(t, err)
require.NotNil(t, usageRepo.lastLog)
require.Equal(t, "gpt-5.6-sol", usageRepo.lastLog.RequestedModel)
require.Equal(t, "gpt-5.6-terra", usageRepo.lastLog.Model)
require.NotNil(t, usageRepo.lastLog.UpstreamModel)
require.Equal(t, "gpt-5.6-sol", *usageRepo.lastLog.UpstreamModel)
}
func TestGatewayServiceRecordUsage_EmptyImageSizeDefaultsBeforeBillingAndPersistence(t *testing.T) {
imagePrice2K := 0.19
groupID := int64(901)
@@ -986,7 +986,7 @@ func (s *GatewayService) buildRecordUsageLog(
RequestID: requestID,
Model: result.Model,
RequestedModel: requestedModel,
UpstreamModel: optionalNonEqualStringPtr(result.UpstreamModel, requestedModel),
UpstreamModel: optionalTrimmedStringPtr(result.UpstreamModel),
ReasoningEffort: result.ReasoningEffort,
InboundEndpoint: optionalTrimmedStringPtr(input.InboundEndpoint),
UpstreamEndpoint: optionalTrimmedStringPtr(input.UpstreamEndpoint),
@@ -1369,6 +1369,37 @@ func TestOpenAIGatewayServiceRecordUsage_PreservesChannelMappedUpstreamModel(t *
require.Equal(t, "gpt-5.6-terra", *usageRepo.lastLog.UpstreamModel)
}
func TestOpenAIGatewayServiceRecordUsage_PreservesLoopedChannelAndAccountUpstreamModel(t *testing.T) {
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
userRepo := &openAIRecordUsageUserRepoStub{}
subRepo := &openAIRecordUsageSubRepoStub{}
svc := newOpenAIRecordUsageServiceForTest(usageRepo, userRepo, subRepo, nil)
err := svc.RecordUsage(context.Background(), &OpenAIRecordUsageInput{
Result: &OpenAIForwardResult{
RequestID: "openai_looped_mapping_models",
Model: "gpt-5.6-terra",
UpstreamModel: "gpt-5.6-sol",
Usage: OpenAIUsage{InputTokens: 20, OutputTokens: 10},
Duration: time.Second,
},
APIKey: &APIKey{ID: 10},
User: &User{ID: 20},
Account: &Account{ID: 30},
ChannelUsageFields: ChannelUsageFields{
OriginalModel: "gpt-5.6-sol",
ChannelMappedModel: "gpt-5.6-terra",
},
})
require.NoError(t, err)
require.NotNil(t, usageRepo.lastLog)
require.Equal(t, "gpt-5.6-sol", usageRepo.lastLog.RequestedModel)
require.Equal(t, "gpt-5.6-terra", usageRepo.lastLog.Model)
require.NotNil(t, usageRepo.lastLog.UpstreamModel)
require.Equal(t, "gpt-5.6-sol", *usageRepo.lastLog.UpstreamModel)
}
func TestOpenAIGatewayServiceRecordUsage_BillsMappedRequestsUsingRequestedModel(t *testing.T) {
usageRepo := &openAIRecordUsageLogRepoStub{inserted: true}
userRepo := &openAIRecordUsageUserRepoStub{}
@@ -256,7 +256,7 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec
RequestID: requestID,
Model: result.Model,
RequestedModel: requestedModel,
UpstreamModel: optionalNonEqualStringPtr(result.UpstreamModel, requestedModel),
UpstreamModel: optionalTrimmedStringPtr(result.UpstreamModel),
ServiceTier: result.ServiceTier,
ReasoningEffort: result.ReasoningEffort,
InboundEndpoint: optionalTrimmedStringPtr(input.InboundEndpoint),
@@ -10,16 +10,6 @@ func optionalTrimmedStringPtr(raw string) *string {
return &trimmed
}
// optionalNonEqualStringPtr returns a pointer to value if it is non-empty and
// differs from compare; otherwise nil. Usage logging passes the requested
// model as compare so a channel mapping still records its effective upstream.
func optionalNonEqualStringPtr(value, compare string) *string {
if value == "" || value == compare {
return nil
}
return &value
}
func forwardResultBillingModel(requestedModel, upstreamModel string) string {
if trimmed := strings.TrimSpace(requestedModel); trimmed != "" {
return trimmed