Merge pull request #5676 from Perfecto23/agent/openai-capacity-failover

fix(openai): recover message-only capacity failures before output
This commit is contained in:
Wesley Liddick
2026-08-19 11:11:23 +08:00
committed by GitHub
14 changed files with 625 additions and 90 deletions
@@ -574,7 +574,7 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
// Forward request
service.SetOpsLatencyMs(c, service.OpsRoutingLatencyMsKey, time.Since(routingStart).Milliseconds())
forwardStart := time.Now()
// 用扣除 compact 心跳字节的口径快照:心跳注释不构成语义响应,
// 用扣除非语义心跳字节的口径快照:心跳注释不构成语义响应,
// 不能因心跳字节变化而放弃 failover 换号(#3887)。
writerSizeBeforeForward := service.OpenAICompactKeepaliveAdjustedWrittenSize(c)
// 跨 passthrough 边界的 failover:从 Kiro 等透传账号切到 Bedrock 等非透传账号前,
@@ -669,7 +669,9 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
h.handleFailoverExhausted(c, failoverErr, true)
return
}
if failoverErr.SafeToFailoverAfterWrite && c.Writer.Written() {
// openAIForwardMayFailover 已确认写出的字节不含语义输出,
// 但重试耗尽时仍须按已提交的 SSE 响应返回流内错误。
if c.Writer.Written() {
streamStarted = true
}
if failoverErr.ShouldReportAccountScheduleFailure() {
@@ -55,6 +55,11 @@ func (s *OpenAIGatewayService) handleOpenAIAccountUpstreamError(ctx context.Cont
if s != nil {
scheduleOllamaCloudUsageActivity(s.deferredService, account)
}
// Capacity shedding describes this request, not account health. Keep the
// account schedulable while the request-local retry budget handles recovery.
if account != nil && account.Platform == PlatformOpenAI && isOpenAIRequestScopedCapacityShed("", responseBody) {
return false
}
stateCtx, cancel := openAIAccountStateContext(ctx)
defer cancel()
@@ -77,6 +77,37 @@ func TestStreamFailedEventCapacityShedRetriesOnSameAccount(t *testing.T) {
require.False(t, openAIStreamFailedEventRetryableOnSameAccount(nonPool, other, "boom"))
}
func TestOpenAIHTTPCapacityShedIsRequestScopedForOAuthAccounts(t *testing.T) {
payload := []byte(`{"error":{"type":"server_error","message":"Our servers are currently overloaded. Please try again later."}}`)
failoverErr := newOpenAIUpstreamFailoverError(
http.StatusBadRequest,
http.Header{"X-Request-Id": []string{"rid-http-capacity"}},
payload,
"Our servers are currently overloaded. Please try again later.",
false,
)
require.True(t, failoverErr.RetryableOnSameAccount)
require.True(t, failoverErr.RequestScopedTransient)
repo := &capacityShedAccountRepoStub{}
(&GatewayService{accountRepo: repo}).TempUnscheduleRetryableError(context.Background(), 1, failoverErr)
require.Zero(t, repo.tempUnschedCalls)
rateLimitService := NewRateLimitService(repo, nil, &config.Config{}, nil, nil)
gateway := &OpenAIGatewayService{rateLimitService: rateLimitService}
account := &Account{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth}
require.False(t, gateway.handleOpenAIAccountUpstreamError(
context.Background(),
account,
http.StatusBadRequest,
nil,
payload,
"gpt-5",
))
require.Zero(t, repo.tempUnschedCalls)
}
// 上游降载的真实序列是「event: error → event: response.failed」。error 帧不算
// 客户端输出:若把它当首输出 flush,clientOutputStarted 被固化,随后的 failed
// 事件就进不了 pre-output failover 分支,只能把致命错误原样转发给客户端。
@@ -94,6 +125,11 @@ func TestOpenAIStreamErrorFrameDoesNotStartClientOutput(t *testing.T) {
{`{"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_item.added","item":{"type":"reasoning","summary":[]}}`, "response.output_item.added", false},
{`{"type":"response.output_item.added","item":{"type":"reasoning","encrypted_content":"ciphertext"}}`, "response.output_item.added", true},
{`{"type":"response.reasoning_summary_part.added","part":{"type":"summary_text","text":""}}`, "response.reasoning_summary_part.added", false},
{`{"type":"response.reasoning_summary_part.added","part":{"type":"summary_text","text":"thinking"}}`, "response.reasoning_summary_part.added", true},
{`{"type":"response.content_part.added","part":{"type":"output_text","text":""}}`, "response.content_part.added", false},
{`{"type":"response.output_text.delta","delta":"hi"}`, "response.output_text.delta", true},
{`[DONE]`, "", true},
}
@@ -102,6 +138,69 @@ func TestOpenAIStreamErrorFrameDoesNotStartClientOutput(t *testing.T) {
}
}
func TestOpenAIStreamMetadataPreambleAndMessageOnlyOverloadFailOver(t *testing.T) {
gin.SetMode(gin.TestMode)
largeMetadata := strings.Repeat("x", 16*1024)
stream := strings.Join([]string{
"event: response.created",
`data: {"type":"response.created","response":{"id":"resp_1","metadata":{"padding":"` + largeMetadata + `"}}}`,
"",
"event: response.output_item.added",
`data: {"type":"response.output_item.added","item":{"type":"reasoning","summary":[]}}`,
"",
"event: response.reasoning_summary_part.added",
`data: {"type":"response.reasoning_summary_part.added","part":{"type":"summary_text","text":""}}`,
"",
"event: error",
`data: {"type":"error","error":{"type":"service_unavailable_error","message":"Our servers are currently overloaded. Please try again later."}}`,
"",
}, "\n")
tests := []struct {
name string
run func(*OpenAIGatewayService, *gin.Context, *http.Response, *Account) error
}{
{
name: "native",
run: func(svc *OpenAIGatewayService, c *gin.Context, resp *http.Response, account *Account) error {
_, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, account, time.Now(), "model", "model")
return err
},
},
{
name: "passthrough",
run: func(svc *OpenAIGatewayService, c *gin.Context, resp *http.Response, account *Account) error {
_, err := svc.handleStreamingResponsePassthrough(c.Request.Context(), resp, c, account, time.Now(), "model", "model")
return err
},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
svc := &OpenAIGatewayService{cfg: &config.Config{Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize}}}
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(stream)),
Header: http.Header{"X-Request-Id": []string{"rid-message-only-overload"}},
}
account := &Account{ID: 1, Platform: PlatformOpenAI, Type: AccountTypeOAuth, Name: "acc"}
err := tt.run(svc, c, resp, account)
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())
})
}
}
// 回归用例(真实上游降载序列):created → in_progress → error 帧 → response.failed。
// 期望仍然走 pre-output failover(同账号重试 + 请求级瞬时标记),且不向客户端写出任何字节。
func TestOpenAIStreamCapacityShedErrorFramePrecedingFailedStillFailsOver(t *testing.T) {
@@ -149,6 +248,8 @@ func TestOpenAIStreamCapacityShedErrorFramePrecedingFailedStillFailsOver(t *test
// 并终止会话,对其余错误码执行内置退避重试。消息原样保留。
func TestOpenAIStreamCapacityShedAfterOutputRewritesCodeForClient(t *testing.T) {
gin.SetMode(gin.TestMode)
logSink, restore := captureStructuredLog(t)
defer restore()
cfg := &config.Config{
Gateway: config.GatewayConfig{MaxLineSize: defaultMaxLineSize},
}
@@ -188,6 +289,9 @@ func TestOpenAIStreamCapacityShedAfterOutputRewritesCodeForClient(t *testing.T)
require.Contains(t, body, `"code":"server_error"`)
require.NotContains(t, body, "server_is_overloaded")
require.Contains(t, body, "Our servers are currently overloaded")
require.True(t, logSink.ContainsMessage("gateway.failover_suppressed_after_semantic_output"))
require.True(t, logSink.ContainsFieldValue("path", "native_sse"))
require.True(t, logSink.ContainsFieldValue("upstream_request_id", "rid-shed-after-output"))
}
// helper 单测:只有降载码被改写,其余错误码(尤其 rate_limit_exceeded,客户端
@@ -211,6 +315,18 @@ func TestSanitizeOpenAICapacityShedErrorCodeForClient(t *testing.T) {
wantChanged: true,
wantContain: `"code":"server_error"`,
},
{
name: "failed事件只有过载文案时补充code",
payload: `{"type":"response.failed","response":{"error":{"message":"Our servers are currently overloaded. Please try again later."}}}`,
wantChanged: true,
wantContain: `"code":"server_error"`,
},
{
name: "error帧只有过载文案时补充code",
payload: `{"type":"error","error":{"message":"Server is overloaded. Please try again later."}}`,
wantChanged: true,
wantContain: `"code":"server_error"`,
},
{
name: "rate_limit不改写",
payload: `{"type":"response.failed","response":{"error":{"code":"rate_limit_exceeded","message":"try again in 3s"}}}`,
@@ -161,21 +161,28 @@ func OpenAICompactKeepaliveAdjustedWrittenSize(c *gin.Context) int {
if c == nil || c.Writer == nil {
return -1
}
value, ok := c.Get(openAICompactSSEKeepaliveKey)
if !ok {
return c.Writer.Size()
streamKeepaliveBytes := 0
if value, ok := c.Get(openAIStreamKeepaliveBytesKey); ok {
streamKeepaliveBytes, _ = value.(int)
}
k, ok := value.(*openAICompactSSEKeepalive)
if !ok || k == nil {
return c.Writer.Size()
size := c.Writer.Size()
compactKeepaliveBytes := 0
if value, ok := c.Get(openAICompactSSEKeepaliveKey); ok {
if k, valid := value.(*openAICompactSSEKeepalive); valid && k != nil {
k.mu.Lock()
size = k.writer.Size()
compactKeepaliveBytes = k.bytes
k.mu.Unlock()
}
}
k.mu.Lock()
defer k.mu.Unlock()
size := k.writer.Size()
if size < 0 {
return size
}
if real := size - k.bytes; real > 0 {
keepaliveBytes := compactKeepaliveBytes + streamKeepaliveBytes
if keepaliveBytes <= 0 {
return size
}
if real := size - keepaliveBytes; real > 0 {
return real
}
return -1
@@ -71,6 +71,20 @@ func TestOpenAICompactSSEKeepalive_StopBeforeFirstBeatKeepsWriterUntouched(t *te
require.False(t, StopOpenAICompactSSEKeepaliveCommitted(c))
}
func TestOpenAIAdjustedWrittenSizeExcludesResponsesStreamKeepalive(t *testing.T) {
c, rec := newCompactBridgeTestContext(t, false)
n, err := c.Writer.Write([]byte(":\n\n"))
require.NoError(t, err)
recordOpenAIStreamKeepaliveBytes(c, n)
require.Equal(t, -1, OpenAICompactKeepaliveAdjustedWrittenSize(c))
_, err = c.Writer.Write([]byte("data: semantic\n\n"))
require.NoError(t, err)
require.Equal(t, len("data: semantic\n\n"), OpenAICompactKeepaliveAdjustedWrittenSize(c))
require.Equal(t, ":\n\ndata: semantic\n\n", rec.Body.String())
}
// 心跳已提交后,2xx 桥接续写事件而不重复提交响应头。
func TestWriteOpenAICompactSSEBridge_AfterKeepaliveCommitAppendsEvents(t *testing.T) {
c, rec := newCompactBridgeTestContext(t, true)
@@ -542,7 +542,7 @@ func TestOpenAINativeFirstOutputScannerAllowsLargeEventAfterSemanticBoundary(t *
require.Equal(t, "request-large-image", rec.Result().Header.Get("X-Request-Id"))
}
func TestOpenAINativeFirstOutputTimeoutDisabledPreservesKeepaliveFlush(t *testing.T) {
func TestOpenAINativeFirstOutputTimeoutDisabledKeepsPreamblePrivateAcrossKeepalive(t *testing.T) {
svc := &OpenAIGatewayService{cfg: &config.Config{Gateway: config.GatewayConfig{
StreamKeepaliveInterval: 1,
MaxLineSize: defaultMaxLineSize,
@@ -552,7 +552,7 @@ func TestOpenAINativeFirstOutputTimeoutDisabledPreservesKeepaliveFlush(t *testin
defer func() { _ = pw.Close() }()
_, _ = pw.Write([]byte("data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_stalled\"}}\n\n"))
_, _ = pw.Write([]byte("data: {\"type\":\"response.in_progress\",\"response\":{\"id\":\"resp_stalled\"}}\n\n"))
time.Sleep(1100 * time.Millisecond)
time.Sleep(2100 * time.Millisecond)
}()
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
@@ -561,10 +561,11 @@ func TestOpenAINativeFirstOutputTimeoutDisabledPreservesKeepaliveFlush(t *testin
_, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI}, time.Now(), "model", "model")
require.Error(t, err)
var failoverErr *UpstreamFailoverError
require.ErrorAs(t, err, &failoverErr)
require.Contains(t, rec.Body.String(), ":\n\n")
require.Contains(t, rec.Body.String(), "response.created")
require.Contains(t, rec.Body.String(), "response.in_progress")
require.NotContains(t, rec.Body.String(), "response.created")
require.NotContains(t, rec.Body.String(), "response.in_progress")
}
func TestOpenAINativeFirstOutputFailoverKeepsAttemptHeadersPrivateAfterKeepaliveCommit(t *testing.T) {
@@ -908,6 +908,19 @@ type openaiNonStreamingResultPassthrough struct {
imageOutputSizes []string
}
const openAIStreamKeepaliveBytesKey = "openai_stream_keepalive_bytes"
func recordOpenAIStreamKeepaliveBytes(c *gin.Context, written int) {
if c == nil || written <= 0 {
return
}
current := 0
if value, ok := c.Get(openAIStreamKeepaliveBytesKey); ok {
current, _ = value.(int)
}
c.Set(openAIStreamKeepaliveBytesKey, current+written)
}
func openAIStreamClientOutputStarted(c *gin.Context, localStarted bool) bool {
if localStarted {
return true
@@ -930,6 +943,85 @@ func openAIStreamEventIsPreamble(eventType string) bool {
}
}
func openAIStreamAddedEventStartsClientOutput(payload []byte, eventType string) bool {
if len(payload) == 0 || !gjson.ValidBytes(payload) {
return true
}
switch strings.TrimSpace(eventType) {
case "response.output_item.added":
item := gjson.GetBytes(payload, "item")
if !item.Exists() || !item.IsObject() {
return true
}
switch strings.TrimSpace(item.Get("type").String()) {
case "reasoning":
if item.Get("encrypted_content").String() != "" {
return true
}
summary := item.Get("summary")
if !summary.IsArray() {
return false
}
for _, part := range summary.Array() {
if strings.TrimSpace(part.Get("type").String()) != "summary_text" || part.Get("text").String() != "" {
return true
}
}
return false
case "message":
content := item.Get("content")
if !content.IsArray() {
return false
}
for _, part := range content.Array() {
switch strings.TrimSpace(part.Get("type").String()) {
case "output_text":
if part.Get("text").String() != "" {
return true
}
case "refusal":
if part.Get("refusal").String() != "" {
return true
}
default:
return true
}
}
return false
case "function_call":
return item.Get("arguments").String() != ""
case "custom_tool_call":
return item.Get("input").String() != ""
case "compaction":
return item.Get("encrypted_content").String() != ""
default:
return true
}
case "response.content_part.added":
part := gjson.GetBytes(payload, "part")
if !part.Exists() || !part.IsObject() {
return true
}
switch strings.TrimSpace(part.Get("type").String()) {
case "output_text":
return part.Get("text").String() != ""
case "refusal":
return part.Get("refusal").String() != ""
default:
return true
}
case "response.reasoning_summary_part.added":
part := gjson.GetBytes(payload, "part")
if !part.Exists() || !part.IsObject() || strings.TrimSpace(part.Get("type").String()) != "summary_text" {
return true
}
return part.Get("text").String() != ""
default:
return true
}
}
func openAIStreamDataStartsClientOutput(data, eventType string) bool {
trimmed := strings.TrimSpace(data)
if trimmed == "" {
@@ -946,6 +1038,8 @@ func openAIStreamDataStartsClientOutput(data, eventType string) bool {
// (content_policy / invalid_request 等)维持原样转发,保留上游错误细节。
payload := []byte(trimmed)
return !openAIStreamFailedEventShouldFailover(payload, extractOpenAISSEErrorMessage(payload))
case "response.output_item.added", "response.content_part.added", "response.reasoning_summary_part.added":
return openAIStreamAddedEventStartsClientOutput([]byte(trimmed), eventType)
}
return !openAIStreamEventIsPreamble(eventType)
}
@@ -1024,9 +1118,34 @@ func isOpenAIUpstreamCapacityShedEvent(payload []byte) bool {
switch openAIStreamFailedEventErrorCode(payload) {
case "server_is_overloaded", "slow_down":
return true
default:
return false
}
for _, path := range []string{"response.error.message", "error.message", "message"} {
if isOpenAICapacityShedMessage(gjson.GetBytes(payload, path).String()) {
return true
}
}
return false
}
func logOpenAICapacityFailoverSuppressed(
ctx context.Context,
account *Account,
path string,
upstreamRequestID string,
eventType string,
) {
fields := []zap.Field{
zap.String("path", path),
zap.String("event_type", strings.TrimSpace(eventType)),
zap.String("upstream_request_id", strings.TrimSpace(upstreamRequestID)),
}
if account != nil {
fields = append(fields,
zap.Int64("account_id", account.ID),
zap.String("platform", account.Platform),
)
}
logger.FromContext(ctx).Warn("gateway.failover_suppressed_after_semantic_output", fields...)
}
// openAICapacityShedRetryableClientCode 是把上游容量降载错误转发给客户端时改写
@@ -1049,9 +1168,12 @@ func sanitizeOpenAICapacityShedErrorCodeForClient(payload []byte) ([]byte, bool)
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:
parent := strings.TrimSuffix(path, ".code")
if !gjson.GetBytes(updated, parent).Exists() {
continue
}
code := strings.ToLower(strings.TrimSpace(gjson.GetBytes(updated, path).String()))
if code != "" && code != "server_is_overloaded" && code != "slow_down" {
continue
}
next, err := sjson.SetBytes(updated, path, openAICapacityShedRetryableClientCode)
@@ -1084,7 +1206,7 @@ func openAIStreamFailedEventSemanticStatus(payload []byte, message string) int {
return http.StatusUnauthorized
case strings.Contains(combined, "permission") || strings.Contains(combined, "forbidden") || strings.Contains(combined, "access denied"):
return http.StatusForbidden
case code == "server_is_overloaded" || code == "slow_down":
case isOpenAIUpstreamCapacityShedEvent(payload):
return http.StatusServiceUnavailable
default:
return http.StatusBadGateway
@@ -1218,6 +1340,16 @@ func openAIStreamFailedEventShouldFailover(payload []byte, message string) bool
return true
}
func openAIStreamErrorEventShouldFailover(payload []byte, message string) bool {
if strings.TrimSpace(gjson.GetBytes(payload, "type").String()) != "error" {
return false
}
if isOpenAIContextWindowError(message, payload) {
return false
}
return isOpenAITransientProcessingError(http.StatusBadRequest, message, payload)
}
func openAIStreamFailedEventRetryableOnSameAccount(account *Account, payload []byte, message string) bool {
if account == nil {
return false
@@ -1360,6 +1492,7 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough(
sawTerminalEvent := false
sawFailedEvent := false
semanticOutputSeen := false
capacityFailoverSuppressedLogged := false
failedMessage := ""
clientOutputStarted := false
upstreamRequestID := strings.TrimSpace(resp.Header.Get("x-request-id"))
@@ -1446,6 +1579,32 @@ func (s *OpenAIGatewayService) handleStreamingResponsePassthrough(
}
}
eventType := strings.TrimSpace(gjson.Get(trimmedData, "type").String())
if !capacityFailoverSuppressedLogged && account != nil && account.Platform == PlatformOpenAI &&
(eventType == "error" || eventType == "response.failed") &&
openAIStreamClientOutputStarted(c, clientOutputStarted) &&
isOpenAIUpstreamCapacityShedEvent(dataBytes) {
logOpenAICapacityFailoverSuppressed(ctx, account, "passthrough_sse", upstreamRequestID, eventType)
capacityFailoverSuppressedLogged = true
}
if eventType == "error" && !openAIStreamClientOutputStarted(c, clientOutputStarted) {
errorMessage := extractOpenAISSEErrorMessage(dataBytes)
if status, errType, errMsg, matched := applyOpenAIStreamFailedErrorPassthroughRule(c, account.Platform, dataBytes, errorMessage); matched {
s.recordOpenAIStreamUpstreamError(c, account, true, upstreamRequestID, "http_error", dataBytes, errorMessage)
MarkResponseCommitted(c)
c.Writer.Header().Set("Content-Type", "application/json; charset=utf-8")
c.JSON(status, gin.H{
"error": gin.H{
"type": errType,
"message": errMsg,
},
})
return resultWithUsage(), fmt.Errorf("upstream error event: passthrough rule matched message=%s", errMsg)
}
if openAIStreamErrorEventShouldFailover(dataBytes, errorMessage) {
return resultWithUsage(),
s.newOpenAIStreamFailoverError(c, account, true, upstreamRequestID, dataBytes, errorMessage, resp.Header)
}
}
if eventType == "response.failed" {
failedMessage = extractOpenAISSEErrorMessage(dataBytes)
// response.failed 自带上游已消耗的 usage(input token 通常已扣);必须先解析
@@ -54,8 +54,9 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
firstOutputTimeout = s.openAIFirstOutputTimeout(reasoningEffort)
}
guardFirstOutput := firstOutputTimeout > 0
stageFirstOutput := account != nil && account.Platform == PlatformOpenAI
var attemptResponseHeaders http.Header
if guardFirstOutput {
if stageFirstOutput {
if s.responseHeaderFilter != nil {
attemptResponseHeaders = responseheaders.FilterHeaders(resp.Header, s.responseHeaderFilter)
} else if requestID := strings.TrimSpace(resp.Header.Get("x-request-id")); requestID != "" {
@@ -66,8 +67,8 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
}
// x-codex-turn-state 不在通用响应头白名单内,按 Codex 协议显式回传:
// 客户端会在同回合的后续请求中回带(openai_codex_turn_state.go)。
// 首输出守卫模式下只暂存,溯源在 applyAttemptResponseHeaders 真正提交时记录。
if guardFirstOutput {
// OpenAI 首个语义输出前只暂存,溯源在 applyAttemptResponseHeaders 真正提交时记录。
if stageFirstOutput {
stageOpenAICodexTurnState(&attemptResponseHeaders, resp.Header)
} else {
s.relayOpenAICodexTurnState(c, account, resp.Header)
@@ -80,12 +81,12 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
c.Header("X-Accel-Buffering", "no")
// Pass through other headers
if !guardFirstOutput && resp.Header.Get("x-request-id") != "" {
if !stageFirstOutput && resp.Header.Get("x-request-id") != "" {
v := resp.Header.Get("x-request-id")
c.Header("x-request-id", v)
}
applyAttemptResponseHeaders := func() {
if !guardFirstOutput || len(attemptResponseHeaders) == 0 || c.Writer.Written() {
if !stageFirstOutput || len(attemptResponseHeaders) == 0 || c.Writer.Written() {
return
}
for key, values := range attemptResponseHeaders {
@@ -117,7 +118,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
firstOutputProgressObserved := false
bufferedWriter := bufio.NewWriterSize(w, 4*1024)
var firstOutputStage *openAIFirstOutputStage
if guardFirstOutput {
if stageFirstOutput {
firstOutputStage = newDefaultOpenAIFirstOutputStage()
defer func() {
if err := firstOutputStage.Close(); err != nil {
@@ -126,19 +127,19 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
}()
}
writePendingString := func(value string) (int, error) {
if firstOutputStage != nil && !firstOutputProgressObserved && !firstOutputStage.closed {
if firstOutputStage != nil && !firstOutputStage.closed {
return firstOutputStage.WriteString(value)
}
return bufferedWriter.WriteString(value)
}
pendingBytes := func() int64 {
if firstOutputStage != nil && !firstOutputProgressObserved && !firstOutputStage.closed {
if firstOutputStage != nil && !firstOutputStage.closed {
return firstOutputStage.Buffered()
}
return int64(bufferedWriter.Buffered())
}
flushBuffered := func() error {
if firstOutputStage != nil && !firstOutputProgressObserved && !firstOutputStage.closed {
if firstOutputStage != nil && !firstOutputStage.closed {
if err := firstOutputStage.CommitTo(w); err != nil {
return err
}
@@ -155,11 +156,11 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
imageCounter := newOpenAIImageOutputCounter()
responseID := ""
var firstOutputScanGuard atomic.Bool
firstOutputScanGuard.Store(guardFirstOutput)
firstOutputScanGuard.Store(stageFirstOutput)
scanner := bufio.NewScanner(resp.Body)
scanBuf := getSSEScannerBuf64K()
scanner.Buffer(scanBuf[:0], maxLineSize)
if guardFirstOutput {
if stageFirstOutput {
scanner.Split(openAIFirstOutputDynamicScanLines(&firstOutputScanGuard))
}
documentScanner := newOpenAISSEJSONDocumentScanner(scanner)
@@ -241,6 +242,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
sawTerminalEvent := false
sawFailedEvent := false
responsesSemanticOutputSeen := false
capacityFailoverSuppressedLogged := false
failedMessage := ""
clientOutputStarted := false
upstreamRequestID := strings.TrimSpace(resp.Header.Get("x-request-id"))
@@ -250,7 +252,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
eventStartsVisibleOutput := false
eventShouldFlush := false
handlePendingWriteError := func(err error) {
if firstOutputStage != nil && !firstOutputProgressObserved && !firstOutputStage.closed {
if firstOutputStage != nil && !firstOutputStage.closed {
message := "OpenAI first-output staging failed"
if errors.Is(err, errOpenAIFirstOutputStageLimit) {
message = "OpenAI first-output staging limit exceeded"
@@ -350,7 +352,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
lastDownstreamWriteAt = time.Now()
}
finalizeStream := func() (*openaiStreamingResult, error) {
if guardFirstOutput && eventInProgress {
if stageFirstOutput && eventInProgress {
// EOF dispatches the final SSE event even without a trailing blank line.
completeGuardedEvent(true)
}
@@ -392,7 +394,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
failoverErr.SafeToFailoverAfterWrite = true
return resultWithUsage(), failoverErr, true
}
if errors.Is(scanErr, bufio.ErrTooLong) && guardFirstOutput && !firstOutputProgressObserved {
if errors.Is(scanErr, bufio.ErrTooLong) && stageFirstOutput && !firstOutputProgressObserved {
logger.LegacyPrintf("service.openai_gateway", "SSE line too long before first output: account=%d max_size=%d error=%v", account.ID, maxLineSize, scanErr)
failoverErr := s.newOpenAIStreamFailoverError(
c, account, false, upstreamRequestID, nil,
@@ -455,6 +457,33 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
responseID = extractOpenAIResponseIDFromJSONBytes(dataBytes)
}
forceFlushFailedEvent := false
if !capacityFailoverSuppressedLogged && account != nil && account.Platform == PlatformOpenAI &&
(eventType == "error" || eventType == "response.failed") &&
openAIStreamClientOutputStarted(c, clientOutputStarted) &&
isOpenAIUpstreamCapacityShedEvent(dataBytes) {
logOpenAICapacityFailoverSuppressed(ctx, account, "native_sse", upstreamRequestID, eventType)
capacityFailoverSuppressedLogged = true
}
if eventType == "error" && !openAIStreamClientOutputStarted(c, clientOutputStarted) {
errorMessage := extractOpenAISSEErrorMessage(dataBytes)
if status, errType, errMsg, matched := applyOpenAIStreamFailedErrorPassthroughRule(c, account.Platform, dataBytes, errorMessage); matched {
s.recordOpenAIStreamUpstreamError(c, account, false, upstreamRequestID, "http_error", dataBytes, errorMessage)
MarkResponseCommitted(c)
c.Writer.Header().Set("Content-Type", "application/json; charset=utf-8")
c.JSON(status, gin.H{
"error": gin.H{
"type": errType,
"message": errMsg,
},
})
streamEarlyErr = fmt.Errorf("upstream error event: passthrough rule matched message=%s", errMsg)
return
}
if openAIStreamErrorEventShouldFailover(dataBytes, errorMessage) {
streamEarlyErr = s.newOpenAIStreamFailoverError(c, account, false, upstreamRequestID, dataBytes, errorMessage, resp.Header)
return
}
}
if eventType == "response.failed" {
failedMessage = extractOpenAISSEErrorMessage(dataBytes)
// response.failed 自带上游已消耗的 usage(input token 通常已扣);必须先解析
@@ -559,9 +588,12 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
}
startsClientOutput := forceFlushFailedEvent || openAIStreamDataStartsClientOutput(data, eventType)
startsVisibleOutput := openAIStreamDataStartsVisibleOutput(data, eventType)
if guardFirstOutput {
if stageFirstOutput {
eventStartsClientOutput = eventStartsClientOutput || startsClientOutput
eventStartsVisibleOutput = eventStartsVisibleOutput || startsVisibleOutput
if startsClientOutput {
firstOutputScanGuard.Store(false)
}
}
if startsClientOutput && !openAIStreamEventTypeIsTerminal(eventType) {
responsesSemanticOutputSeen = true
@@ -607,7 +639,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
}
// A blank line dispatches a guarded event from the attempt-local stage.
if guardFirstOutput && line == "" {
if stageFirstOutput && line == "" {
if !clientDisconnected {
if _, err := writePendingString("\n"); err != nil {
handlePendingWriteError(err)
@@ -672,7 +704,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
events := make(chan scanEvent, openAIFirstOutputEventQueueSize(guardFirstOutput))
done := make(chan struct{})
sendEvent := func(ev scanEvent) bool {
if guardFirstOutput {
if firstOutputScanGuard.Load() {
ev.processed = make(chan struct{})
}
select {
@@ -716,7 +748,7 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
select {
case ev, ok := <-events:
if !ok {
if guardFirstOutput && eventInProgress {
if stageFirstOutput && eventInProgress {
// EOF dispatches the final SSE event even without a trailing blank
// line. Do not synthesize extra bytes on the downstream wire.
completeGuardedEvent(true)
@@ -783,10 +815,12 @@ func (s *OpenAIGatewayService) handleStreamingResponseWithReasoning(ctx context.
if time.Since(lastDownstreamWriteAt) < keepaliveInterval {
continue
}
if guardFirstOutput {
if stageFirstOutput {
// Bypass attempt-local buffered frames. The stable SSE headers may be
// committed here, but account headers remain private until semantic output.
if _, err := w.Write([]byte(":\n\n")); err != nil {
n, err := w.Write([]byte(":\n\n"))
recordOpenAIStreamKeepaliveBytes(c, n)
if err != nil {
clientDisconnected = true
logger.LegacyPrintf("service.openai_gateway", "Client disconnected during streaming, continuing to drain upstream for billing")
continue
@@ -287,6 +287,24 @@ func TestIsOpenAITransientProcessingError(t *testing.T) {
[]byte(`{"error":{"code":"slow_down","message":"Please retry later."}}`),
))
require.True(t, isOpenAITransientProcessingError(
http.StatusBadRequest,
"",
[]byte(`{"error":{"message":"Our servers are currently overloaded. Please try again later."}}`),
))
require.True(t, isOpenAITransientProcessingError(
http.StatusServiceUnavailable,
"Server is overloaded. Please try again later.",
nil,
))
require.True(t, isOpenAITransientProcessingError(
http.StatusBadGateway,
"",
[]byte(`{"error":{"message":"Our servers are currently overloaded. Please try again later."}}`),
))
require.True(t, isOpenAITransientProcessingError(
http.StatusBadRequest,
"",
@@ -115,7 +115,7 @@ func (r stubOpenAIAccountRepo) ListSchedulableUngroupedByPlatform(ctx context.Co
return r.ListSchedulableByPlatform(ctx, platform)
}
func TestOpenAIGatewayService_ForwardAsAnthropic_TempUnschedulableReturnsFailoverWithoutCommit(t *testing.T) {
func TestOpenAIGatewayService_ForwardAsAnthropic_CapacityShedReturnsRequestScopedFailoverWithoutCommit(t *testing.T) {
gin.SetMode(gin.TestMode)
body := []byte(`{"model":"gpt-5.4","max_tokens":32,"messages":[{"role":"user","content":"hello"}],"stream":false}`)
@@ -173,8 +173,10 @@ func TestOpenAIGatewayService_ForwardAsAnthropic_TempUnschedulableReturnsFailove
require.ErrorAs(t, err, &failoverErr)
require.Equal(t, http.StatusBadRequest, failoverErr.StatusCode)
require.True(t, failoverErr.ShouldRetryNextAccount())
require.Equal(t, account.ID, repo.modelRateLimitAccountID, "temporary unschedulability should exclude this account from reselection")
require.Equal(t, "gpt-5.4", repo.modelRateLimitKey)
require.True(t, failoverErr.RetryableOnSameAccount)
require.True(t, failoverErr.RequestScopedTransient)
require.Zero(t, repo.modelRateLimitAccountID, "request-scoped capacity shedding must not change account health")
require.Empty(t, repo.modelRateLimitKey)
require.False(t, IsResponseCommitted(c))
require.Equal(t, http.StatusOK, rec.Code)
require.Empty(t, rec.Body.String())
@@ -204,17 +206,17 @@ func TestFailoverOpenAIUpstreamHTTPError_NilContextSkipsTempUnschedulablePolicy(
"temp_unschedulable_enabled": true,
"temp_unschedulable_rules": []any{map[string]any{
"error_code": float64(http.StatusBadRequest),
"keywords": []any{"servers are currently overloaded"},
"keywords": []any{"custom temporary outage"},
"duration_minutes": float64(1),
}},
},
}
body := []byte(`{"error":{"message":"Our servers are currently overloaded."}}`)
body := []byte(`{"error":{"message":"Custom temporary outage."}}`)
resp := &http.Response{StatusCode: http.StatusBadRequest, Header: http.Header{}}
got := svc.failoverOpenAIUpstreamHTTPError(
context.Background(), nil, account, resp, body,
"Our servers are currently overloaded.", "gpt-5.4",
"Custom temporary outage.", "gpt-5.4",
)
require.Nil(t, got)
@@ -2288,7 +2290,7 @@ func TestOpenAIStreamingMissingTerminalEventReturnsIncompleteError(t *testing.T)
go func() {
defer func() { _ = pw.Close() }()
_, _ = pw.Write([]byte("data: {\"type\":\"response.output_item.added\",\"item\":{\"type\":\"message\"},\"output_index\":0}\n\n"))
_, _ = pw.Write([]byte("data: {\"type\":\"response.output_text.delta\",\"delta\":\"partial\",\"output_index\":0}\n\n"))
}()
_, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, &Account{ID: 1}, time.Now(), "model", "model")
@@ -2320,7 +2322,7 @@ func TestOpenAIStreamingPassthroughMissingTerminalEventReturnsIncompleteError(t
go func() {
defer func() { _ = pw.Close() }()
_, _ = pw.Write([]byte("data: {\"type\":\"response.output_item.added\",\"item\":{\"type\":\"message\"},\"output_index\":0}\n\n"))
_, _ = pw.Write([]byte("data: {\"type\":\"response.output_text.delta\",\"delta\":\"partial\",\"output_index\":0}\n\n"))
}()
_, err := svc.handleStreamingResponsePassthrough(c.Request.Context(), resp, c, &Account{ID: 1}, time.Now(), "", "")
@@ -117,7 +117,7 @@ func isOpenAIInstructionsRequiredError(upstreamStatusCode int, upstreamMsg strin
}
func isOpenAITransientProcessingError(upstreamStatusCode int, upstreamMsg string, upstreamBody []byte) bool {
if upstreamStatusCode != http.StatusBadRequest && upstreamStatusCode != http.StatusServiceUnavailable {
if upstreamStatusCode < http.StatusBadRequest {
return false
}
@@ -132,6 +132,15 @@ func isOpenAITransientProcessingError(upstreamStatusCode int, upstreamMsg string
if len(upstreamBody) > 0 && hasOpenAIServerOverloadedCode(upstreamBody) {
return true
}
if isOpenAICapacityShedMessage(upstreamMsg) ||
isOpenAICapacityShedMessage(gjson.GetBytes(upstreamBody, "error.message").String()) ||
isOpenAICapacityShedMessage(gjson.GetBytes(upstreamBody, "response.error.message").String()) ||
isOpenAICapacityShedMessage(string(upstreamBody)) {
return true
}
if upstreamStatusCode != http.StatusBadRequest && upstreamStatusCode != http.StatusServiceUnavailable {
return false
}
if upstreamStatusCode != http.StatusBadRequest {
return false
}
@@ -164,6 +173,19 @@ func isOpenAITransientProcessingError(upstreamStatusCode int, upstreamMsg string
return match(string(upstreamBody))
}
func isOpenAICapacityShedMessage(text string) bool {
lower := strings.ToLower(strings.TrimSpace(text))
return strings.Contains(lower, "server is overloaded") ||
strings.Contains(lower, "servers are overloaded") ||
strings.Contains(lower, "servers are currently overloaded")
}
func isOpenAIRequestScopedCapacityShed(upstreamMsg string, upstreamBody []byte) bool {
return isOpenAIUpstreamCapacityShedEvent(upstreamBody) ||
isOpenAICapacityShedMessage(upstreamMsg) ||
isOpenAICapacityShedMessage(string(upstreamBody))
}
func isOpenAIContextWindowError(upstreamMsg string, upstreamBody []byte) bool {
match := func(text string) bool {
lower := strings.ToLower(strings.TrimSpace(text))
@@ -248,14 +270,17 @@ func newOpenAIUpstreamFailoverError(
upstreamMsg string,
retryableOnSameAccount bool,
) *UpstreamFailoverError {
requestScopedCapacity := isOpenAIRequestScopedCapacityShed(upstreamMsg, responseBody)
failoverErr := &UpstreamFailoverError{
StatusCode: statusCode,
ResponseBody: responseBody,
ResponseHeaders: responseHeaders.Clone(),
RetryableOnSameAccount: retryableOnSameAccount,
RetryableOnSameAccount: retryableOnSameAccount || requestScopedCapacity,
RequestScopedTransient: requestScopedCapacity,
}
if isOpenAIRequestBodyTooLargeError(statusCode, upstreamMsg, responseBody) {
failoverErr.RetryableOnSameAccount = false
failoverErr.RequestScopedTransient = false
failoverErr.Scope = GatewayFailureScopeAccount
failoverErr.Reason = openAIRequestBodyTooLargeReason
failoverErr.NextAccountAction = NextAccountRetry
@@ -71,11 +71,38 @@ func TestOpenAIResponsesTTFTStartsAtCompletedImage(t *testing.T) {
}
}
func TestOpenAINativeProgressDisarmsTimeoutWithoutStartingTTFT(t *testing.T) {
result := runSyntheticVisibleTTFTStream(t, false, 1200*time.Millisecond, 1,
`{"type":"response.output_text.delta","delta":"test output"}`)
require.NotNil(t, result.firstTokenMs)
require.GreaterOrEqual(t, *result.firstTokenMs, 1100)
func TestOpenAINativeMetadataDoesNotDisarmFirstOutputTimeout(t *testing.T) {
gin.SetMode(gin.TestMode)
svc := &OpenAIGatewayService{cfg: &config.Config{Gateway: config.GatewayConfig{
MaxLineSize: defaultMaxLineSize,
OpenAIFirstOutputTimeoutSeconds: 1,
}}}
reader, writer := io.Pipe()
writerDone := make(chan struct{})
go func() {
defer close(writerDone)
defer func() { _ = writer.Close() }()
_, _ = io.WriteString(writer, "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_test\"}}\n\n")
_, _ = io.WriteString(writer, "data: {\"type\":\"response.output_item.added\",\"item\":{\"id\":\"item_test\",\"type\":\"reasoning\",\"summary\":[]}}\n\n")
time.Sleep(1200 * time.Millisecond)
}()
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
resp := &http.Response{StatusCode: http.StatusOK, Header: http.Header{}, Body: reader}
account := &Account{ID: 1, Name: "account_test", Platform: PlatformOpenAI}
_, err := svc.handleStreamingResponse(context.Background(), resp, c, account, time.Now(), "test-model", "test-model")
var failoverErr *UpstreamFailoverError
require.ErrorAs(t, err, &failoverErr)
require.True(t, failoverErr.SafeToFailoverAfterWrite)
require.Empty(t, recorder.Body.String())
select {
case <-writerDone:
case <-time.After(time.Second):
t.Fatal("synthetic upstream writer did not exit")
}
}
func runSyntheticVisibleTTFTStream(t *testing.T, passthrough bool, visibleDelay time.Duration, timeoutSeconds int, visibleEvent string) *openaiStreamingResult {
@@ -290,6 +290,9 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
upstreamTerminalEvent := ""
sawDone := false
wroteDownstream := false
pendingClientMessages := make([][]byte, 0, 4)
pendingClientMessageBytes := int64(0)
capacityFailoverSuppressedLogged := false
clientDisconnected := false
mappedModel := ""
needModelReplace := false
@@ -403,15 +406,20 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
replayCollector.AddEvent(eventType, upstreamMessage)
var upstreamEventErr error
if eventType == "error" {
errCodeRaw, errTypeRaw, errMsgRaw := parseOpenAIWSErrorEventFields(upstreamMessage)
errMessage := strings.TrimSpace(errMsgRaw)
if eventType == "error" || eventType == "response.failed" {
errMessage := extractOpenAISSEErrorMessage(upstreamMessage)
if errMessage == "" {
errMessage = "upstream error event"
}
statusCode := openAIWSErrorHTTPStatusFromRaw(errCodeRaw, errTypeRaw)
shouldFailover := s.shouldFailoverOpenAIUpstreamResponse(statusCode, errMessage, upstreamMessage)
if account.Platform == PlatformGrok {
statusCode := openAIStreamFailureStatus(upstreamMessage, errMessage)
shouldFailover := openAIStreamFailedEventShouldFailover(upstreamMessage, errMessage)
if eventType == "error" {
errCodeRaw, errTypeRaw, _ := parseOpenAIWSErrorEventFields(upstreamMessage)
statusCode = openAIWSErrorHTTPStatusFromRaw(errCodeRaw, errTypeRaw)
shouldFailover = s.shouldFailoverOpenAIUpstreamResponse(statusCode, errMessage, upstreamMessage)
}
requestScopedCapacity := isOpenAIUpstreamCapacityShedEvent(upstreamMessage)
if account.Platform == PlatformGrok && eventType == "error" {
// SSE error events do not carry an HTTP status. The local status
// mapper therefore defaults unknown xAI codes (for example
// new_sensitive) to 502; classify the body as a request-scoped
@@ -422,7 +430,7 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
shouldFailover = s.shouldFailoverGrokUpstreamError(statusCode, upstreamMessage)
s.handleGrokAccountUpstreamError(ctx, account, statusCode, resp.Header, upstreamMessage)
}
} else if shouldFailover {
} else if eventType == "error" && shouldFailover && !requestScopedCapacity {
accountStatus := statusCode
if transientStatus := openAIWSPayloadTransientStatus(upstreamMessage); transientStatus != 0 {
accountStatus = transientStatus
@@ -431,9 +439,18 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
s.handleOpenAIAccountUpstreamError(ctx, account, accountStatus, resp.Header, upstreamMessage, canonicalModel)
}
if turn == 1 && !wroteDownstream && shouldFailover {
return nil, newOpenAIUpstreamFailoverError(statusCode, resp.Header, upstreamMessage, errMessage, false)
if account.Platform == PlatformGrok {
return nil, newOpenAIUpstreamFailoverError(statusCode, resp.Header, upstreamMessage, errMessage, false)
}
return nil, s.newOpenAIStreamFailoverError(c, account, true, resp.Header.Get("x-request-id"), upstreamMessage, errMessage, resp.Header)
}
if wroteDownstream && requestScopedCapacity && !capacityFailoverSuppressedLogged {
logOpenAICapacityFailoverSuppressed(ctx, account, "ws_http_bridge", resp.Header.Get("x-request-id"), eventType)
capacityFailoverSuppressedLogged = true
}
if eventType == "error" {
upstreamEventErr = errors.New(errMessage)
}
upstreamEventErr = errors.New(errMessage)
}
// 客户端写出副本改写容量降载码:Codex 对 error/response.failed 中的
@@ -447,26 +464,50 @@ func (s *OpenAIGatewayService) proxyOpenAIWSHTTPBridgeTurn(
}
}
if !clientDisconnected {
if err := writeClientMessage(clientMessage); err != nil {
if isOpenAIWSClientDisconnectError(err) {
clientDisconnected = true
closeStatus, closeReason := summarizeOpenAIWSReadCloseError(err)
logOpenAIWSModeInfo(
"ingress_ws_http_bridge_client_disconnected_drain account_id=%d turn=%d close_status=%s close_reason=%s",
account.ID,
turn,
closeStatus,
truncateOpenAIWSLogValue(closeReason, openAIWSHeaderValueMaxLen),
)
} else {
return nil, wrapOpenAIWSIngressTurnError(
"write_client",
fmt.Errorf("write client websocket event: %w", err),
wroteDownstream,
stageBeforeSemanticOutput := turn == 1 && account.Platform == PlatformOpenAI && !wroteDownstream
commitStagedMessages := !stageBeforeSemanticOutput ||
openAIStreamDataStartsClientOutput(string(clientMessage), eventType) ||
isOpenAIWSTerminalEvent(eventType)
if stageBeforeSemanticOutput && !commitStagedMessages {
if pendingClientMessageBytes+int64(len(clientMessage)) > openAIFirstOutputStageMaxBytes {
return nil, s.newOpenAIStreamFailoverError(
c,
account,
true,
resp.Header.Get("x-request-id"),
nil,
"OpenAI WS HTTP bridge first-output staging limit exceeded",
resp.Header,
)
}
pendingClientMessages = append(pendingClientMessages, append([]byte(nil), clientMessage...))
pendingClientMessageBytes += int64(len(clientMessage))
} else {
wroteDownstream = true
messages := append(pendingClientMessages, clientMessage)
pendingClientMessages = nil
pendingClientMessageBytes = 0
for _, message := range messages {
if err := writeClientMessage(message); err != nil {
if isOpenAIWSClientDisconnectError(err) {
clientDisconnected = true
closeStatus, closeReason := summarizeOpenAIWSReadCloseError(err)
logOpenAIWSModeInfo(
"ingress_ws_http_bridge_client_disconnected_drain account_id=%d turn=%d close_status=%s close_reason=%s",
account.ID,
turn,
closeStatus,
truncateOpenAIWSLogValue(closeReason, openAIWSHeaderValueMaxLen),
)
break
}
return nil, wrapOpenAIWSIngressTurnError(
"write_client",
fmt.Errorf("write client websocket event: %w", err),
wroteDownstream,
)
}
wroteDownstream = true
}
}
}
@@ -368,10 +368,9 @@ func TestProxyOpenAIWSHTTPBridgeTurnRewritesCapacityShedCodeForClient(t *testing
wantErr: true,
},
{
// response.failed 不走 error 事件分支:即便 turn 1 也会被当终止事件
// 原样转发(不 failover),因此改写必须在这里同样生效。
name: "turn1_bare_response_failed",
turn: 1,
// 后续 turn 不允许 replay,容量错误必须改写后交给客户端重试。
name: "turn2_bare_response_failed",
turn: 2,
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",
},
}
@@ -412,6 +411,91 @@ func TestProxyOpenAIWSHTTPBridgeTurnRewritesCapacityShedCodeForClient(t *testing
}
}
func TestProxyOpenAIWSHTTPBridgeTurnStagesMetadataBeforeCapacityFailover(t *testing.T) {
gin.SetMode(gin.TestMode)
body := strings.Join([]string{
`data: {"type":"response.created","response":{"id":"resp_shed"}}`,
"",
`data: {"type":"response.in_progress","response":{"id":"resp_shed"}}`,
"",
`data: {"type":"response.failed","response":{"id":"resp_shed","status":"failed","error":{"message":"Our servers are currently overloaded. Please try again later."}}}`,
"",
}, "\n")
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"X-Request-Id": []string{"rid-ws-bridge-capacity"}},
Body: io.NopCloser(strings.NewReader(body)),
}}
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
account := &Account{ID: 12, 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
result, err := svc.proxyOpenAIWSHTTPBridgeTurn(
context.Background(), c, account, "sk-test", payload, len(payload),
"gpt-5", "", "", "", "", 1,
func(message []byte) error {
writes = append(writes, append([]byte(nil), message...))
return nil
},
)
require.Nil(t, result)
var failoverErr *UpstreamFailoverError
require.ErrorAs(t, err, &failoverErr)
require.True(t, failoverErr.RetryableOnSameAccount)
require.True(t, failoverErr.RequestScopedTransient)
require.Empty(t, writes)
}
func TestProxyOpenAIWSHTTPBridgeTurnDoesNotReplayCapacityAfterSemanticOutput(t *testing.T) {
gin.SetMode(gin.TestMode)
logSink, restore := captureStructuredLog(t)
defer restore()
body := strings.Join([]string{
`data: {"type":"response.created","response":{"id":"resp_partial"}}`,
"",
`data: {"type":"response.output_text.delta","delta":"partial"}`,
"",
`data: {"type":"response.failed","response":{"id":"resp_partial","status":"failed","error":{"code":"server_is_overloaded","message":"Our servers are currently overloaded. Please try again later."}}}`,
"",
}, "\n")
upstream := &httpUpstreamRecorder{resp: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"X-Request-Id": []string{"rid-ws-bridge-post-output"}},
Body: io.NopCloser(strings.NewReader(body)),
}}
svc := &OpenAIGatewayService{cfg: &config.Config{}, httpUpstream: upstream}
account := &Account{ID: 13, 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
result, err := svc.proxyOpenAIWSHTTPBridgeTurn(
context.Background(), c, account, "sk-test", payload, len(payload),
"gpt-5", "", "", "", "", 1,
func(message []byte) error {
writes = append(writes, append([]byte(nil), message...))
return nil
},
)
require.NotNil(t, result)
require.NoError(t, err)
require.Len(t, writes, 3)
var failoverErr *UpstreamFailoverError
require.False(t, errors.As(err, &failoverErr))
require.Contains(t, string(writes[2]), `"code":"server_error"`)
require.NotContains(t, string(writes[2]), "server_is_overloaded")
require.True(t, logSink.ContainsMessage("gateway.failover_suppressed_after_semantic_output"))
require.True(t, logSink.ContainsFieldValue("path", "ws_http_bridge"))
}
func TestProxyOpenAIWSHTTPBridgeTurnRequiresTerminalEvent(t *testing.T) {
gin.SetMode(gin.TestMode)
@@ -423,10 +507,10 @@ func TestProxyOpenAIWSHTTPBridgeTurnRequiresTerminalEvent(t *testing.T) {
}{
{name: "done_without_events_fails_over", body: "data: [DONE]\n\n", wantFailover: true},
{
name: "created_then_done_is_truncated_not_success",
name: "created_then_done_fails_over_before_semantic_output",
body: "data: {\"type\":\"response.created\",\"response\":{\"id\":\"resp_truncated\"}}\n\n" +
"data: [DONE]\n\n",
wantWrites: 1,
wantFailover: true,
},
}
for _, tt := range tests {