fix(openai): resume later websocket turns after 429

Port the safe current-turn replay from #4622 onto current main.

Co-authored-by: Kinso <kinsolee@users.noreply.github.com>
This commit is contained in:
Vincent Cui
2026-08-19 18:51:21 +08:00
co-authored by Kinso
parent ae62854abc
commit 82cbe6aff7
9 changed files with 461 additions and 20 deletions
@@ -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")
@@ -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)
}
@@ -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
@@ -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 ""
@@ -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 != "" {
@@ -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,
@@ -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)
}
@@ -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))
}
@@ -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)
})
}
}