mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 17:08:33 +08:00
Merge pull request #5398 from Wei-Shaw/fix/openai-capacity-shed-stream-recovery
fix(gateway): 流内降载错误恢复 pre-output failover 并对客户端改写为可重试错误码
This commit is contained in:
@@ -2,12 +2,17 @@ package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/openai"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
@@ -72,6 +77,172 @@ func TestStreamFailedEventCapacityShedRetriesOnSameAccount(t *testing.T) {
|
||||
require.False(t, openAIStreamFailedEventRetryableOnSameAccount(nonPool, other, "boom"))
|
||||
}
|
||||
|
||||
// 上游降载的真实序列是「event: error → event: response.failed」。error 帧不算
|
||||
// 客户端输出:若把它当首输出 flush,clientOutputStarted 被固化,随后的 failed
|
||||
// 事件就进不了 pre-output failover 分支,只能把致命错误原样转发给客户端。
|
||||
func TestOpenAIStreamErrorFrameDoesNotStartClientOutput(t *testing.T) {
|
||||
cases := []struct {
|
||||
data string
|
||||
eventType string
|
||||
want bool
|
||||
}{
|
||||
{`{"type":"error","error":{"code":"server_is_overloaded","message":"overloaded"}}`, "error", false},
|
||||
{`{"type":"error","error":{"code":"slow_down","message":"slow down"}}`, "error", false},
|
||||
{`{"type":"error","error":{"type":"rate_limit_error","code":"rate_limit_exceeded","message":"limited"}}`, "error", false},
|
||||
// 不可重试类错误帧维持原样转发(不进 failover),保留上游错误细节。
|
||||
{`{"type":"error","error":{"type":"invalid_request_error","code":"content_policy_violation","message":"blocked"}}`, "error", true},
|
||||
{`{"type":"response.failed","response":{"error":{"code":"server_is_overloaded"}}}`, "response.failed", false},
|
||||
{`{"type":"response.created","response":{"id":"resp_1"}}`, "response.created", false},
|
||||
{`{"type":"response.in_progress","response":{"id":"resp_1"}}`, "response.in_progress", false},
|
||||
{`{"type":"response.output_text.delta","delta":"hi"}`, "response.output_text.delta", true},
|
||||
{`[DONE]`, "", true},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
require.Equal(t, tc.want, openAIStreamDataStartsClientOutput(tc.data, tc.eventType), "data=%s type=%s", tc.data, tc.eventType)
|
||||
}
|
||||
}
|
||||
|
||||
// 回归用例(真实上游降载序列):created → in_progress → error 帧 → response.failed。
|
||||
// 期望仍然走 pre-output failover(同账号重试 + 请求级瞬时标记),且不向客户端写出任何字节。
|
||||
func TestOpenAIStreamCapacityShedErrorFramePrecedingFailedStillFailsOver(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
cfg := &config.Config{
|
||||
Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize},
|
||||
}
|
||||
svc := &OpenAIGatewayService{cfg: cfg}
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/", nil)
|
||||
|
||||
resp := &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Body: io.NopCloser(strings.NewReader(strings.Join([]string{
|
||||
"event: response.created",
|
||||
`data: {"type":"response.created","response":{"id":"resp_1"},"sequence_number":0}`,
|
||||
"",
|
||||
"event: response.in_progress",
|
||||
`data: {"type":"response.in_progress","response":{"id":"resp_1"},"sequence_number":1}`,
|
||||
"",
|
||||
"event: error",
|
||||
`data: {"type":"error","error":{"type":"service_unavailable_error","code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."},"sequence_number":2}`,
|
||||
"",
|
||||
"event: response.failed",
|
||||
`data: {"type":"response.failed","response":{"id":"resp_1","status":"failed","error":{"code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."}},"sequence_number":3}`,
|
||||
"",
|
||||
}, "\n"))),
|
||||
Header: http.Header{"X-Request-Id": []string{"rid-shed-error-then-failed"}},
|
||||
}
|
||||
|
||||
_, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Name: "acc"}, time.Now(), "model", "model")
|
||||
require.Error(t, err)
|
||||
var failoverErr *UpstreamFailoverError
|
||||
require.ErrorAs(t, err, &failoverErr)
|
||||
require.True(t, failoverErr.RetryableOnSameAccount)
|
||||
require.True(t, failoverErr.RequestScopedTransient)
|
||||
require.False(t, c.Writer.Written())
|
||||
require.Empty(t, rec.Body.String())
|
||||
}
|
||||
|
||||
// 流中途(已有真实输出)降载时无法再 failover,此时必须把降载码改写为客户端
|
||||
// 可重试的 server_error 再转发——Codex 对 server_is_overloaded/slow_down 判致命
|
||||
// 并终止会话,对其余错误码执行内置退避重试。消息原样保留。
|
||||
func TestOpenAIStreamCapacityShedAfterOutputRewritesCodeForClient(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
cfg := &config.Config{
|
||||
Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize},
|
||||
}
|
||||
svc := &OpenAIGatewayService{cfg: cfg}
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/", nil)
|
||||
|
||||
resp := &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Body: io.NopCloser(strings.NewReader(strings.Join([]string{
|
||||
"event: response.created",
|
||||
`data: {"type":"response.created","response":{"id":"resp_1"}}`,
|
||||
"",
|
||||
"event: response.output_text.delta",
|
||||
`data: {"type":"response.output_text.delta","delta":"partial"}`,
|
||||
"",
|
||||
"event: error",
|
||||
`data: {"type":"error","error":{"type":"service_unavailable_error","code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."},"sequence_number":2}`,
|
||||
"",
|
||||
"event: response.failed",
|
||||
`data: {"type":"response.failed","response":{"id":"resp_1","status":"failed","error":{"code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."}},"sequence_number":3}`,
|
||||
"",
|
||||
}, "\n"))),
|
||||
Header: http.Header{"X-Request-Id": []string{"rid-shed-after-output"}},
|
||||
}
|
||||
|
||||
_, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Name: "acc"}, time.Now(), "model", "model")
|
||||
require.Error(t, err)
|
||||
var failoverErr *UpstreamFailoverError
|
||||
require.False(t, errors.As(err, &failoverErr))
|
||||
|
||||
body := rec.Body.String()
|
||||
require.Contains(t, body, "partial")
|
||||
require.Contains(t, body, "event: response.failed")
|
||||
require.Contains(t, body, `"code":"server_error"`)
|
||||
require.NotContains(t, body, "server_is_overloaded")
|
||||
require.Contains(t, body, "Our servers are currently overloaded")
|
||||
}
|
||||
|
||||
// helper 单测:只有降载码被改写,其余错误码(尤其 rate_limit_exceeded,客户端
|
||||
// 依赖其原码解析重试延时)必须原样保留。
|
||||
func TestSanitizeOpenAICapacityShedErrorCodeForClient(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
payload string
|
||||
wantChanged bool
|
||||
wantContain string
|
||||
}{
|
||||
{
|
||||
name: "failed事件嵌套code改写",
|
||||
payload: `{"type":"response.failed","response":{"error":{"code":"server_is_overloaded","message":"overloaded"}}}`,
|
||||
wantChanged: true,
|
||||
wantContain: `"code":"server_error"`,
|
||||
},
|
||||
{
|
||||
name: "error帧裸code改写",
|
||||
payload: `{"type":"error","error":{"code":"slow_down","message":"slow down"}}`,
|
||||
wantChanged: true,
|
||||
wantContain: `"code":"server_error"`,
|
||||
},
|
||||
{
|
||||
name: "rate_limit不改写",
|
||||
payload: `{"type":"response.failed","response":{"error":{"code":"rate_limit_exceeded","message":"try again in 3s"}}}`,
|
||||
wantChanged: false,
|
||||
wantContain: `"code":"rate_limit_exceeded"`,
|
||||
},
|
||||
{
|
||||
name: "普通server_error不改写",
|
||||
payload: `{"type":"response.failed","response":{"error":{"code":"server_error","message":"boom"}}}`,
|
||||
wantChanged: false,
|
||||
wantContain: `"code":"server_error"`,
|
||||
},
|
||||
{
|
||||
name: "非JSON不改写",
|
||||
payload: `not-json`,
|
||||
wantChanged: false,
|
||||
wantContain: `not-json`,
|
||||
},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
out, changed := sanitizeOpenAICapacityShedErrorCodeForClient([]byte(tc.payload))
|
||||
require.Equal(t, tc.wantChanged, changed)
|
||||
require.Contains(t, string(out), tc.wantContain)
|
||||
if changed {
|
||||
require.NotContains(t, string(out), "server_is_overloaded")
|
||||
require.NotContains(t, string(out), "slow_down")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// 出站身份的版本声明只能有一个来源:UA 的版本段、version 头、探针版本三处必须同源,
|
||||
// 各自硬编码会漂移成互相矛盾的身份,而自相矛盾或陈旧的身份会被上游优先降载。
|
||||
func TestCodexOutboundVersionHasSingleSource(t *testing.T) {
|
||||
|
||||
@@ -759,8 +759,17 @@ func openAIStreamDataStartsClientOutput(data, eventType string) bool {
|
||||
if trimmed == "" {
|
||||
return false
|
||||
}
|
||||
if strings.TrimSpace(eventType) == "response.failed" {
|
||||
switch strings.TrimSpace(eventType) {
|
||||
case "response.failed":
|
||||
return false
|
||||
case "error":
|
||||
// 上游降载/瞬时故障会先推 {"type":"error"} 帧、再以 response.failed 收尾。
|
||||
// 可重试类错误帧不能算客户端输出:一旦把它当首输出 flush,
|
||||
// clientOutputStarted 即被固化,随后的 failed 事件永远进不了 pre-output
|
||||
// failover 分支,只能把致命错误原样转发给客户端。不可重试类
|
||||
// (content_policy / invalid_request 等)维持原样转发,保留上游错误细节。
|
||||
payload := []byte(trimmed)
|
||||
return !openAIStreamFailedEventShouldFailover(payload, extractOpenAISSEErrorMessage(payload))
|
||||
}
|
||||
return !openAIStreamEventIsPreamble(eventType)
|
||||
}
|
||||
@@ -787,6 +796,41 @@ func isOpenAIUpstreamCapacityShedEvent(payload []byte) bool {
|
||||
}
|
||||
}
|
||||
|
||||
// openAICapacityShedRetryableClientCode 是把上游容量降载错误转发给客户端时改写
|
||||
// 使用的错误码。Codex CLI 按闭集对错误码分类:server_is_overloaded / slow_down
|
||||
// 被判为致命错误(客户端提示 "Selected model is at capacity. Please try a
|
||||
// different model." 并直接终止会话),而 server_error 等致命集之外的错误码会进入
|
||||
// 客户端内置的退避重试。
|
||||
const openAICapacityShedRetryableClientCode = "server_error"
|
||||
|
||||
// sanitizeOpenAICapacityShedErrorCodeForClient 把即将写给下游客户端的
|
||||
// error / response.failed 事件中的容量降载错误码改写为客户端可重试的错误码。
|
||||
// 走到转发这一步说明网关侧 failover 已不可用(流中途)或已用尽;保留原始降载码
|
||||
// 只会让客户端就地终止会话。错误消息原样保留;监控与账号状态判定都基于改写前
|
||||
// 的原始 payload,不受影响。rate_limit 等其他错误码一律不动(客户端依赖
|
||||
// rate_limit_exceeded 原码解析重试延时)。
|
||||
func sanitizeOpenAICapacityShedErrorCodeForClient(payload []byte) ([]byte, bool) {
|
||||
if len(payload) == 0 || !gjson.ValidBytes(payload) || !isOpenAIUpstreamCapacityShedEvent(payload) {
|
||||
return payload, false
|
||||
}
|
||||
updated := payload
|
||||
changed := false
|
||||
for _, path := range []string{"response.error.code", "error.code"} {
|
||||
switch strings.ToLower(strings.TrimSpace(gjson.GetBytes(updated, path).String())) {
|
||||
case "server_is_overloaded", "slow_down":
|
||||
default:
|
||||
continue
|
||||
}
|
||||
next, err := sjson.SetBytes(updated, path, openAICapacityShedRetryableClientCode)
|
||||
if err != nil {
|
||||
return payload, false
|
||||
}
|
||||
updated = next
|
||||
changed = true
|
||||
}
|
||||
return updated, changed
|
||||
}
|
||||
|
||||
func openAIStreamFailedEventSemanticStatus(payload []byte, message string) int {
|
||||
if isOpenAIContextWindowError(message, payload) {
|
||||
return http.StatusBadRequest
|
||||
|
||||
@@ -387,13 +387,30 @@ func TestOpenAIResponseFlush_FailedAndErrorEventsFlushAtBoundaries(t *testing.T)
|
||||
require.Contains(t, flushes[1], "response.failed")
|
||||
})
|
||||
|
||||
t.Run("error event", func(t *testing.T) {
|
||||
t.Run("retryable error event buffered until terminal", func(t *testing.T) {
|
||||
// 可重试类 error 帧不算客户端输出:保持在 attempt 缓冲中不单独 flush,
|
||||
// 为随后可能到达的 response.failed 保留 pre-output failover 能力,
|
||||
// 与终止帧一起出站。
|
||||
body := "data: {\"type\":\"error\",\"error\":{\"message\":\"failed\"}}\n\n" +
|
||||
"data: [DONE]\n\n"
|
||||
recorder := newOpenAIResponseFlushRecorder()
|
||||
|
||||
result, err := runOpenAIResponseFlushTest(recorder, io.NopCloser(strings.NewReader(body)), config.GatewayConfig{})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
gotBody, flushes := recorder.snapshot()
|
||||
require.Equal(t, body, gotBody)
|
||||
require.Len(t, flushes, 1)
|
||||
})
|
||||
|
||||
t.Run("non-retryable error event flushes at boundary", func(t *testing.T) {
|
||||
body := "data: {\"type\":\"error\",\"error\":{\"code\":\"invalid_request\",\"message\":\"bad request\"}}\n\n" +
|
||||
"data: [DONE]\n\n"
|
||||
recorder := newOpenAIResponseFlushRecorder()
|
||||
|
||||
result, err := runOpenAIResponseFlushTest(recorder, io.NopCloser(strings.NewReader(body)), config.GatewayConfig{})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, result)
|
||||
gotBody, flushes := recorder.snapshot()
|
||||
|
||||
@@ -1324,10 +1324,21 @@ func extractOpenAISSEErrorMessage(payload []byte) string {
|
||||
}
|
||||
|
||||
func sanitizeOpenAIResponseFailedEventForClient(payload []byte, eventType string, clientOutputStarted bool) ([]byte, bool) {
|
||||
if eventType != "response.failed" || len(payload) == 0 || !gjson.ValidBytes(payload) {
|
||||
eventType = strings.TrimSpace(eventType)
|
||||
isFailedEvent := eventType == "response.failed"
|
||||
if (!isFailedEvent && eventType != "error") || len(payload) == 0 || !gjson.ValidBytes(payload) {
|
||||
return payload, false
|
||||
}
|
||||
updated := payload
|
||||
// 容量降载码对 Codex CLI 是致命错误;事件既然要写给客户端(failover 已不可用),
|
||||
// 就改写为客户端可重试的错误码。error 帧与 response.failed 都要改:上游降载
|
||||
// 总是先推 error 帧再收 failed,两帧携带同一个错误。
|
||||
if rewritten, changed := sanitizeOpenAICapacityShedErrorCodeForClient(updated); changed {
|
||||
updated = rewritten
|
||||
}
|
||||
if !isFailedEvent {
|
||||
return updated, !bytes.Equal(updated, payload)
|
||||
}
|
||||
if clientOutputStarted && isOpenAIContextWindowError(extractOpenAISSEErrorMessage(payload), payload) {
|
||||
errorPath := ""
|
||||
switch {
|
||||
|
||||
@@ -425,8 +425,18 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
|
||||
upstreamEventErr = errors.New(errMessage)
|
||||
}
|
||||
|
||||
// 客户端写出副本改写容量降载码:Codex 对 error/response.failed 中的
|
||||
// server_is_overloaded / slow_down 判致命并终止会话,改写后走客户端内置
|
||||
// 重试。账号状态与终止事件判定(下方 handleOpenAIWSTerminalTransientFailure)
|
||||
// 仍使用未改写的 upstreamMessage。
|
||||
clientMessage := upstreamMessage
|
||||
if eventType == "error" || eventType == "response.failed" {
|
||||
if rewritten, changed := sanitizeOpenAICapacityShedErrorCodeForClient(clientMessage); changed {
|
||||
clientMessage = rewritten
|
||||
}
|
||||
}
|
||||
if !clientDisconnected {
|
||||
if err := writeClientMessage(upstreamMessage); err != nil {
|
||||
if err := writeClientMessage(clientMessage); err != nil {
|
||||
if isOpenAIWSClientDisconnectError(err) {
|
||||
clientDisconnected = true
|
||||
closeStatus, closeReason := summarizeOpenAIWSReadCloseError(err)
|
||||
|
||||
@@ -220,6 +220,69 @@ func TestProxyOpenAIWSHTTPBridgeTurnSSEErrorFailoverSafety(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// 桥接转发 error / response.failed 给 WS 客户端前必须把容量降载码改写为可重试
|
||||
// 的 server_error:Codex 对 server_is_overloaded/slow_down 判致命并终止会话。
|
||||
// 账号状态判定使用改写前的原始事件,不受影响。
|
||||
func TestProxyOpenAIWSHTTPBridgeTurnRewritesCapacityShedCodeForClient(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
turn int
|
||||
body string
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "turn2_error_frame",
|
||||
turn: 2,
|
||||
body: "data: {\"type\":\"error\",\"error\":{\"type\":\"service_unavailable_error\",\"code\":\"server_is_overloaded\",\"message\":\"Our servers are currently overloaded. Please try again later.\"}}\n\n",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
// response.failed 不走 error 事件分支:即便 turn 1 也会被当终止事件
|
||||
// 原样转发(不 failover),因此改写必须在这里同样生效。
|
||||
name: "turn1_bare_response_failed",
|
||||
turn: 1,
|
||||
body: "data: {\"type\":\"response.failed\",\"response\":{\"id\":\"resp_shed\",\"status\":\"failed\",\"error\":{\"code\":\"server_is_overloaded\",\"message\":\"Our servers are currently overloaded. Please try again later.\"}}}\n\n",
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
upstream := &httpUpstreamRecorder{resp: &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: make(http.Header),
|
||||
Body: io.NopCloser(strings.NewReader(tt.body)),
|
||||
}}
|
||||
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
|
||||
account := &Account{ID: 11, 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","input":"hi"}`)
|
||||
var writes [][]byte
|
||||
|
||||
_, err := svc.proxyOpenAIWSHTTPBridgeTurn(
|
||||
context.Background(), c, account, "sk-test", payload, len(payload),
|
||||
"gpt-5", "", "", "", "", tt.turn,
|
||||
func(message []byte) error {
|
||||
writes = append(writes, append([]byte(nil), message...))
|
||||
return nil
|
||||
},
|
||||
)
|
||||
|
||||
if tt.wantErr {
|
||||
require.Error(t, err)
|
||||
} else {
|
||||
require.NoError(t, err)
|
||||
}
|
||||
require.Len(t, writes, 1)
|
||||
require.Contains(t, string(writes[0]), `"code":"server_error"`)
|
||||
require.NotContains(t, string(writes[0]), "server_is_overloaded")
|
||||
require.Contains(t, string(writes[0]), "Our servers are currently overloaded")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxyOpenAIWSHTTPBridgeTurnRequiresTerminalEvent(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user