From 1e2193c3d27b7d770158fc7e133772b218bbe4dd Mon Sep 17 00:00:00 2001 From: gsh Date: Sun, 31 May 2026 15:09:06 +0800 Subject: [PATCH] fix: avoid websocket usage dedup conflicts --- .../openai_gateway_record_usage_test.go | 31 +++++++++++++++++++ .../service/openai_gateway_service.go | 5 +++ 2 files changed, 36 insertions(+) diff --git a/backend/internal/service/openai_gateway_record_usage_test.go b/backend/internal/service/openai_gateway_record_usage_test.go index 9769a82e..318c0861 100644 --- a/backend/internal/service/openai_gateway_record_usage_test.go +++ b/backend/internal/service/openai_gateway_record_usage_test.go @@ -721,6 +721,37 @@ func TestOpenAIGatewayServiceRecordUsage_PrefersClientRequestIDOverUpstreamReque require.Equal(t, "client:openai-client-stable-123", usageRepo.lastLog.RequestID) } +func TestOpenAIGatewayServiceRecordUsage_WSModePrefersUpstreamRequestIDOverClientRequestID(t *testing.T) { + usageRepo := &openAIRecordUsageLogRepoStub{} + billingRepo := &openAIRecordUsageBillingRepoStub{result: &UsageBillingApplyResult{Applied: true}} + userRepo := &openAIRecordUsageUserRepoStub{} + subRepo := &openAIRecordUsageSubRepoStub{} + svc := newOpenAIRecordUsageServiceWithBillingRepoForTest(usageRepo, billingRepo, userRepo, subRepo, nil) + + ctx := context.WithValue(context.Background(), ctxkey.ClientRequestID, "openai-ws-connection-123") + err := svc.RecordUsage(ctx, &OpenAIRecordUsageInput{ + Result: &OpenAIForwardResult{ + RequestID: "resp_openai_ws_turn_456", + OpenAIWSMode: true, + Usage: OpenAIUsage{ + InputTokens: 8, + OutputTokens: 4, + }, + Model: "gpt-5.1", + Duration: time.Second, + }, + APIKey: &APIKey{ID: 10050}, + User: &User{ID: 20050}, + Account: &Account{ID: 30050}, + }) + + require.NoError(t, err) + require.NotNil(t, billingRepo.lastCmd) + require.Equal(t, "resp_openai_ws_turn_456", billingRepo.lastCmd.RequestID) + require.NotNil(t, usageRepo.lastLog) + require.Equal(t, "resp_openai_ws_turn_456", usageRepo.lastLog.RequestID) +} + func TestOpenAIGatewayServiceRecordUsage_GeneratesRequestIDWhenAllSourcesMissing(t *testing.T) { usageRepo := &openAIRecordUsageLogRepoStub{} billingRepo := &openAIRecordUsageBillingRepoStub{result: &UsageBillingApplyResult{Applied: true}} diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index cd5a4015..10080f31 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -5761,6 +5761,11 @@ func (s *OpenAIGatewayService) RecordUsage(ctx context.Context, input *OpenAIRec durationMs := int(result.Duration.Milliseconds()) accountRateMultiplier := account.BillingRateMultiplier() requestID := resolveUsageBillingRequestID(ctx, result.RequestID) + if result.OpenAIWSMode { + if upstreamRequestID := strings.TrimSpace(result.RequestID); upstreamRequestID != "" { + requestID = upstreamRequestID + } + } // 确定 RequestedModel(渠道映射前的原始模型) requestedModel := result.Model