Merge pull request #3548 from dftian478/codex/中文-上下文窗口不切号

修复 OpenAI 上下文窗口错误误触发账号切换
This commit is contained in:
Wesley Liddick
2026-06-30 10:57:26 +08:00
committed by GitHub
7 changed files with 218 additions and 17 deletions
@@ -4,6 +4,7 @@ import (
"bytes"
"context"
"encoding/json"
"errors"
"io"
"net/http"
"net/http/httptest"
@@ -88,6 +89,35 @@ func TestOpenAIHandleErrorResponse_NoRuleKeepsDefault(t *testing.T) {
assert.Equal(t, "Upstream request failed", errField["message"])
}
func TestOpenAIHandleErrorResponse_ContextWindow502KeepsMessageWithoutFailover(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/", nil)
svc := &OpenAIGatewayService{}
respBody := []byte(`{"error":{"message":"Your input exceeds the context window of this model. Please adjust your input and try again.","type":"upstream_error","code":null}}`)
resp := &http.Response{
StatusCode: http.StatusBadGateway,
Body: io.NopCloser(bytes.NewReader(respBody)),
Header: http.Header{},
}
account := &Account{ID: 14, Platform: PlatformOpenAI, Type: AccountTypeAPIKey}
_, err := svc.handleErrorResponse(context.Background(), resp, c, account, nil)
require.Error(t, err)
var failoverErr *UpstreamFailoverError
require.False(t, errors.As(err, &failoverErr))
assert.Equal(t, http.StatusBadGateway, rec.Code)
var payload map[string]any
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &payload))
errField, ok := payload["error"].(map[string]any)
require.True(t, ok)
assert.Equal(t, "upstream_error", errField["type"])
assert.Equal(t, "Your input exceeds the context window of this model. Please adjust your input and try again.", errField["message"])
}
func TestGeminiWriteGeminiMappedError_NoRuleKeepsDefault(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
@@ -39,6 +39,10 @@ func (s *OpenAIGatewayService) handleOpenAIAccountUpstreamError(ctx context.Cont
stateCtx, cancel := openAIAccountStateContext(ctx)
defer cancel()
if account != nil && account.Platform == PlatformOpenAI && isOpenAIContextWindowError("", responseBody) {
return false
}
if isOpenAIImageRateLimitError(statusCode, responseBody) {
if s != nil && s.rateLimitService != nil {
_ = s.rateLimitService.HandleOpenAIImageRateLimit(stateCtx, account, statusCode, headers, responseBody)
@@ -450,7 +450,13 @@ func (s *OpenAIGatewayService) handleChatBufferedStreamingResponse(
writeChatCompletionsError(c, http.StatusBadRequest, "invalid_request_error", clientMsg)
return nil, fmt.Errorf("openai cyber_policy: %s", msg)
}
return nil, s.newOpenAIStreamFailoverError(c, account, false, requestID, payload, openAICompatFailedResponseMessage(finalResponse))
message := openAICompatFailedResponseMessage(finalResponse)
if openAIStreamFailedEventShouldFailover(payload, message) {
return nil, s.newOpenAIStreamFailoverError(c, account, false, requestID, payload, message)
}
message = s.recordOpenAIStreamUpstreamError(c, account, false, requestID, "http_error", payload, message)
writeChatCompletionsError(c, http.StatusBadGateway, "upstream_error", message)
return nil, fmt.Errorf("upstream response failed: %s", message)
}
// When the terminal event has an empty output array, reconstruct from
@@ -524,6 +530,7 @@ func (s *OpenAIGatewayService) handleChatStreamingResponse(
pendingSSE := make([]string, 0, 4)
refusalDetector := newOpenAIChatSilentRefusalDetector(requestBodyLen)
var streamFailoverErr *UpstreamFailoverError
var streamNonFailoverErr error
scanner := bufio.NewScanner(resp.Body)
maxLineSize := defaultMaxLineSize
@@ -618,10 +625,34 @@ func (s *OpenAIGatewayService) handleChatStreamingResponse(
clientDisconnected = true
}
return true
} else {
}
if openAIStreamFailedEventShouldFailover(payloadBytes, message) {
streamFailoverErr = s.newOpenAIStreamFailoverError(c, account, false, requestID, payloadBytes, message)
return true
}
message = s.recordOpenAIStreamUpstreamError(c, account, false, requestID, "http_error", payloadBytes, message)
errorPayload, _ := json.Marshal(gin.H{
"error": gin.H{
"type": "upstream_error",
"message": message,
},
})
if c != nil && c.Writer != nil && !c.Writer.Written() {
writeChatCompletionsError(c, http.StatusBadGateway, "upstream_error", message)
clientOutputStarted = true
} else if c != nil && c.Writer != nil && !clientDisconnected {
if _, err := fmt.Fprintf(c.Writer, "data: %s\n\n", errorPayload); err != nil {
clientDisconnected = true
logger.L().Info("openai chat_completions stream: client disconnected while writing upstream error",
zap.String("request_id", requestID),
)
}
}
if !clientDisconnected {
c.Writer.Flush()
}
streamNonFailoverErr = fmt.Errorf("upstream response failed: %s", message)
return true
}
chunks := apicompat.ResponsesEventToChatChunks(&event, state)
@@ -679,6 +710,9 @@ func (s *OpenAIGatewayService) handleChatStreamingResponse(
}
return resultWithUsage(), streamFailoverErr
}
if streamNonFailoverErr != nil {
return resultWithUsage(), streamNonFailoverErr
}
if finalChunks := apicompat.FinalizeResponsesChatStream(state); len(finalChunks) > 0 && !clientDisconnected {
for _, chunk := range finalChunks {
refusalDetector.ObserveChatChunk(chunk)
@@ -269,7 +269,7 @@ func TestForwardAsChatCompletions_ClientDisconnectDrainsUpstreamUsage(t *testing
require.Equal(t, 4, result.Usage.CacheReadInputTokens)
}
func TestForwardAsChatCompletions_BufferedResponseFailedTriggersFailover(t *testing.T) {
func TestForwardAsChatCompletions_BufferedContextWindowResponseFailedReturnsErrorWithoutFailover(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
@@ -306,13 +306,13 @@ func TestForwardAsChatCompletions_BufferedResponseFailedTriggersFailover(t *test
require.Error(t, err)
require.Nil(t, result)
var failoverErr *UpstreamFailoverError
require.ErrorAs(t, err, &failoverErr)
require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode)
require.Contains(t, string(failoverErr.ResponseBody), "input exceeds the context window")
require.False(t, c.Writer.Written())
require.False(t, errors.As(err, &failoverErr))
require.True(t, c.Writer.Written())
require.Equal(t, http.StatusBadGateway, rec.Code)
require.Contains(t, rec.Body.String(), "input exceeds the context window")
}
func TestForwardAsChatCompletions_StreamResponseFailedTriggersFailoverBeforeFlush(t *testing.T) {
func TestForwardAsChatCompletions_StreamContextWindowResponseFailedReturnsErrorWithoutFailover(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
@@ -349,12 +349,14 @@ func TestForwardAsChatCompletions_StreamResponseFailedTriggersFailoverBeforeFlus
result, err := svc.ForwardAsChatCompletions(context.Background(), c, account, body, "", "gpt-5.5")
require.Error(t, err)
require.Nil(t, result)
require.NotNil(t, result)
var failoverErr *UpstreamFailoverError
require.ErrorAs(t, err, &failoverErr)
require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode)
require.Contains(t, string(failoverErr.ResponseBody), "input exceeds the context window")
require.False(t, c.Writer.Written())
require.False(t, errors.As(err, &failoverErr))
require.True(t, c.Writer.Written())
require.Equal(t, http.StatusBadGateway, rec.Code)
require.Contains(t, rec.Header().Get("Content-Type"), "application/json")
require.Contains(t, rec.Body.String(), "input exceeds the context window")
require.NotContains(t, rec.Body.String(), "[DONE]")
}
func TestForwardAsChatCompletions_StreamCyberPolicyNoFailover(t *testing.T) {
@@ -1196,6 +1196,51 @@ func isOpenAITransientProcessingError(upstreamStatusCode int, upstreamMsg string
return match(string(upstreamBody))
}
func isOpenAIContextWindowError(upstreamMsg string, upstreamBody []byte) bool {
match := func(text string) bool {
lower := strings.ToLower(strings.TrimSpace(text))
if lower == "" {
return false
}
if strings.Contains(lower, "context_too_large") || strings.Contains(lower, "context_length_exceeded") {
return true
}
if strings.Contains(lower, "maximum context length") || strings.Contains(lower, "max context length") {
return true
}
hasExceeded := strings.Contains(lower, "exceed") || strings.Contains(lower, "too large") || strings.Contains(lower, "too long")
if strings.Contains(lower, "context window") && hasExceeded {
return true
}
if strings.Contains(lower, "context length") && hasExceeded {
return true
}
return strings.Contains(lower, "token limit") &&
strings.Contains(lower, "context") &&
hasExceeded
}
if match(upstreamMsg) {
return true
}
if len(upstreamBody) == 0 {
return false
}
for _, path := range []string{
"error.message",
"response.error.message",
"message",
"error.code",
"response.error.code",
"code",
} {
if match(gjson.GetBytes(upstreamBody, path).String()) {
return true
}
}
return match(string(upstreamBody))
}
// ExtractSessionID extracts the raw session ID from headers or body without hashing.
// Used by ForwardAsAnthropic to pass as prompt_cache_key for upstream cache.
func (s *OpenAIGatewayService) ExtractSessionID(c *gin.Context, body []byte) string {
@@ -2450,6 +2495,9 @@ func (s *OpenAIGatewayService) shouldFailoverUpstreamError(statusCode int) bool
}
func (s *OpenAIGatewayService) shouldFailoverOpenAIUpstreamResponse(statusCode int, upstreamMsg string, upstreamBody []byte) bool {
if isOpenAIContextWindowError(upstreamMsg, upstreamBody) {
return false
}
if s.shouldFailoverUpstreamError(statusCode) {
return true
}
@@ -3854,6 +3902,9 @@ func openAIStreamDataStartsClientOutput(data, eventType string) bool {
}
func openAIStreamFailedEventShouldFailover(payload []byte, message string) bool {
if isOpenAIContextWindowError(message, payload) {
return false
}
if isOpenAITransientProcessingError(http.StatusBadRequest, message, payload) {
return true
}
@@ -3886,17 +3937,18 @@ func openAIStreamFailedEventShouldFailover(payload []byte, message string) bool
return true
}
func (s *OpenAIGatewayService) newOpenAIStreamFailoverError(
func (s *OpenAIGatewayService) recordOpenAIStreamUpstreamError(
c *gin.Context,
account *Account,
passthrough bool,
upstreamRequestID string,
kind string,
payload []byte,
message string,
) *UpstreamFailoverError {
) string {
message = sanitizeUpstreamErrorMessage(strings.TrimSpace(message))
if message == "" {
message = "OpenAI stream disconnected before completion"
message = "OpenAI upstream response failed"
}
detail := ""
if len(payload) > 0 && s != nil && s.cfg != nil && s.cfg.Gateway.LogUpstreamErrorBody {
@@ -3913,7 +3965,7 @@ func (s *OpenAIGatewayService) newOpenAIStreamFailoverError(
UpstreamStatusCode: http.StatusBadGateway,
UpstreamRequestID: strings.TrimSpace(upstreamRequestID),
Passthrough: passthrough,
Kind: "failover",
Kind: kind,
Message: message,
Detail: detail,
}
@@ -3924,6 +3976,22 @@ func (s *OpenAIGatewayService) newOpenAIStreamFailoverError(
}
appendOpsUpstreamError(c, event)
}
return message
}
func (s *OpenAIGatewayService) newOpenAIStreamFailoverError(
c *gin.Context,
account *Account,
passthrough bool,
upstreamRequestID string,
payload []byte,
message string,
) *UpstreamFailoverError {
message = sanitizeUpstreamErrorMessage(strings.TrimSpace(message))
if message == "" {
message = "OpenAI stream disconnected before completion"
}
message = s.recordOpenAIStreamUpstreamError(c, account, passthrough, upstreamRequestID, "failover", payload, message)
body, _ := json.Marshal(gin.H{
"error": gin.H{
"type": "upstream_error",
@@ -4603,6 +4671,9 @@ func (s *OpenAIGatewayService) handleErrorResponse(
errType = "upstream_error"
errMsg = "Upstream request failed"
}
if isOpenAIContextWindowError(upstreamMsg, body) && upstreamMsg != "" {
errMsg = upstreamMsg
}
c.JSON(statusCode, gin.H{
"error": gin.H{
@@ -253,6 +253,29 @@ func TestIsOpenAITransientProcessingError(t *testing.T) {
))
}
func TestIsOpenAIContextWindowError(t *testing.T) {
require.True(t, isOpenAIContextWindowError(
"",
[]byte(`{"error":{"message":"Your input exceeds the context window of this model. Please adjust your input and try again.","type":"upstream_error","code":null}}`),
))
require.True(t, isOpenAIContextWindowError(
"maximum context length exceeded",
nil,
))
require.False(t, isOpenAIContextWindowError(
"context canceled",
nil,
))
}
func TestShouldFailoverOpenAIUpstreamResponseContextWindow502(t *testing.T) {
svc := &OpenAIGatewayService{}
body := []byte(`{"error":{"message":"Your input exceeds the context window of this model. Please adjust your input and try again.","type":"upstream_error","code":null}}`)
require.False(t, svc.shouldFailoverOpenAIUpstreamResponse(http.StatusBadGateway, "", body))
require.True(t, svc.shouldFailoverOpenAIUpstreamResponse(http.StatusBadGateway, "temporary upstream outage", []byte(`{"error":{"message":"temporary upstream outage"}}`)))
}
func TestOpenAIGatewayService_Forward_LogsInstructionsRequiredDetails(t *testing.T) {
gin.SetMode(gin.TestMode)
logSink, restore := captureStructuredLog(t)
@@ -1417,6 +1417,43 @@ func TestOpenAIStreamingResponseFailedAfterOutputSanitizesVerboseResponseForClie
require.NotContains(t, body, `"usage"`)
}
func TestOpenAIStreamingContextWindowResponseFailedBeforeOutputPassesThrough(t *testing.T) {
gin.SetMode(gin.TestMode)
cfg := &config.Config{
Gateway: config.GatewayConfig{
StreamDataIntervalTimeout: 0,
StreamKeepaliveInterval: 0,
MaxLineSize: defaultMaxLineSize,
},
}
svc := &OpenAIGatewayService{cfg: cfg}
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/", nil)
resp := &http.Response{
StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(strings.Join([]string{
"event: response.created",
`data: {"type":"response.created","response":{"id":"resp_1"}}`,
"",
"event: response.failed",
`data: {"type":"response.failed","error":{"type":"upstream_error","message":"Your input exceeds the context window of this model. Please adjust your input and try again.","code":null}}`,
"",
}, "\n"))),
Header: http.Header{"X-Request-Id": []string{"rid-context-window-failed"}},
}
_, err := svc.handleStreamingResponse(c.Request.Context(), resp, c, &Account{ID: 1, Platform: PlatformOpenAI, Name: "acc"}, time.Now(), "model", "model")
require.Error(t, err)
var failoverErr *UpstreamFailoverError
require.False(t, errors.As(err, &failoverErr))
require.True(t, c.Writer.Written())
require.Contains(t, rec.Body.String(), "response.failed")
require.Contains(t, rec.Body.String(), "Your input exceeds the context window")
}
func TestOpenAIStreamingPreambleOnlyMissingTerminalReturnsFailover(t *testing.T) {
gin.SetMode(gin.TestMode)
cfg := &config.Config{