Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions controller/relay.go
Original file line number Diff line number Diff line change
Expand Up @@ -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())
Expand Down
1 change: 1 addition & 0 deletions dto/openai_response.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
101 changes: 89 additions & 12 deletions relay/channel/openai/relay_responses.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand Down Expand Up @@ -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":
Expand All @@ -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:
Expand All @@ -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()
Expand Down
139 changes: 139 additions & 0 deletions relay/channel/openai/relay_responses_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
})
}
}
8 changes: 8 additions & 0 deletions relay/common/relay_info.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
Expand Down
20 changes: 20 additions & 0 deletions relay/common/relay_info_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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())
}
23 changes: 23 additions & 0 deletions relay/helper/common.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
Loading