diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index e7272c8188..f632803b7d 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -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() { diff --git a/backend/internal/service/openai_account_runtime_block_fastpath.go b/backend/internal/service/openai_account_runtime_block_fastpath.go index 2822374257..6853cc8487 100644 --- a/backend/internal/service/openai_account_runtime_block_fastpath.go +++ b/backend/internal/service/openai_account_runtime_block_fastpath.go @@ -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() diff --git a/backend/internal/service/openai_capacity_shed_test.go b/backend/internal/service/openai_capacity_shed_test.go index eef8261bb6..3f4f273055 100644 --- a/backend/internal/service/openai_capacity_shed_test.go +++ b/backend/internal/service/openai_capacity_shed_test.go @@ -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"}}}`, diff --git a/backend/internal/service/openai_compact_sse_keepalive.go b/backend/internal/service/openai_compact_sse_keepalive.go index 70e6af9454..360c1d55a5 100644 --- a/backend/internal/service/openai_compact_sse_keepalive.go +++ b/backend/internal/service/openai_compact_sse_keepalive.go @@ -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 diff --git a/backend/internal/service/openai_compact_sse_keepalive_test.go b/backend/internal/service/openai_compact_sse_keepalive_test.go index 95ee4ce76c..3703ea9f2a 100644 --- a/backend/internal/service/openai_compact_sse_keepalive_test.go +++ b/backend/internal/service/openai_compact_sse_keepalive_test.go @@ -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) diff --git a/backend/internal/service/openai_first_output_timeout_test.go b/backend/internal/service/openai_first_output_timeout_test.go index 799f7af12b..6783de90d1 100644 --- a/backend/internal/service/openai_first_output_timeout_test.go +++ b/backend/internal/service/openai_first_output_timeout_test.go @@ -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) { diff --git a/backend/internal/service/openai_gateway_passthrough.go b/backend/internal/service/openai_gateway_passthrough.go index f23a656f17..c15a78ef48 100644 --- a/backend/internal/service/openai_gateway_passthrough.go +++ b/backend/internal/service/openai_gateway_passthrough.go @@ -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 通常已扣);必须先解析 diff --git a/backend/internal/service/openai_gateway_response_handling.go b/backend/internal/service/openai_gateway_response_handling.go index 82cd89f695..5577e90e4b 100644 --- a/backend/internal/service/openai_gateway_response_handling.go +++ b/backend/internal/service/openai_gateway_response_handling.go @@ -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 diff --git a/backend/internal/service/openai_gateway_service_codex_cli_only_test.go b/backend/internal/service/openai_gateway_service_codex_cli_only_test.go index 846529828b..c8acab0358 100644 --- a/backend/internal/service/openai_gateway_service_codex_cli_only_test.go +++ b/backend/internal/service/openai_gateway_service_codex_cli_only_test.go @@ -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, "", diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index 617c07c066..6dfb88d803 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -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(), "", "") diff --git a/backend/internal/service/openai_gateway_upstream_errors.go b/backend/internal/service/openai_gateway_upstream_errors.go index 00a62f3f3e..32f04b2c6a 100644 --- a/backend/internal/service/openai_gateway_upstream_errors.go +++ b/backend/internal/service/openai_gateway_upstream_errors.go @@ -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 diff --git a/backend/internal/service/openai_visible_ttft_test.go b/backend/internal/service/openai_visible_ttft_test.go index 8fd243c39e..4d811dfbb3 100644 --- a/backend/internal/service/openai_visible_ttft_test.go +++ b/backend/internal/service/openai_visible_ttft_test.go @@ -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 { diff --git a/backend/internal/service/openai_ws_http_bridge.go b/backend/internal/service/openai_ws_http_bridge.go index 98eee9c7db..1575318526 100644 --- a/backend/internal/service/openai_ws_http_bridge.go +++ b/backend/internal/service/openai_ws_http_bridge.go @@ -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 + } } } diff --git a/backend/internal/service/openai_ws_http_bridge_test.go b/backend/internal/service/openai_ws_http_bridge_test.go index 5478672316..2f0b78d8ba 100644 --- a/backend/internal/service/openai_ws_http_bridge_test.go +++ b/backend/internal/service/openai_ws_http_bridge_test.go @@ -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 {