From 82cbe6aff7d963e5096d26df73671236a707ad24 Mon Sep 17 00:00:00 2001 From: Vincent Cui Date: Wed, 19 Aug 2026 18:51:21 +0800 Subject: [PATCH] fix(openai): resume later websocket turns after 429 Port the safe current-turn replay from #4622 onto current main. Co-authored-by: Kinso --- .../handler/openai_gateway_handler.go | 28 +- .../openai_gateway_ws_failover_resume_test.go | 35 +++ .../service/openai_gateway_service.go | 5 +- .../internal/service/openai_ws_forwarder.go | 36 +++ .../service/openai_ws_forwarder_ingress.go | 36 +++ .../service/openai_ws_forwarder_payload.go | 27 ++ .../internal/service/openai_ws_http_bridge.go | 47 +++- .../openai_ws_http_bridge_resume_test.go | 252 ++++++++++++++++++ .../service/openai_ws_http_bridge_test.go | 15 +- 9 files changed, 461 insertions(+), 20 deletions(-) create mode 100644 backend/internal/handler/openai_gateway_ws_failover_resume_test.go create mode 100644 backend/internal/service/openai_ws_http_bridge_resume_test.go diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index 84b6417e5c..aebb68f543 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -1898,6 +1898,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { failedAccountIDs := make(map[int64]struct{}) var lastFailoverErr *service.UpstreamFailoverError var oauth429FailoverState service.OpenAIOAuth429FailoverState + wsAttemptMessage := append([]byte(nil), firstMessage...) handleWSFailover := func(account *service.Account, failoverErr *service.UpstreamFailoverError) bool { if ctx.Err() != nil { return false @@ -2306,7 +2307,7 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { }, } - wsFirstMessage := firstMessage + wsFirstMessage := wsAttemptMessage // 切组/会话失配防护:previous_response_id 未在当前分组命中粘连账号(StickyPreviousHit=false), // 说明该会话链不属于本次调度到的账号,原样转发会触发上游会话链鉴权失败(“鉴权失败,请检查 API Key”)。 // 故剥离首包里的 previous_response_id,改用首包内 input 重建上下文;带 function_call_output 的 @@ -2325,6 +2326,21 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) { if err := h.gatewayService.ProxyResponsesWebSocketFromClient(ctx, c, wsConn, account, token, wsFirstMessage, hooks); err != nil { var failoverErr *service.UpstreamFailoverError if errors.As(err, &failoverErr) { + retryPayload, retryCurrentTurn := service.OpenAIWSCurrentTurnRetryPayload(err) + nextAttemptMessage, retrySafe := openAIWSNextAttemptMessage(wsAttemptMessage, retryPayload, retryCurrentTurn) + if !retrySafe { + closeOpenAIWSFailoverExhausted(wsConn, failoverErr) + return + } + wsAttemptMessage = nextAttemptMessage + if retryCurrentTurn { + previousResponseID = "" + reqLog.Warn("openai.websocket_current_turn_failover_retry", + zap.Int64("account_id", account.ID), + zap.Int("upstream_status", failoverErr.StatusCode), + zap.Int("retry_payload_bytes", len(retryPayload)), + ) + } if handleWSFailover(account, failoverErr) { continue } @@ -2967,6 +2983,16 @@ func closeOpenAIClientWS(conn *coderws.Conn, status coderws.StatusCode, reason s _ = conn.CloseNow() } +func openAIWSNextAttemptMessage(current, retryPayload []byte, retryCurrentTurn bool) ([]byte, bool) { + if !retryCurrentTurn { + return append([]byte(nil), current...), true + } + if len(retryPayload) == 0 { + return nil, false + } + return append([]byte(nil), retryPayload...), true +} + func closeOpenAIWSFailoverExhausted(conn *coderws.Conn, failoverErr *service.UpstreamFailoverError) { if failoverErr == nil { closeOpenAIClientWS(conn, coderws.StatusInternalError, "upstream websocket proxy failed") diff --git a/backend/internal/handler/openai_gateway_ws_failover_resume_test.go b/backend/internal/handler/openai_gateway_ws_failover_resume_test.go new file mode 100644 index 0000000000..d62f83ee62 --- /dev/null +++ b/backend/internal/handler/openai_gateway_ws_failover_resume_test.go @@ -0,0 +1,35 @@ +package handler + +import ( + "testing" + + "github.com/stretchr/testify/require" +) + +func TestOpenAIWSNextAttemptMessageUsesCurrentTurnPayload(t *testing.T) { + firstMessage := []byte(`{"type":"response.create","input":"first"}`) + currentTurn := []byte(`{"type":"response.create","input":"turn-281"}`) + + next, ok := openAIWSNextAttemptMessage(firstMessage, currentTurn, true) + + require.True(t, ok) + require.Equal(t, currentTurn, next) + next[0] = 'x' + require.Equal(t, byte('{'), currentTurn[0], "retry payload must be cloned") +} + +func TestOpenAIWSNextAttemptMessageRejectsMissingCurrentTurnPayload(t *testing.T) { + next, ok := openAIWSNextAttemptMessage([]byte(`{"type":"response.create"}`), nil, true) + + require.False(t, ok) + require.Nil(t, next) +} + +func TestOpenAIWSNextAttemptMessageKeepsInitialMessageForFirstTurnFailover(t *testing.T) { + firstMessage := []byte(`{"type":"response.create","input":"first"}`) + + next, ok := openAIWSNextAttemptMessage(firstMessage, nil, false) + + require.True(t, ok) + require.Equal(t, firstMessage, next) +} diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 25b3333f5a..16082a16a4 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -284,8 +284,9 @@ type OpenAIForwardResult struct { // AudioUsage carries Voice billing units when present. AudioUsage *AudioUsage - wsReplayInput []json.RawMessage - wsReplayInputExists bool + wsReplayInput []json.RawMessage + wsReplayInputExists bool + wsAccountFailoverReplayInput []json.RawMessage } // SucceededForScheduling reports whether this result is an upstream success diff --git a/backend/internal/service/openai_ws_forwarder.go b/backend/internal/service/openai_ws_forwarder.go index 1928c1d6e4..0994446624 100644 --- a/backend/internal/service/openai_ws_forwarder.go +++ b/backend/internal/service/openai_ws_forwarder.go @@ -96,6 +96,42 @@ type openAIWSIngressTurnError struct { wroteDownstream bool } +type openAIWSCurrentTurnFailoverError struct { + cause error + retryPayload []byte +} + +func (e *openAIWSCurrentTurnFailoverError) Error() string { + if e == nil || e.cause == nil { + return "openai websocket current-turn failover" + } + return e.cause.Error() +} + +func (e *openAIWSCurrentTurnFailoverError) Unwrap() error { + if e == nil { + return nil + } + return e.cause +} + +func newOpenAIWSCurrentTurnFailoverError(cause error, retryPayload []byte) error { + return &openAIWSCurrentTurnFailoverError{ + cause: cause, + retryPayload: append([]byte(nil), retryPayload...), + } +} + +// OpenAIWSCurrentTurnRetryPayload returns an isolated copy of the payload that +// may be retried on a replacement account without replaying the first turn. +func OpenAIWSCurrentTurnRetryPayload(err error) ([]byte, bool) { + var retryErr *openAIWSCurrentTurnFailoverError + if !errors.As(err, &retryErr) || retryErr == nil { + return nil, false + } + return append([]byte(nil), retryErr.retryPayload...), true +} + func (e *openAIWSIngressTurnError) Error() string { if e == nil { return "" diff --git a/backend/internal/service/openai_ws_forwarder_ingress.go b/backend/internal/service/openai_ws_forwarder_ingress.go index bdd59f5e1c..4b0fe5a012 100644 --- a/backend/internal/service/openai_ws_forwarder_ingress.go +++ b/backend/internal/service/openai_ws_forwarder_ingress.go @@ -493,6 +493,8 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( grokCacheSeedPayload := firstPayload.payloadRaw var bridgeReplayInput []json.RawMessage bridgeReplayInputExists := false + var bridgeAccountFailoverInput []json.RawMessage + bridgeAccountFailoverInputExists := false for turn := 1; ; turn++ { if turn > 1 && hooks != nil && hooks.BeforeRequest != nil { if err := hooks.BeforeRequest(turn, currentBridgePayload.payloadRaw, currentBridgePayload.originalModel); err != nil { @@ -519,6 +521,15 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( if replayInputErr != nil { return fmt.Errorf("build websocket http bridge replay input: %w", replayInputErr) } + turnAccountFailoverInput, turnAccountFailoverInputExists, failoverInputErr := buildOpenAIWSReplayInputSequence( + bridgeAccountFailoverInput, + bridgeAccountFailoverInputExists, + currentBridgePayload.payloadRaw, + needsBridgeReplay, + ) + if failoverInputErr != nil { + return fmt.Errorf("build websocket account failover input: %w", failoverInputErr) + } if needsBridgeReplay && turnReplayInputExists { updatedPayload, setInputErr := setOpenAIWSPayloadInputSequence( currentBridgePayload.payloadRaw, @@ -571,6 +582,22 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( hooks.AfterTurn(turn, result, bridgeErr) } if bridgeErr != nil { + var failoverErr *UpstreamFailoverError + if turn > 1 && errors.As(bridgeErr, &failoverErr) && failoverErr != nil { + retryPayload, retrySafe, retryPayloadErr := buildOpenAIWSCurrentTurnRetryPayload( + bridgePayloadRaw, + turnAccountFailoverInput, + turnAccountFailoverInputExists, + currentBridgePayload.originalModel, + ) + if retryPayloadErr != nil { + return fmt.Errorf("build websocket current-turn failover payload: %w", retryPayloadErr) + } + if !retrySafe { + retryPayload = nil + } + return newOpenAIWSCurrentTurnFailoverError(bridgeErr, retryPayload) + } return bridgeErr } if result == nil { @@ -582,6 +609,15 @@ func (s *OpenAIGatewayService) ProxyResponsesWebSocketFromClient( bridgeReplayInput = append(bridgeReplayInput, cloneOpenAIWSRawMessages(result.wsReplayInput)...) bridgeReplayInputExists = true } + bridgeAccountFailoverInput = cloneOpenAIWSRawMessages(turnAccountFailoverInput) + bridgeAccountFailoverInputExists = turnAccountFailoverInputExists + if len(result.wsAccountFailoverReplayInput) > 0 { + bridgeAccountFailoverInput = append( + bridgeAccountFailoverInput, + cloneOpenAIWSRawMessages(result.wsAccountFailoverReplayInput)..., + ) + bridgeAccountFailoverInputExists = true + } if bridgeTurnState := strings.TrimSpace(result.ResponseHeaders.Get(openAIWSTurnStateHeader)); bridgeTurnState != "" { turnState = bridgeTurnState if stateStore != nil && sessionHash != "" { diff --git a/backend/internal/service/openai_ws_forwarder_payload.go b/backend/internal/service/openai_ws_forwarder_payload.go index 7eabe7d31f..3d6fef10be 100644 --- a/backend/internal/service/openai_ws_forwarder_payload.go +++ b/backend/internal/service/openai_ws_forwarder_payload.go @@ -674,6 +674,33 @@ func setOpenAIWSPayloadInputSequence( return sjson.SetRawBytes(payload, "input", inputRaw) } +func buildOpenAIWSCurrentTurnRetryPayload( + payload []byte, + fullInput []json.RawMessage, + fullInputExists bool, + originalModel string, +) ([]byte, bool, error) { + if !fullInputExists { + return nil, false, nil + } + retryPayload, err := setOpenAIWSPayloadInputSequence(payload, fullInput, true) + if err != nil { + return nil, false, err + } + retryPayload = RemovePreviousResponseIDFromBody(retryPayload) + if model := strings.TrimSpace(originalModel); model != "" { + retryPayload, err = sjson.SetBytes(retryPayload, "model", model) + if err != nil { + return nil, false, err + } + } + coverage := AnalyzeToolCallOutputContextCoverageBytes(retryPayload) + if coverage.HasFunctionCallOutput && !coverage.ContextCoversAllCallIDs { + return nil, false, nil + } + return retryPayload, true, nil +} + func shouldKeepIngressPreviousResponseID( previousPayload []byte, currentPayload []byte, diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go index 1ea3f2a935..1fab9be1a5 100644 --- a/backend/internal/service/openai_ws_http_bridge.go +++ b/backend/internal/service/openai_ws_http_bridge.go @@ -105,20 +105,25 @@ func prepareOpenAIWSHTTPBridgeBody(payload []byte) ([]byte, error) { } type openAIWSToolCallReplayCollector struct { - items []json.RawMessage - seen map[string]struct{} + items []json.RawMessage + seen map[string]struct{} + allItems []json.RawMessage + allSeen map[string]struct{} } func (c *openAIWSToolCallReplayCollector) AddEvent(eventType string, message []byte) { switch strings.TrimSpace(eventType) { case "response.output_item.done": - c.addItem(gjson.GetBytes(message, "item")) + item := gjson.GetBytes(message, "item") + c.addAllItem(item) + c.addItem(item) case "response.completed", "response.done": output := gjson.GetBytes(message, "response.output") if !output.IsArray() { return } for _, item := range output.Array() { + c.addAllItem(item) c.addItem(item) } } @@ -128,6 +133,35 @@ func (c *openAIWSToolCallReplayCollector) Items() []json.RawMessage { return cloneOpenAIWSRawMessages(c.items) } +func (c *openAIWSToolCallReplayCollector) AllItems() []json.RawMessage { + return cloneOpenAIWSRawMessages(c.allItems) +} + +func (c *openAIWSToolCallReplayCollector) addAllItem(item gjson.Result) { + if !item.Exists() || item.Type != gjson.JSON { + return + } + raw := strings.TrimSpace(item.Raw) + if raw == "" || !strings.HasPrefix(raw, "{") || strings.TrimSpace(item.Get("type").String()) == "" { + return + } + key := strings.TrimSpace(item.Get("id").String()) + if key == "" { + key = strings.TrimSpace(item.Get("call_id").String()) + } + if key == "" { + key = raw + } + if c.allSeen == nil { + c.allSeen = make(map[string]struct{}) + } + if _, ok := c.allSeen[key]; ok { + return + } + c.allSeen[key] = struct{}{} + c.allItems = append(c.allItems, json.RawMessage(raw)) +} + func (c *openAIWSToolCallReplayCollector) addItem(item gjson.Result) { if !item.Exists() || item.Type != gjson.JSON { return @@ -303,10 +337,10 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( if account.Platform == PlatformGrok { shouldFailover = s.shouldFailoverGrokUpstreamError(resp.StatusCode, respBody) s.handleGrokAccountUpstreamError(withGrokTeamRateLimitModel(ctx, resolveGrokWSUpstreamModel(account, body, originalModel)), account, resp.StatusCode, resp.Header, respBody) - if turn == 1 && shouldFailover { + if shouldFailover && (turn == 1 || resp.StatusCode == http.StatusTooManyRequests) { return nil, newOpenAIUpstreamFailoverError(resp.StatusCode, resp.Header, respBody, upstreamMsg, false) } - } else if turn == 1 && shouldFailover { + } else if shouldFailover && (turn == 1 || resp.StatusCode == http.StatusTooManyRequests) { return nil, s.handleFailoverErrorResponsePassthrough(ctx, resp, c, account, body, respBody) } if account.Platform != PlatformGrok && (shouldFailover || shouldCooldownOpenAITransientUpstreamError(resp.StatusCode, respBody)) { @@ -374,6 +408,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( result.wsReplayInput = replayInput result.wsReplayInputExists = true } + result.wsAccountFailoverReplayInput = replayCollector.AllItems() if imageCount > 0 { result.ImageCount = imageCount result.ImageSize = imageSizeTier @@ -482,7 +517,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn( canonicalModel := canonicalOpenAIAccountSchedulingModel(account, originalModel) s.handleOpenAIAccountUpstreamError(ctx, account, accountStatus, resp.Header, upstreamMessage, canonicalModel) } - if turn == 1 && !wroteDownstream && shouldFailover { + if !wroteDownstream && shouldFailover && (turn == 1 || statusCode == http.StatusTooManyRequests) { if account.Platform == PlatformGrok { return nil, newOpenAIUpstreamFailoverError(statusCode, resp.Header, upstreamMessage, errMessage, false) } diff --git a/backend/internal/service/openai_ws_http_bridge_resume_test.go b/backend/internal/service/openai_ws_http_bridge_resume_test.go new file mode 100644 index 0000000000..b6719d493b --- /dev/null +++ b/backend/internal/service/openai_ws_http_bridge_resume_test.go @@ -0,0 +1,252 @@ +package service + +import ( + "context" + "encoding/json" + "errors" + "io" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/Wei-Shaw/sub2api/internal/config" + "github.com/coder/websocket" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/require" + "github.com/tidwall/gjson" +) + +func TestBuildOpenAIWSCurrentTurnRetryPayloadRejectsOrphanToolOutput(t *testing.T) { + payload := []byte(`{"type":"response.create","model":"mapped-model","previous_response_id":"resp_old"}`) + fullInput := []json.RawMessage{ + json.RawMessage(`{"type":"function_call_output","call_id":"missing_call","output":"done"}`), + } + + retryPayload, retrySafe, err := buildOpenAIWSCurrentTurnRetryPayload(payload, fullInput, true, "gpt-5.6-sol") + + require.NoError(t, err) + require.False(t, retrySafe) + require.Nil(t, retryPayload) +} + +func TestProxyOpenAIWSHTTPBridgeTurnLaterTurn429FailsOverBeforeClientWrite(t *testing.T) { + gin.SetMode(gin.TestMode) + + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusTooManyRequests, + Header: http.Header{"Retry-After": []string{"60"}}, + Body: io.NopCloser(strings.NewReader(`{"error":{"type":"usage_limit_reached","message":"The usage limit has been reached"}}`)), + }} + svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} + account := &Account{ID: 129, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Concurrency: 1} + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) + payload := []byte(`{"type":"response.create","model":"gpt-5.6-sol","previous_response_id":"resp_old","input":[{"role":"user","content":"continue"}]}`) + writes := 0 + + result, err := svc.proxyOpenAIWSHTTPBridgeTurn( + context.Background(), c, account, "access-token", payload, len(payload), + "gpt-5.6-sol", "", "", "", "", 281, + func([]byte) error { + writes++ + return nil + }, + ) + + require.Nil(t, result) + var failoverErr *UpstreamFailoverError + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, http.StatusTooManyRequests, failoverErr.StatusCode) + require.Zero(t, writes) +} + +func TestProxyOpenAIWSHTTPBridgeTurnLaterTurnDoesNotFailOverAfterDownstreamOutput(t *testing.T) { + gin.SetMode(gin.TestMode) + + upstream := &httpUpstreamRecorder{resp: &http.Response{ + StatusCode: http.StatusOK, + Header: make(http.Header), + Body: io.NopCloser(strings.NewReader( + "data: {\"type\":\"response.output_text.delta\",\"delta\":\"partial\"}\n\n" + + "data: {\"type\":\"error\",\"error\":{\"type\":\"rate_limit_error\",\"message\":\"limited\"}}\n\n", + )), + }} + svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream} + account := &Account{ID: 10, Platform: PlatformOpenAI, Type: AccountTypeAPIKey, Concurrency: 1} + recorder := httptest.NewRecorder() + c, _ := gin.CreateTestContext(recorder) + c.Request = httptest.NewRequest(http.MethodGet, "/v1/responses", nil) + payload := []byte(`{"type":"response.create","model":"gpt-5","input":"hi"}`) + var writes [][]byte + + result, err := svc.proxyOpenAIWSHTTPBridgeTurn( + context.Background(), c, account, "sk-test", payload, len(payload), + "gpt-5", "", "", "", "", 281, + func(message []byte) error { + writes = append(writes, append([]byte(nil), message...)) + return nil + }, + ) + + require.NotNil(t, result) + require.Error(t, err) + var failoverErr *UpstreamFailoverError + require.False(t, errors.As(err, &failoverErr)) + require.Len(t, writes, 2) + require.Equal(t, "response.output_text.delta", gjson.GetBytes(writes[0], "type").String()) + require.Equal(t, "error", gjson.GetBytes(writes[1], "type").String()) +} + +func TestOpenAIWSHTTPBridgeLaterTurn429RetriesCurrentTurnOnReplacementAccount(t *testing.T) { + gin.SetMode(gin.TestMode) + + cfg := &config.Config{} + cfg.Security.URLAllowlist.Enabled = false + cfg.Security.URLAllowlist.AllowInsecureHTTP = true + cfg.Gateway.OpenAIWS.Enabled = true + cfg.Gateway.OpenAIWS.OAuthEnabled = true + cfg.Gateway.OpenAIWS.ResponsesWebsocketsV2 = true + cfg.Gateway.OpenAIWS.ModeRouterV2Enabled = true + cfg.Gateway.OpenAIWS.ReadTimeoutSeconds = 3 + cfg.Gateway.OpenAIWS.WriteTimeoutSeconds = 3 + + upstream := &httpUpstreamRecorder{responses: []*http.Response{ + { + StatusCode: http.StatusOK, + Header: http.Header{ + "Content-Type": []string{"text/event-stream"}, + openAIWSTurnStateHeader: []string{"old-account-state"}, + }, + Body: io.NopCloser(strings.NewReader( + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_first\",\"output\":[{\"id\":\"msg_1\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[{\"type\":\"output_text\",\"text\":\"first-ok\"}]},{\"id\":\"fc_1\",\"type\":\"function_call\",\"call_id\":\"call_1\",\"name\":\"inspect\",\"arguments\":\"{}\"}],\"usage\":{\"input_tokens\":1,\"output_tokens\":1}}}\n\n", + )), + }, + { + StatusCode: http.StatusTooManyRequests, + Header: http.Header{"Retry-After": []string{"60"}}, + Body: io.NopCloser(strings.NewReader(`{"error":{"type":"usage_limit_reached","message":"The usage limit has been reached"}}`)), + }, + { + StatusCode: http.StatusOK, + Header: http.Header{"Content-Type": []string{"text/event-stream"}}, + Body: io.NopCloser(strings.NewReader( + "data: {\"type\":\"response.completed\",\"response\":{\"id\":\"resp_second\",\"output\":[{\"id\":\"msg_2\",\"type\":\"message\",\"role\":\"assistant\",\"content\":[{\"type\":\"output_text\",\"text\":\"second-ok\"}]}],\"usage\":{\"input_tokens\":4,\"output_tokens\":1}}}\n\n", + )), + }, + }} + svc := &OpenAIGatewayService{ + cfg: cfg, + httpUpstream: upstream, + cache: &stubGatewayCache{}, + openaiWSResolver: NewOpenAIWSProtocolResolver(cfg), + toolCorrector: NewCodexToolCorrector(), + } + account := &Account{ + ID: 129, Name: "limited", Platform: PlatformOpenAI, Type: AccountTypeOAuth, + Status: StatusActive, Schedulable: true, Concurrency: 1, + Extra: map[string]any{"openai_oauth_responses_websockets_v2_mode": OpenAIWSIngressModeHTTPBridge}, + } + nextAccount := *account + nextAccount.ID = 130 + nextAccount.Name = "replacement" + + serverErrCh := make(chan error, 1) + failoverCh := make(chan []byte, 1) + wsServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + conn, err := websocket.Accept(w, r, nil) + if err != nil { + serverErrCh <- err + return + } + defer func() { _ = conn.CloseNow() }() + + rec := httptest.NewRecorder() + ginCtx, _ := gin.CreateTestContext(rec) + ginCtx.Request = r.Clone(r.Context()) + readCtx, cancel := context.WithTimeout(r.Context(), 3*time.Second) + _, firstMessage, readErr := conn.Read(readCtx) + cancel() + if readErr != nil { + serverErrCh <- readErr + return + } + proxyErr := svc.ProxyResponsesWebSocketFromClient(r.Context(), ginCtx, conn, account, "access-token-a", firstMessage, nil) + var failoverErr *UpstreamFailoverError + if !errors.As(proxyErr, &failoverErr) { + serverErrCh <- proxyErr + return + } + retryPayload, retryCurrentTurn := OpenAIWSCurrentTurnRetryPayload(proxyErr) + if !retryCurrentTurn || len(retryPayload) == 0 { + serverErrCh <- errors.New("missing current-turn retry payload") + return + } + failoverCh <- retryPayload + serverErrCh <- svc.ProxyResponsesWebSocketFromClient( + r.Context(), ginCtx, conn, &nextAccount, "access-token-b", retryPayload, nil, + ) + })) + defer wsServer.Close() + + dialCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + clientConn, _, err := websocket.Dial(dialCtx, "ws"+strings.TrimPrefix(wsServer.URL, "http"), nil) + cancel() + require.NoError(t, err) + defer func() { _ = clientConn.CloseNow() }() + + writeCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + err = clientConn.Write(writeCtx, websocket.MessageText, []byte(`{"type":"response.create","model":"gpt-5.6-sol","input":[{"role":"user","content":"first"}]}`)) + cancel() + require.NoError(t, err) + + readCtx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + _, completed, err := clientConn.Read(readCtx) + cancel() + require.NoError(t, err) + require.Equal(t, "response.completed", gjson.GetBytes(completed, "type").String()) + + writeCtx, cancel = context.WithTimeout(context.Background(), 3*time.Second) + err = clientConn.Write(writeCtx, websocket.MessageText, []byte(`{"type":"response.create","model":"gpt-5.6-sol","previous_response_id":"resp_first","input":[{"type":"function_call_output","call_id":"call_1","output":"second"}]}`)) + cancel() + require.NoError(t, err) + + readCtx, cancel = context.WithTimeout(context.Background(), 3*time.Second) + _, retriedCompleted, err := clientConn.Read(readCtx) + cancel() + require.NoError(t, err) + require.Equal(t, "response.completed", gjson.GetBytes(retriedCompleted, "type").String()) + require.Equal(t, "resp_second", gjson.GetBytes(retriedCompleted, "response.id").String()) + _ = clientConn.Close(websocket.StatusNormalClosure, "done") + + select { + case retryPayload := <-failoverCh: + require.NotEmpty(t, retryPayload) + require.False(t, gjson.GetBytes(retryPayload, "previous_response_id").Exists()) + require.Equal(t, "gpt-5.6-sol", gjson.GetBytes(retryPayload, "model").String()) + input := gjson.GetBytes(retryPayload, "input") + require.True(t, input.IsArray()) + require.Len(t, input.Array(), 4) + require.Contains(t, input.Raw, "first") + require.Contains(t, input.Raw, "first-ok") + require.Contains(t, input.Raw, "second") + require.Equal(t, 1, strings.Count(input.Raw, `"id":"fc_1"`)) + require.Equal(t, 2, strings.Count(input.Raw, `"call_id":"call_1"`)) + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for current-turn failover") + } + + select { + case proxyErr := <-serverErrCh: + require.NoError(t, proxyErr) + case <-time.After(5 * time.Second): + t.Fatal("timed out waiting for replacement-account completion") + } + require.Len(t, upstream.bodies, 3) + require.Contains(t, string(upstream.bodies[0]), "first") + require.NotContains(t, string(upstream.bodies[2]), "previous_response_id") + require.Contains(t, string(upstream.bodies[2]), "second") + require.Empty(t, upstream.requests[2].Header.Get(openAIWSTurnStateHeader)) +} diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go index b283c36ce6..9617bd464d 100644 --- a/backend/internal/service/openai_ws_http_bridge_test.go +++ b/backend/internal/service/openai_ws_http_bridge_test.go @@ -452,17 +452,10 @@ func TestProxyOpenAIWSHTTPBridgeTurnSSEErrorFailoverSafety(t *testing.T) { ) var failoverErr *UpstreamFailoverError - if turn == 1 { - require.Nil(t, result) - require.ErrorAs(t, err, &failoverErr) - require.Equal(t, http.StatusTooManyRequests, failoverErr.StatusCode) - require.Empty(t, writes) - } else { - require.NotNil(t, result) - require.Error(t, err) - require.False(t, errors.As(err, &failoverErr)) - require.Len(t, writes, 1) - } + require.Nil(t, result) + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, http.StatusTooManyRequests, failoverErr.StatusCode) + require.Empty(t, writes) }) } }