mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:18:18 +08:00
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:
@@ -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)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user