diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index d381d7a3bf..062a54bf40 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -3066,6 +3066,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco // Handle normal response var usage *OpenAIUsage var firstTokenMs *int + responseID := "" imageCount := 0 var imageOutputSizes []string if reqStream { @@ -3075,6 +3076,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco } usage = streamResult.usage firstTokenMs = streamResult.firstTokenMs + responseID = strings.TrimSpace(streamResult.responseID) imageCount = streamResult.imageCount imageOutputSizes = streamResult.imageOutputSizes } else { @@ -3083,9 +3085,11 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco return nil, err } usage = nonStreamResult.usage + responseID = strings.TrimSpace(nonStreamResult.responseID) imageCount = nonStreamResult.imageCount imageOutputSizes = nonStreamResult.imageOutputSizes } + s.bindHTTPResponseAccount(ctx, c, account, responseID) // Extract and save Codex usage snapshot from response headers (for OAuth accounts) if account.Type == AccountTypeOAuth { @@ -3100,6 +3104,7 @@ func (s *OpenAIGatewayService) Forward(ctx context.Context, c *gin.Context, acco forwardResult := &OpenAIForwardResult{ RequestID: resp.Header.Get("x-request-id"), + ResponseID: responseID, Usage: *usage, Model: originalModel, UpstreamModel: upstreamModel, @@ -3312,6 +3317,7 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough( var usage *OpenAIUsage var firstTokenMs *int + responseID := "" imageCount := 0 var imageOutputSizes []string if reqStream { @@ -3321,6 +3327,7 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough( } usage = result.usage firstTokenMs = result.firstTokenMs + responseID = strings.TrimSpace(result.responseID) imageCount = result.imageCount imageOutputSizes = result.imageOutputSizes } else { @@ -3329,9 +3336,11 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough( return nil, err } usage = result.usage + responseID = strings.TrimSpace(result.responseID) imageCount = result.imageCount imageOutputSizes = result.imageOutputSizes } + s.bindHTTPResponseAccount(ctx, c, account, responseID) if snapshot := ParseCodexRateLimitHeaders(resp.Header); snapshot != nil { s.updateCodexUsageSnapshot(ctx, account.ID, snapshot) @@ -3343,6 +3352,7 @@ func (s *OpenAIGatewayService) forwardOpenAIPassthrough( forwardResult := &OpenAIForwardResult{ RequestID: resp.Header.Get("x-request-id"), + ResponseID: responseID, Usage: *usage, Model: reqModel, UpstreamModel: upstreamPassthroughModel, @@ -3657,6 +3667,7 @@ func collectOpenAIPassthroughTimeoutHeaders(h http.Header) []string { type openaiStreamingResultPassthrough struct { usage *OpenAIUsage firstTokenMs *int + responseID string imageCount int imageOutputSizes []string } @@ -3664,6 +3675,7 @@ type openaiStreamingResultPassthrough struct { type openaiNonStreamingResultPassthrough struct { *OpenAIUsage usage *OpenAIUsage + responseID string imageCount int imageOutputSizes []string } @@ -3807,6 +3819,7 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough( usage := &OpenAIUsage{} imageCounter := newOpenAIImageOutputCounter() var firstTokenMs *int + responseID := "" clientDisconnected := false sawDone := false sawTerminalEvent := false @@ -3841,6 +3854,7 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough( return &openaiStreamingResultPassthrough{ usage: usage, firstTokenMs: firstTokenMs, + responseID: responseID, imageCount: imageCounter.Count(), imageOutputSizes: imageCounter.Sizes(), } @@ -3876,6 +3890,9 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough( if openAIStreamEventIsTerminal(trimmedData) { sawTerminalEvent = true } + if responseID == "" { + responseID = extractOpenAIResponseIDFromJSONBytes(dataBytes) + } imageCounter.AddSSEData(dataBytes) lineStartsClientOutput = forceFlushFailedEvent || openAIStreamDataStartsClientOutput(trimmedData, eventType) if firstTokenMs == nil && lineStartsClientOutput && trimmedData != "[DONE]" { @@ -4002,6 +4019,7 @@ func (s *OpenAIGatewayService) handleNonStreamingResponsePassthrough( return &openaiNonStreamingResultPassthrough{ OpenAIUsage: usage, usage: usage, + responseID: extractOpenAIResponseIDFromJSONBytes(body), imageCount: countOpenAIResponseImageOutputsFromJSONBytes(body), imageOutputSizes: collectOpenAIResponseImageOutputSizesFromJSONBytes(body), }, nil @@ -4065,6 +4083,7 @@ func (s *OpenAIGatewayService) handlePassthroughSSEToJSON(resp *http.Response, c return &openaiNonStreamingResultPassthrough{ OpenAIUsage: usage, usage: usage, + responseID: extractOpenAIResponseIDFromJSONBytes(body), imageCount: countOpenAIImageOutputsFromSSEBody(bodyText), imageOutputSizes: collectOpenAIImageOutputSizesFromSSEBody(bodyText), }, nil @@ -4526,6 +4545,7 @@ func (s *OpenAIGatewayService) handleCompatErrorResponse( type openaiStreamingResult struct { usage *OpenAIUsage firstTokenMs *int + responseID string imageCount int imageOutputSizes []string } @@ -4533,6 +4553,7 @@ type openaiStreamingResult struct { type openaiNonStreamingResult struct { *OpenAIUsage usage *OpenAIUsage + responseID string imageCount int imageOutputSizes []string } @@ -4570,6 +4591,7 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp usage := &OpenAIUsage{} imageCounter := newOpenAIImageOutputCounter() var firstTokenMs *int + responseID := "" scanner := bufio.NewScanner(resp.Body) maxLineSize := defaultMaxLineSize if s.cfg != nil && s.cfg.Gateway.MaxLineSize > 0 { @@ -4653,6 +4675,7 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp return &openaiStreamingResult{ usage: usage, firstTokenMs: firstTokenMs, + responseID: responseID, imageCount: imageCounter.Count(), imageOutputSizes: imageCounter.Sizes(), } @@ -4732,6 +4755,9 @@ func (s *OpenAIGatewayService) handleStreamingResponse(ctx context.Context, resp sawTerminalEvent = true } eventType := strings.TrimSpace(gjson.GetBytes(dataBytes, "type").String()) + if responseID == "" { + responseID = extractOpenAIResponseIDFromJSONBytes(dataBytes) + } forceFlushFailedEvent := false if eventType == "response.failed" { failedMessage = extractOpenAISSEErrorMessage(dataBytes) @@ -5089,6 +5115,33 @@ func extractOpenAIUsageFromJSONBytes(body []byte) (OpenAIUsage, bool) { return openAIUsageFromGJSON(gjson.GetBytes(body, "response.usage")) } +func extractOpenAIResponseIDFromJSONBytes(body []byte) string { + if len(body) == 0 || !gjson.ValidBytes(body) { + return "" + } + if id := strings.TrimSpace(gjson.GetBytes(body, "id").String()); id != "" { + return id + } + return strings.TrimSpace(gjson.GetBytes(body, "response.id").String()) +} + +func (s *OpenAIGatewayService) bindHTTPResponseAccount(ctx context.Context, c *gin.Context, account *Account, responseID string) { + if s == nil || account == nil || account.ID <= 0 { + return + } + responseID = strings.TrimSpace(responseID) + if responseID == "" { + return + } + store := s.getOpenAIWSStateStore() + if store == nil { + return + } + groupID := getOpenAIGroupIDFromContext(c) + ttl := s.openAIWSResponseStickyTTL() + logOpenAIWSBindResponseAccountWarn(groupID, account.ID, responseID, store.BindResponseAccount(ctx, groupID, responseID, account.ID, ttl)) +} + func openAIUsageFromGJSON(value gjson.Result) (OpenAIUsage, bool) { if !value.Exists() || !value.IsObject() { return OpenAIUsage{}, false @@ -5169,6 +5222,7 @@ func (s *OpenAIGatewayService) handleNonStreamingResponse(ctx context.Context, r return &openaiNonStreamingResult{ OpenAIUsage: usage, usage: usage, + responseID: extractOpenAIResponseIDFromJSONBytes(body), imageCount: countOpenAIResponseImageOutputsFromJSONBytes(body), imageOutputSizes: collectOpenAIResponseImageOutputSizesFromJSONBytes(body), }, nil @@ -5234,6 +5288,7 @@ func (s *OpenAIGatewayService) handleSSEToJSON(resp *http.Response, c *gin.Conte return &openaiNonStreamingResult{ OpenAIUsage: usage, usage: usage, + responseID: extractOpenAIResponseIDFromJSONBytes(body), imageCount: countOpenAIImageOutputsFromSSEBody(bodyText), imageOutputSizes: collectOpenAIImageOutputSizesFromSSEBody(bodyText), }, nil diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 440f378055..e00a3dc2e2 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -227,6 +227,30 @@ func TestOpenAIGatewayService_GenerateSessionHash_AttachesLegacyHashToContext(t require.NotEmpty(t, openAILegacySessionHashFromContext(c.Request.Context())) } +func TestExtractOpenAIResponseIDFromJSONBytes(t *testing.T) { + require.Equal(t, "resp_json", extractOpenAIResponseIDFromJSONBytes([]byte(`{"id":"resp_json"}`))) + require.Equal(t, "resp_sse", extractOpenAIResponseIDFromJSONBytes([]byte(`{"type":"response.completed","response":{"id":"resp_sse"}}`))) + require.Empty(t, extractOpenAIResponseIDFromJSONBytes([]byte(`{"response":{}}`))) + require.Empty(t, extractOpenAIResponseIDFromJSONBytes([]byte(`not-json`))) +} + +func TestOpenAIGatewayService_BindHTTPResponseAccount(t *testing.T) { + gin.SetMode(gin.TestMode) + rec := httptest.NewRecorder() + c, _ := gin.CreateTestContext(rec) + c.Request = httptest.NewRequest(http.MethodPost, "/openai/v1/responses", nil) + groupID := int64(4201) + c.Set("api_key", &APIKey{ID: 501, GroupID: &groupID}) + + svc := &OpenAIGatewayService{} + account := &Account{ID: 37001, Platform: PlatformOpenAI, Type: AccountTypeAPIKey} + svc.bindHTTPResponseAccount(context.Background(), c, account, "resp_http_001") + + got, err := svc.getOpenAIWSStateStore().GetResponseAccount(context.Background(), groupID, "resp_http_001") + require.NoError(t, err) + require.Equal(t, account.ID, got) +} + func TestOpenAIGatewayService_GenerateExplicitSessionHash_SkipsContentFallback(t *testing.T) { gin.SetMode(gin.TestMode) svc := &OpenAIGatewayService{}