mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:48:43 +08:00
Merge pull request #5676 from Perfecto23/agent/openai-capacity-failover
fix(openai): recover message-only capacity failures before output
This commit is contained in:
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user