mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-24 16:05:44 +08:00
Merge pull request #3548 from dftian478/codex/中文-上下文窗口不切号
修复 OpenAI 上下文窗口错误误触发账号切换
This commit is contained in:
@@ -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{
|
||||
|
||||
Reference in New Issue
Block a user