From 7e96c8ef0babc9d885860b7c049a375debb9fadb Mon Sep 17 00:00:00 2001 From: chenlingzhi Date: Thu, 30 Apr 2026 14:58:25 +0800 Subject: [PATCH] fix: aggregate SSE upstream into chat.completion JSON for non-stream clients (Co-Authored-By: Claude Opus 4.7 ) --- relay/channel/openai/chat_via_responses.go | 231 +++++++++++++++++++++ relay/chat_completions_via_responses.go | 20 +- 2 files changed, 249 insertions(+), 2 deletions(-) diff --git a/relay/channel/openai/chat_via_responses.go b/relay/channel/openai/chat_via_responses.go index 2c0752275daa..01937c315f29 100644 --- a/relay/channel/openai/chat_via_responses.go +++ b/relay/channel/openai/chat_via_responses.go @@ -548,3 +548,234 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo } return usage, nil } + +// OaiResponsesSSEToChatJSON handles the case where the client requested a +// non-streaming chat completion but the upstream /v1/responses endpoint +// returned an SSE stream (common for reasoning models). It parses the SSE +// events, accumulates output text / tool calls / usage, builds a single +// dto.OpenAITextResponse, and writes it to the client as one JSON body. +// +// This avoids the bug where, in the original code path, when upstream returns +// SSE for a non-stream client request, raw SSE chunks (with empty choices) get +// forwarded directly to the client. +func OaiResponsesSSEToChatJSON(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { + if resp == nil || resp.Body == nil { + return nil, types.NewOpenAIError(fmt.Errorf("invalid response"), types.ErrorCodeBadResponse, http.StatusInternalServerError) + } + defer service.CloseResponseBodyGracefully(resp) + + responseId := helper.GetResponseID(c) + createAt := time.Now().Unix() + model := info.UpstreamModelName + + var ( + usage = &dto.Usage{} + outputText strings.Builder + usageText strings.Builder + streamErr *types.NewAPIError + ) + + toolCallIndexByID := make(map[string]int) + toolCallNameByID := make(map[string]string) + toolCallArgsByID := make(map[string]string) + toolCallOrder := make([]string, 0) + toolCallCanonicalIDByItemID := make(map[string]string) + + registerToolCall := func(callID string) { + if _, ok := toolCallIndexByID[callID]; ok { + return + } + toolCallIndexByID[callID] = len(toolCallOrder) + toolCallOrder = append(toolCallOrder, callID) + } + + helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) { + if streamErr != nil { + sr.Stop(streamErr) + return + } + + var streamResp dto.ResponsesStreamResponse + if err := common.UnmarshalJsonStr(data, &streamResp); err != nil { + logger.LogError(c, "failed to unmarshal responses stream event: "+err.Error()) + return + } + + switch streamResp.Type { + case "response.created": + if streamResp.Response != nil { + if streamResp.Response.Model != "" { + model = streamResp.Response.Model + } + if streamResp.Response.CreatedAt != 0 { + createAt = int64(streamResp.Response.CreatedAt) + } + } + + case "response.output_text.delta": + if streamResp.Delta != "" { + outputText.WriteString(streamResp.Delta) + usageText.WriteString(streamResp.Delta) + } + + case "response.output_item.added", "response.output_item.done": + if streamResp.Item == nil || streamResp.Item.Type != "function_call" { + break + } + itemID := strings.TrimSpace(streamResp.Item.ID) + callID := strings.TrimSpace(streamResp.Item.CallId) + if callID == "" { + callID = itemID + } + if itemID != "" && callID != "" { + toolCallCanonicalIDByItemID[itemID] = callID + } + if callID == "" { + break + } + registerToolCall(callID) + if name := strings.TrimSpace(streamResp.Item.Name); name != "" { + toolCallNameByID[callID] = name + usageText.WriteString(name) + } + if args := streamResp.Item.ArgumentsString(); args != "" { + toolCallArgsByID[callID] = args + } + + case "response.function_call_arguments.delta": + itemID := strings.TrimSpace(streamResp.ItemID) + callID := toolCallCanonicalIDByItemID[itemID] + if callID == "" { + callID = itemID + } + if callID == "" { + break + } + registerToolCall(callID) + toolCallArgsByID[callID] += streamResp.Delta + usageText.WriteString(streamResp.Delta) + + case "response.completed": + if streamResp.Response != nil { + if streamResp.Response.Model != "" { + model = streamResp.Response.Model + } + if streamResp.Response.CreatedAt != 0 { + createAt = int64(streamResp.Response.CreatedAt) + } + if streamResp.Response.Usage != nil { + if streamResp.Response.Usage.InputTokens != 0 { + usage.PromptTokens = streamResp.Response.Usage.InputTokens + usage.InputTokens = streamResp.Response.Usage.InputTokens + } + if streamResp.Response.Usage.OutputTokens != 0 { + usage.CompletionTokens = streamResp.Response.Usage.OutputTokens + usage.OutputTokens = streamResp.Response.Usage.OutputTokens + } + if streamResp.Response.Usage.TotalTokens != 0 { + usage.TotalTokens = streamResp.Response.Usage.TotalTokens + } else { + usage.TotalTokens = usage.PromptTokens + usage.CompletionTokens + } + if streamResp.Response.Usage.InputTokensDetails != nil { + usage.PromptTokensDetails.CachedTokens = streamResp.Response.Usage.InputTokensDetails.CachedTokens + usage.PromptTokensDetails.ImageTokens = streamResp.Response.Usage.InputTokensDetails.ImageTokens + usage.PromptTokensDetails.AudioTokens = streamResp.Response.Usage.InputTokensDetails.AudioTokens + } + if streamResp.Response.Usage.CompletionTokenDetails.ReasoningTokens != 0 { + usage.CompletionTokenDetails.ReasoningTokens = streamResp.Response.Usage.CompletionTokenDetails.ReasoningTokens + } + } + } + + case "response.error", "response.failed": + if streamResp.Response != nil { + if oaiErr := streamResp.Response.GetOpenAIError(); oaiErr != nil && oaiErr.Type != "" { + streamErr = types.WithOpenAIError(*oaiErr, http.StatusInternalServerError) + sr.Stop(streamErr) + return + } + } + streamErr = types.NewOpenAIError(fmt.Errorf("responses stream error: %s", streamResp.Type), types.ErrorCodeBadResponse, http.StatusInternalServerError) + sr.Stop(streamErr) + return + + default: + } + }) + + if streamErr != nil { + return nil, streamErr + } + + if usage.TotalTokens == 0 { + usage = service.ResponseText2Usage(c, usageText.String(), info.UpstreamModelName, info.GetEstimatePromptTokens()) + } + + sawToolCall := len(toolCallOrder) > 0 + finishReason := "stop" + if sawToolCall && outputText.Len() == 0 { + finishReason = "tool_calls" + } + + msg := dto.Message{ + Role: "assistant", + Content: outputText.String(), + } + + if sawToolCall { + toolCalls := make([]dto.ToolCallResponse, 0, len(toolCallOrder)) + for _, callID := range toolCallOrder { + tc := dto.ToolCallResponse{ + ID: callID, + Type: "function", + Function: dto.FunctionResponse{ + Name: toolCallNameByID[callID], + Arguments: toolCallArgsByID[callID], + }, + } + toolCalls = append(toolCalls, tc) + } + toolCallsBytes, err := common.Marshal(toolCalls) + if err != nil { + return nil, types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError) + } + msg.ToolCalls = toolCallsBytes + } + + chatResp := &dto.OpenAITextResponse{ + Id: responseId, + Model: model, + Object: "chat.completion", + Created: createAt, + Choices: []dto.OpenAITextResponseChoice{ + { + Index: 0, + Message: msg, + FinishReason: finishReason, + }, + }, + Usage: *usage, + } + + var ( + responseBody []byte + err error + ) + switch info.RelayFormat { + case types.RelayFormatClaude: + claudeResp := service.ResponseOpenAI2Claude(chatResp, info) + responseBody, err = common.Marshal(claudeResp) + case types.RelayFormatGemini: + geminiResp := service.ResponseOpenAI2Gemini(chatResp, info) + responseBody, err = common.Marshal(geminiResp) + default: + responseBody, err = common.Marshal(chatResp) + } + if err != nil { + return nil, types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError) + } + + service.IOCopyBytesGracefully(c, resp, responseBody) + return usage, nil +} diff --git a/relay/chat_completions_via_responses.go b/relay/chat_completions_via_responses.go index 7a2eb9aa6dac..d325d265ace4 100644 --- a/relay/chat_completions_via_responses.go +++ b/relay/chat_completions_via_responses.go @@ -139,14 +139,17 @@ func chatCompletionsViaResponses(c *gin.Context, info *relaycommon.RelayInfo, ad statusCodeMappingStr := c.GetString("status_code_mapping") httpResp = resp.(*http.Response) - info.IsStream = info.IsStream || strings.HasPrefix(httpResp.Header.Get("Content-Type"), "text/event-stream") + clientWantsStream := info.IsStream + upstreamIsSSE := strings.HasPrefix(httpResp.Header.Get("Content-Type"), "text/event-stream") + info.IsStream = clientWantsStream || upstreamIsSSE if httpResp.StatusCode != http.StatusOK { newApiErr := service.RelayErrorHandler(c.Request.Context(), httpResp, false) service.ResetStatusCode(newApiErr, statusCodeMappingStr) return nil, newApiErr } - if info.IsStream { + if clientWantsStream { + // Client requested stream — forward as SSE regardless of upstream format. usage, newApiErr := openaichannel.OaiResponsesToChatStreamHandler(c, info, httpResp) if newApiErr != nil { service.ResetStatusCode(newApiErr, statusCodeMappingStr) @@ -155,6 +158,19 @@ func chatCompletionsViaResponses(c *gin.Context, info *relaycommon.RelayInfo, ad return usage, nil } + if upstreamIsSSE { + // Client wants non-stream JSON, but upstream returned SSE (common for + // reasoning models). Buffer and aggregate the SSE stream into a full + // chat.completion JSON before writing to the client. + info.IsStream = false + usage, newApiErr := openaichannel.OaiResponsesSSEToChatJSON(c, info, httpResp) + if newApiErr != nil { + service.ResetStatusCode(newApiErr, statusCodeMappingStr) + return nil, newApiErr + } + return usage, nil + } + usage, newApiErr := openaichannel.OaiResponsesToChatHandler(c, info, httpResp) if newApiErr != nil { service.ResetStatusCode(newApiErr, statusCodeMappingStr)