diff --git a/backend/internal/service/error_passthrough_runtime_test.go b/backend/internal/service/error_passthrough_runtime_test.go index 73a4bfab19..b0eff6c8c0 100644 --- a/backend/internal/service/error_passthrough_runtime_test.go +++ b/backend/internal/service/error_passthrough_runtime_test.go @@ -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() diff --git a/backend/internal/service/openai_account_runtime_block_fastpath.go b/backend/internal/service/openai_account_runtime_block_fastpath.go index 0a17f3b938..d22b95b01f 100644 --- a/backend/internal/service/openai_account_runtime_block_fastpath.go +++ b/backend/internal/service/openai_account_runtime_block_fastpath.go @@ -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) diff --git a/backend/internal/service/openai_gateway_chat_completions.go b/backend/internal/service/openai_gateway_chat_completions.go index 10578a2380..13184047c3 100644 --- a/backend/internal/service/openai_gateway_chat_completions.go +++ b/backend/internal/service/openai_gateway_chat_completions.go @@ -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) diff --git a/backend/internal/service/openai_gateway_chat_completions_test.go b/backend/internal/service/openai_gateway_chat_completions_test.go index b9933e7579..b85ee33947 100644 --- a/backend/internal/service/openai_gateway_chat_completions_test.go +++ b/backend/internal/service/openai_gateway_chat_completions_test.go @@ -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) { diff --git a/backend/internal/service/openai_gateway_service.go b/backend/internal/service/openai_gateway_service.go index 4bb726cd25..af0062db75 100644 --- a/backend/internal/service/openai_gateway_service.go +++ b/backend/internal/service/openai_gateway_service.go @@ -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{ 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 9e2ede0cba..23a1750021 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 @@ -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) diff --git a/backend/internal/service/openai_gateway_service_test.go b/backend/internal/service/openai_gateway_service_test.go index abcf53a269..c11d78e55c 100644 --- a/backend/internal/service/openai_gateway_service_test.go +++ b/backend/internal/service/openai_gateway_service_test.go @@ -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{