diff --git a/controller/relay.go b/controller/relay.go index df03619e3bad..d2a8c3e62d44 100644 --- a/controller/relay.go +++ b/controller/relay.go @@ -119,8 +119,16 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) { logger.LogError(c, fmt.Sprintf("relay error: %s", common.LocalLogPreview(newAPIError.Error()))) newAPIError.SetMessage(common.MessageWithRequestId(newAPIError.Error(), requestId)) if relayFormat != types.RelayFormatOpenAIRealtime && c.Writer.Written() { + if relayFormat == types.RelayFormatOpenAIResponses && types.IsServerOverloadedError(newAPIError) { + if err := helper.ResponsesStreamError(c, newAPIError.ToOpenAIError()); err != nil { + logger.LogError(c, fmt.Sprintf("write responses stream error: %s", err.Error())) + } + } return } + if relayFormat == types.RelayFormatOpenAIResponses && types.IsServerOverloadedError(newAPIError) { + c.Writer.Header().Set("Content-Type", "application/json; charset=utf-8") + } switch relayFormat { case types.RelayFormatOpenAIRealtime: helper.WssError(c, ws, newAPIError.ToOpenAIError()) diff --git a/dto/openai_response.go b/dto/openai_response.go index 2de6014f4d05..8dbfae3d6213 100644 --- a/dto/openai_response.go +++ b/dto/openai_response.go @@ -415,6 +415,7 @@ const ( type ResponsesStreamResponse struct { Type string `json:"type"` Response *OpenAIResponsesResponse `json:"response,omitempty"` + Error any `json:"error,omitempty"` Delta string `json:"delta,omitempty"` Item *ResponsesOutput `json:"item,omitempty"` // - response.function_call_arguments.delta diff --git a/relay/channel/openai/relay_responses.go b/relay/channel/openai/relay_responses.go index 9293183168ee..046692c08745 100644 --- a/relay/channel/openai/relay_responses.go +++ b/relay/channel/openai/relay_responses.go @@ -17,6 +17,41 @@ import ( "github.com/gin-gonic/gin" ) +type responsesStreamChunk struct { + response dto.ResponsesStreamResponse + data string +} + +func isResponsesStreamPreludeEvent(eventType string) bool { + switch eventType { + case "response.created", "response.in_progress", "response.queued": + return true + default: + return false + } +} + +func responsesStreamServerOverloadedError(streamResponse dto.ResponsesStreamResponse) *types.NewAPIError { + switch streamResponse.Type { + case "error", "response.error", "response.failed": + default: + return nil + } + + errorFields := []any{streamResponse.Error} + if streamResponse.Response != nil { + errorFields = append(errorFields, streamResponse.Response.Error) + } + for _, errorField := range errorFields { + openAIError := dto.GetOpenAIError(errorField) + if openAIError == nil || !types.IsServerOverloadedCode(openAIError.Code) { + continue + } + return types.WithOpenAIError(*openAIError, http.StatusServiceUnavailable) + } + return nil +} + func OaiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) { defer service.CloseResponseBodyGracefully(resp) @@ -76,19 +111,19 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp } defer service.CloseResponseBodyGracefully(resp) + info.StreamStatus = relaycommon.NewStreamStatus() + info.ReceivedResponseCount = 0 var usage = &dto.Usage{} var responseTextBuilder strings.Builder - - helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) { - - // 检查当前数据是否包含 completed 状态和 usage 信息 - var streamResponse dto.ResponsesStreamResponse - if err := common.UnmarshalJsonStr(data, &streamResponse); err != nil { - logger.LogError(c, "failed to unmarshal stream response: "+err.Error()) - sr.Error(err) - return - } + // Do not expose lifecycle-only events until the stream produces stateful output. + // This leaves the outer relay free to retry an admission/capacity failure + // without duplicating text, reasoning, or tool-call events downstream. + var pendingChunks []responsesStreamChunk + var streamErr *types.NewAPIError + streamCommitted := false + + handleChunk := func(streamResponse dto.ResponsesStreamResponse, data string) { sendResponsesStreamData(c, streamResponse, data) switch streamResponse.Type { case "response.completed": @@ -115,10 +150,8 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp } } case "response.output_text.delta": - // 处理输出文本 responseTextBuilder.WriteString(streamResponse.Delta) case dto.ResponsesOutputTypeItemDone: - // 函数调用处理 if streamResponse.Item != nil { switch streamResponse.Item.Type { case dto.BuildInCallWebSearchCall: @@ -130,8 +163,52 @@ func OaiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp } } } + } + + helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) { + var streamResponse dto.ResponsesStreamResponse + if err := common.UnmarshalJsonStr(data, &streamResponse); err != nil { + logger.LogError(c, "failed to unmarshal stream response: "+err.Error()) + sr.Error(err) + return + } + + if !streamCommitted { + if overloadedErr := responsesStreamServerOverloadedError(streamResponse); overloadedErr != nil { + streamErr = overloadedErr + sr.Stop(overloadedErr) + return + } + + pendingChunks = append(pendingChunks, responsesStreamChunk{ + response: streamResponse, + data: data, + }) + if isResponsesStreamPreludeEvent(streamResponse.Type) { + return + } + + streamCommitted = true + for _, chunk := range pendingChunks { + handleChunk(chunk.response, chunk.data) + } + pendingChunks = nil + return + } + + handleChunk(streamResponse, data) }) + if streamErr != nil { + info.ResetFirstResponseTime() + return nil, streamErr + } + if !streamCommitted { + for _, chunk := range pendingChunks { + handleChunk(chunk.response, chunk.data) + } + } + if usage.CompletionTokens == 0 { // 计算输出文本的 token 数量 tempStr := responseTextBuilder.String() diff --git a/relay/channel/openai/relay_responses_test.go b/relay/channel/openai/relay_responses_test.go new file mode 100644 index 000000000000..e298d828b341 --- /dev/null +++ b/relay/channel/openai/relay_responses_test.go @@ -0,0 +1,139 @@ +package openai + +import ( + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + + "github.com/QuantumNous/new-api/constant" + "github.com/QuantumNous/new-api/dto" + relaycommon "github.com/QuantumNous/new-api/relay/common" + relayconstant "github.com/QuantumNous/new-api/relay/constant" + "github.com/QuantumNous/new-api/types" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func newResponsesStreamHandlerTest(t *testing.T, stream string) (*gin.Context, *httptest.ResponseRecorder, *http.Response, *relaycommon.RelayInfo) { + t.Helper() + previousStreamingTimeout := constant.StreamingTimeout + constant.StreamingTimeout = 30 + t.Cleanup(func() { + constant.StreamingTimeout = previousStreamingTimeout + }) + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + + resp := &http.Response{ + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + }, + Body: io.NopCloser(strings.NewReader(stream)), + } + info := &relaycommon.RelayInfo{ + RelayMode: relayconstant.RelayModeResponses, + RelayFormat: types.RelayFormatOpenAIResponses, + IsStream: true, + DisablePing: true, + ChannelMeta: &relaycommon.ChannelMeta{ + UpstreamModelName: "gpt-5.6-sol", + }, + } + return c, recorder, resp, info +} + +func TestOaiResponsesStreamHandler_RetriesServerOverloadBeforeVisibleOutput(t *testing.T) { + stream := strings.Join([]string{ + `data: {"type":"response.created","response":{"id":"resp_overload","status":"in_progress"}}`, + "", + `data: {"type":"response.in_progress","response":{"id":"resp_overload","status":"in_progress"}}`, + "", + `data: {"type":"error","error":{"type":"service_unavailable_error","code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."}}`, + "", + `data: {"type":"response.failed","response":{"id":"resp_overload","status":"failed","error":{"code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."}}}`, + "", + }, "\n") + c, recorder, resp, info := newResponsesStreamHandlerTest(t, stream) + + usage, apiErr := OaiResponsesStreamHandler(c, info, resp) + + require.Nil(t, usage) + require.NotNil(t, apiErr) + require.Equal(t, http.StatusServiceUnavailable, apiErr.StatusCode) + require.Equal(t, "server_is_overloaded", apiErr.ToOpenAIError().Code) + require.Empty(t, recorder.Body.String(), "retryable prelude events must not be committed downstream") +} + +func TestOaiResponsesStreamHandler_DoesNotRetryServerOverloadAfterVisibleOutput(t *testing.T) { + stream := strings.Join([]string{ + `data: {"type":"response.created","response":{"id":"resp_partial","status":"in_progress"}}`, + "", + `data: {"type":"response.output_item.added","item":{"id":"msg_partial","type":"message","status":"in_progress","role":"assistant","content":[]}}`, + "", + `data: {"type":"response.failed","response":{"id":"resp_partial","status":"failed","error":{"type":"server_error","code":"server_is_overloaded","message":"server is overloaded"}}}`, + "", + }, "\n") + c, recorder, resp, info := newResponsesStreamHandlerTest(t, stream) + + usage, apiErr := OaiResponsesStreamHandler(c, info, resp) + + require.Nil(t, apiErr) + require.NotNil(t, usage) + require.Contains(t, recorder.Body.String(), `"type":"response.created"`) + require.Contains(t, recorder.Body.String(), `"type":"response.output_item.added"`) + require.Contains(t, recorder.Body.String(), `"type":"response.failed"`) +} + +func TestOaiResponsesStreamHandler_FlushesBufferedPreludeOnSuccess(t *testing.T) { + stream := strings.Join([]string{ + `data: {"type":"response.created","response":{"id":"resp_ok","status":"in_progress"}}`, + "", + `data: {"type":"response.completed","response":{"id":"resp_ok","status":"completed","usage":{"input_tokens":10,"output_tokens":2,"total_tokens":12}}}`, + "", + }, "\n") + c, recorder, resp, info := newResponsesStreamHandlerTest(t, stream) + + usage, apiErr := OaiResponsesStreamHandler(c, info, resp) + + require.Nil(t, apiErr) + require.Equal(t, &dto.Usage{PromptTokens: 10, CompletionTokens: 2, TotalTokens: 12}, usage) + require.Contains(t, recorder.Body.String(), `"type":"response.created"`) + require.Contains(t, recorder.Body.String(), `"type":"response.completed"`) +} + +func TestResponsesStreamServerOverloadedError_RecognizesCodexCodes(t *testing.T) { + for _, testCase := range []struct { + code string + want bool + }{ + {code: "server_is_overloaded", want: true}, + {code: "slow_down", want: true}, + {code: "rate_limit_exceeded", want: false}, + } { + t.Run(testCase.code, func(t *testing.T) { + streamResponse := dto.ResponsesStreamResponse{ + Type: "response.failed", + Response: &dto.OpenAIResponsesResponse{ + Error: types.OpenAIError{ + Message: "upstream failure", + Code: testCase.code, + }, + }, + } + + apiErr := responsesStreamServerOverloadedError(streamResponse) + if testCase.want { + require.NotNil(t, apiErr) + require.Equal(t, http.StatusServiceUnavailable, apiErr.StatusCode) + } else { + require.Nil(t, apiErr) + } + }) + } +} diff --git a/relay/common/relay_info.go b/relay/common/relay_info.go index 2858fa1ef773..d2933b8e37ff 100644 --- a/relay/common/relay_info.go +++ b/relay/common/relay_info.go @@ -682,6 +682,14 @@ func (info *RelayInfo) SetTraceContext(ctx context.Context) { info.traceContext = ctx } +func (info *RelayInfo) ResetFirstResponseTime() { + if info == nil { + return + } + info.FirstResponseTime = info.StartTime.Add(-time.Second) + info.isFirstResponse = true +} + func (info *RelayInfo) HasSendResponse() bool { return info.FirstResponseTime.After(info.StartTime) } diff --git a/relay/common/relay_info_test.go b/relay/common/relay_info_test.go index bfae5d5d2a0c..3b4e452ed9c0 100644 --- a/relay/common/relay_info_test.go +++ b/relay/common/relay_info_test.go @@ -67,3 +67,23 @@ func TestRelayInfoFirstResponseRecordsOneTraceEvent(t *testing.T) { require.Len(t, spans[0].Events(), 1) require.Equal(t, "llm.response.first_chunk", spans[0].Events()[0].Name) } + +func TestRelayInfoResetFirstResponseTimeAllowsRetryToRecordNewAttempt(t *testing.T) { + startTime := time.Now().Add(-time.Second) + info := &RelayInfo{ + StartTime: startTime, + FirstResponseTime: startTime.Add(-time.Second), + isFirstResponse: true, + } + + info.SetFirstResponseTime() + require.True(t, info.HasSendResponse()) + + info.ResetFirstResponseTime() + require.Equal(t, startTime.Add(-time.Second), info.FirstResponseTime) + require.False(t, info.HasSendResponse()) + require.True(t, info.isFirstResponse) + + info.SetFirstResponseTime() + require.True(t, info.HasSendResponse()) +} diff --git a/relay/helper/common.go b/relay/helper/common.go index 5b118aef8118..e44bc9474472 100644 --- a/relay/helper/common.go +++ b/relay/helper/common.go @@ -94,6 +94,29 @@ func ResponseChunkData(c *gin.Context, resp dto.ResponsesStreamResponse, data st return FlushWriter(c) } +func ResponsesStreamError(c *gin.Context, openAIError types.OpenAIError) error { + payload := struct { + Type string `json:"type"` + Response struct { + Status string `json:"status"` + Error types.OpenAIError `json:"error"` + } `json:"response"` + }{ + Type: "response.failed", + } + payload.Response.Status = "failed" + payload.Response.Error = openAIError + + data, err := common.Marshal(payload) + if err != nil { + return fmt.Errorf("marshal responses stream error: %w", err) + } + SetEventStreamHeaders(c) + c.Render(-1, common.CustomEvent{Data: "event: response.failed\n"}) + c.Render(-1, common.CustomEvent{Data: "data: " + string(data)}) + return FlushWriter(c) +} + func StringData(c *gin.Context, str string) error { if c == nil || c.Writer == nil { return errors.New("context or writer is nil") diff --git a/relay/helper/common_test.go b/relay/helper/common_test.go new file mode 100644 index 000000000000..30a5bd6f5a33 --- /dev/null +++ b/relay/helper/common_test.go @@ -0,0 +1,34 @@ +package helper + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/QuantumNous/new-api/types" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" +) + +func TestResponsesStreamError_WritesFailureAfterStreamWasCommitted(t *testing.T) { + gin.SetMode(gin.TestMode) + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + SetEventStreamHeaders(c) + _, err := c.Writer.Write([]byte(": PING\n\n")) + require.NoError(t, err) + + err = ResponsesStreamError(c, types.OpenAIError{ + Message: "server is overloaded", + Type: "server_error", + Code: "server_is_overloaded", + }) + + require.NoError(t, err) + require.Equal(t, "text/event-stream", recorder.Header().Get("Content-Type")) + require.Contains(t, recorder.Body.String(), "event: response.failed") + require.Contains(t, recorder.Body.String(), `"status":"failed"`) + require.Contains(t, recorder.Body.String(), `"code":"server_is_overloaded"`) +} diff --git a/relay/helper/stream_scanner.go b/relay/helper/stream_scanner.go index 43ef2649c761..bf6cd44c4cba 100644 --- a/relay/helper/stream_scanner.go +++ b/relay/helper/stream_scanner.go @@ -151,6 +151,7 @@ func StreamScannerHandler(c *gin.Context, resp *http.Response, info *relaycommon if pingEnabled && pingTicker != nil { wg.Add(1) gopool.Go(func() { + defer wg.Done() defer func() { if r := recover(); r != nil { logger.LogError(c, fmt.Sprintf("ping goroutine panic: %v", r)) @@ -158,7 +159,6 @@ func StreamScannerHandler(c *gin.Context, resp *http.Response, info *relaycommon stop() } logger.LogDebug(c, "ping goroutine exited") - wg.Done() }() // 添加超时保护,防止 goroutine 无限运行 @@ -201,13 +201,13 @@ func StreamScannerHandler(c *gin.Context, resp *http.Response, info *relaycommon wg.Add(1) gopool.Go(func() { + defer wg.Done() defer func() { if r := recover(); r != nil { logger.LogError(c, fmt.Sprintf("data handler goroutine panic: %v", r)) info.StreamStatus.SetEndReason(relaycommon.StreamEndReasonPanic, fmt.Errorf("handler panic: %v", r)) } stop() - wg.Done() }() sr := newStreamResult(info.StreamStatus) for data := range dataChan { @@ -227,6 +227,7 @@ func StreamScannerHandler(c *gin.Context, resp *http.Response, info *relaycommon // Scanner goroutine with improved error handling wg.Add(1) common.RelayCtxGo(ctx, func() { + defer wg.Done() defer func() { close(dataChan) if r := recover(); r != nil { @@ -235,7 +236,6 @@ func StreamScannerHandler(c *gin.Context, resp *http.Response, info *relaycommon } stop() logger.LogDebug(c, "scanner goroutine exited") - wg.Done() }() for scanner.Scan() { diff --git a/service/channel.go b/service/channel.go index 81cd83f96bbe..a8f998ac3ec9 100644 --- a/service/channel.go +++ b/service/channel.go @@ -52,6 +52,9 @@ func ShouldDisableChannel(err *types.NewAPIError) bool { if err.GetErrorCode() == types.ErrorCodeUpstreamFirstResponseTimeout { return false } + if types.IsServerOverloadedError(err) { + return false + } if types.IsChannelError(err) { return true } diff --git a/service/channel_first_response_timeout_test.go b/service/channel_first_response_timeout_test.go index b7db5e76cf63..8a0a286fff63 100644 --- a/service/channel_first_response_timeout_test.go +++ b/service/channel_first_response_timeout_test.go @@ -6,6 +6,7 @@ import ( "testing" "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/setting/operation_setting" "github.com/QuantumNous/new-api/types" "github.com/stretchr/testify/require" ) @@ -24,3 +25,23 @@ func TestShouldDisableChannelSkipsFirstResponseTimeout(t *testing.T) { ) require.False(t, ShouldDisableChannel(timeoutErr)) } + +func TestShouldDisableChannelSkipsServerOverload(t *testing.T) { + previousEnabled := common.AutomaticDisableChannelEnabled + previousRanges := operation_setting.AutomaticDisableStatusCodeRanges + common.AutomaticDisableChannelEnabled = true + operation_setting.AutomaticDisableStatusCodeRanges = []operation_setting.StatusCodeRange{ + {Start: http.StatusServiceUnavailable, End: http.StatusServiceUnavailable}, + } + t.Cleanup(func() { + common.AutomaticDisableChannelEnabled = previousEnabled + operation_setting.AutomaticDisableStatusCodeRanges = previousRanges + }) + + overloadedErr := types.WithOpenAIError(types.OpenAIError{ + Message: "server is overloaded", + Type: "server_error", + Code: "server_is_overloaded", + }, http.StatusServiceUnavailable) + require.False(t, ShouldDisableChannel(overloadedErr)) +} diff --git a/types/error.go b/types/error.go index 2594f70c5df6..99aac9f9fa0d 100644 --- a/types/error.go +++ b/types/error.go @@ -371,6 +371,18 @@ func IsChannelError(err *NewAPIError) bool { return strings.HasPrefix(string(err.errorCode), "channel:") } +func IsServerOverloadedCode(code any) bool { + codeString := strings.ToLower(strings.TrimSpace(fmt.Sprint(code))) + return codeString == "server_is_overloaded" || codeString == "slow_down" +} + +func IsServerOverloadedError(err *NewAPIError) bool { + if err == nil { + return false + } + return IsServerOverloadedCode(err.ToOpenAIError().Code) +} + func IsSkipRetryError(err *NewAPIError) bool { if err == nil { return false