From ecc2a94d2ceef7e1018d26291c5ffa5dba2d90a2 Mon Sep 17 00:00:00 2001 From: yucheng Date: Thu, 24 Sep 2026 21:21:08 +0000 Subject: [PATCH] fix(passthrough): relay upstream error streams while the log preview is collected Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../litellm_core_utils/error_normalization.py | 9 +- .../pass_through_endpoints.py | 191 ++++++++++-------- ...t_passthrough_upstream_error_visibility.py | 73 +++++++ .../test_error_normalization.py | 11 + .../test_pass_through_endpoints.py | 74 +++++++ 5 files changed, 268 insertions(+), 90 deletions(-) diff --git a/litellm/litellm_core_utils/error_normalization.py b/litellm/litellm_core_utils/error_normalization.py index be0098ec34b..6f66e8b04ad 100644 --- a/litellm/litellm_core_utils/error_normalization.py +++ b/litellm/litellm_core_utils/error_normalization.py @@ -59,8 +59,9 @@ class _HasProxyErrorType(Protocol): type: str +_UPSTREAM_PASSTHROUGH_PATTERN: Final = re.compile(r"upstream passthrough request failed", re.IGNORECASE) + _MESSAGE_PATTERNS: Final[tuple[tuple[re.Pattern[str], str], ...]] = ( - (re.compile(r"upstream passthrough request failed", re.IGNORECASE), UPSTREAM_PASSTHROUGH), ( re.compile(r"budget has been exceeded|max budget|crossed budget", re.IGNORECASE), BUDGET_EXCEEDED, @@ -192,7 +193,11 @@ def normalize_error(exc: Exception | None, status_code: str, message: str) -> st if by_proxy_type is not None: return by_proxy_type by_message: Final = ( - BUDGET_EXCEEDED if _exceeded_before_budget(message) else _classify_by_message(message, _MESSAGE_PATTERNS) + UPSTREAM_PASSTHROUGH + if _UPSTREAM_PASSTHROUGH_PATTERN.search(message) + else BUDGET_EXCEEDED + if _exceeded_before_budget(message) + else _classify_by_message(message, _MESSAGE_PATTERNS) ) if by_message is not None: return by_message diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index a119335ba46..c81d0ed5d6b 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -5,7 +5,7 @@ import json import posixpath import traceback from base64 import b64encode -from collections.abc import AsyncGenerator, AsyncIterator, Callable, Iterable, Mapping, Sequence +from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterable, Mapping, Sequence from dataclasses import dataclass from datetime import datetime from itertools import count, groupby @@ -37,7 +37,7 @@ from websockets.exceptions import ( from websockets.frames import Close, CloseCode import litellm -from litellm._logging import verbose_proxy_logger +from litellm._logging import redact_secrets, verbose_proxy_logger from litellm._uuid import uuid from litellm.constants import ( MAXIMUM_TRACEBACK_LINES_TO_LOG, @@ -128,6 +128,7 @@ from .upstream_usage_headers import ( if TYPE_CHECKING: from litellm.proxy.proxy_server import ProxyConfig + from litellm.proxy.utils import ProxyLogging router: Final = APIRouter() @@ -865,68 +866,108 @@ def _sanitize_upstream_error_body(body: str) -> str: return " ".join("".join(char if char.isprintable() else " " for char in body).split()) -class _PrefixReplayStream(httpx.AsyncByteStream): - def __init__(self, prefix: bytes, rest: AsyncIterator[bytes], upstream: httpx.Response) -> None: - self._prefix: Final = prefix - self._rest: Final = rest +_ReportPreview = Callable[[bytes], Awaitable[None]] # mutable-ok: Callable arg-list syntax, not a mutable collection + + +class _PreviewReportingStream(httpx.AsyncByteStream): + def __init__( + self, + upstream: httpx.Response, + report: _ReportPreview, + log_warning: Callable[..., None], + ) -> None: self._upstream: Final = upstream + self._report: Final = report + self._log_warning: Final = log_warning + self._collected: Final[list[bytes]] = [] # mutable-ok: preview prefix accumulated while relaying + self._reported = False + + async def _report_once(self) -> None: + if self._reported: + return + self._reported = True + await self._report(b"".join(self._collected)) async def __aiter__(self) -> AsyncIterator[bytes]: - if self._prefix: - yield self._prefix - async for chunk in self._rest: - yield chunk + total = 0 # rebind-ok: running byte count against the preview budget + try: + async for chunk in self._upstream.aiter_bytes(): + if not self._reported: + self._collected.append(chunk) + total += len(chunk) + if total > PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS: + await self._report_once() + yield chunk + except httpx.HTTPError as err: + self._log_warning( + "pass_through_endpoint: upstream error body read failed after %d bytes: %s", + sum(len(part) for part in self._collected), + type(err).__name__, + ) + finally: + await self._report_once() async def aclose(self) -> None: + await self._report_once() await self._upstream.aclose() -async def _no_more_chunks() -> AsyncIterator[bytes]: - return - yield b"" - - -async def _read_error_body_preview( - stream: AsyncIterator[bytes], -) -> tuple[bytes, AsyncIterator[bytes]]: - collected: Final[list[bytes]] = [] # mutable-ok: accumulated until the preview byte budget, then joined once - total = 0 # rebind-ok: running byte count against the preview budget - try: - async for chunk in stream: - collected.append(chunk) - total += len(chunk) - if total > PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS: - break - except httpx.HTTPError as err: - partial: Final = b"".join(collected) - verbose_proxy_logger.warning( - "pass_through_endpoint: upstream error body read failed after %d bytes: %s", - len(partial), - type(err).__name__, - ) - return partial, _no_more_chunks() - return b"".join(collected), stream - - def _headers_without_body_framing(headers: httpx.Headers) -> httpx.Headers: return httpx.Headers( [(name, value) for name, value in headers.raw if name.lower() not in (b"content-encoding", b"content-length")] ) -async def _error_body_preview_and_relay(response: httpx.Response) -> tuple[str, httpx.Response]: - if response.is_stream_consumed: - return response.text, response - body_iter: Final = response.aiter_bytes() - prefix, rest = await _read_error_body_preview(body_iter) - preview_text: Final = prefix.decode(response.encoding or "utf-8", errors="replace") - return preview_text, httpx.Response( - status_code=response.status_code, - headers=_headers_without_body_framing(response.headers), - stream=_PrefixReplayStream(prefix=prefix, rest=rest, upstream=response), - request=response.request, - extensions=response.extensions, - ) +def _passthrough_upstream_failure_reporter( + response: httpx.Response, + user_api_key_dict: UserAPIKeyAuth, + request_payload: dict, + logging_obj: LiteLLMLoggingObj, + proxy_logging: "ProxyLogging", + log_warning: Callable[..., None], +) -> _ReportPreview: + async def report(preview: bytes) -> None: + preview_text: Final = preview.decode(response.encoding or "utf-8", errors="replace") + upstream_error_body: Final = ( + REDACTED_BY_LITELLM + if should_redact_message_logging(logging_obj.model_call_details) + else _truncate_upstream_error_body(_sanitize_upstream_error_body(redact_secrets(preview_text))) + ) + log_warning( + "pass_through_endpoint: upstream %s %s returned %s: %s", + response.request.method, + response.url.copy_with(query=None, fragment=None), + response.status_code, + upstream_error_body, + ) + try: + response.raise_for_status() + except httpx.HTTPStatusError: + # Reported as an HTTPException, not the raw httpx error: ProxyLogging's + # alerting path only excludes HTTPException/ProxyException from its + # "High" severity llm_exceptions alert, treating everything else as an + # operational LLM-API failure. An upstream 4xx/5xx returned unchanged + # to the client is a user-facing error like any other, not something + # ops needs paged for, so it must be excluded the same way auth and + # rate-limit errors already are. + synthetic_exception: Final = HTTPException( + status_code=response.status_code, + detail=f"Upstream passthrough request failed with status {response.status_code}: {upstream_error_body}", + ) + try: + await proxy_logging.post_call_failure_hook( + user_api_key_dict=user_api_key_dict, + original_exception=synthetic_exception, + request_data=request_payload, + traceback_str=traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG), + ) + except Exception: # noqa: BLE001 - a failing logging callback must never break the passthrough response + log_warning( + "pass_through_endpoint: post_call_failure_hook raised for upstream error", + exc_info=True, + ) + + return report async def _log_passthrough_upstream_failure( @@ -939,46 +980,20 @@ async def _log_passthrough_upstream_failure( return response from litellm.proxy.proxy_server import proxy_logging_obj - preview_text, relay_response = await _error_body_preview_and_relay(response) - upstream_error_body: Final = ( - REDACTED_BY_LITELLM - if should_redact_message_logging(logging_obj.model_call_details) - else _truncate_upstream_error_body(_sanitize_upstream_error_body(preview_text)) + log_warning: Final = verbose_proxy_logger.warning + report: Final = _passthrough_upstream_failure_reporter( + response, user_api_key_dict, request_payload, logging_obj, proxy_logging_obj, log_warning ) - verbose_proxy_logger.warning( - "pass_through_endpoint: upstream %s %s returned %s: %s", - response.request.method, - response.url.copy_with(query=None, fragment=None), - response.status_code, - upstream_error_body, + if response.is_stream_consumed: + await report(response.content) + return response + return httpx.Response( + status_code=response.status_code, + headers=_headers_without_body_framing(response.headers), + stream=_PreviewReportingStream(upstream=response, report=report, log_warning=log_warning), + request=response.request, + extensions=response.extensions, ) - try: - response.raise_for_status() - except httpx.HTTPStatusError: - # Reported as an HTTPException, not the raw httpx error: ProxyLogging's - # alerting path only excludes HTTPException/ProxyException from its - # "High" severity llm_exceptions alert, treating everything else as an - # operational LLM-API failure. An upstream 4xx/5xx returned unchanged - # to the client is a user-facing error like any other, not something - # ops needs paged for, so it must be excluded the same way auth and - # rate-limit errors already are. - synthetic_exception: Final = HTTPException( - status_code=response.status_code, - detail=f"Upstream passthrough request failed with status {response.status_code}: {upstream_error_body}", - ) - try: - await proxy_logging_obj.post_call_failure_hook( - user_api_key_dict=user_api_key_dict, - original_exception=synthetic_exception, - request_data=request_payload, - traceback_str=traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG), - ) - except Exception: # noqa: BLE001 - a failing logging callback must never break the passthrough response - verbose_proxy_logger.warning( - "pass_through_endpoint: post_call_failure_hook raised for upstream error", - exc_info=True, - ) - return relay_response async def _relay_reporting_failures( diff --git a/tests/integration/observability/test_passthrough_upstream_error_visibility.py b/tests/integration/observability/test_passthrough_upstream_error_visibility.py index bb18add2f2f..103474ba985 100644 --- a/tests/integration/observability/test_passthrough_upstream_error_visibility.py +++ b/tests/integration/observability/test_passthrough_upstream_error_visibility.py @@ -1,5 +1,6 @@ import gzip import json +import threading from hashlib import sha256 from pathlib import Path from typing import Final @@ -598,3 +599,75 @@ def test_budget_rejected_call_keeps_budget_normalized_error(gateway: Gateway, tm ) ) assert len(budget_rows) == 1, budget_rows + + +_LEAKED_UPSTREAM_KEY: Final = "sk-" + "leak0" * 8 + + +def test_gemini_passthrough_streaming_429_first_frame_reaches_client_while_upstream_holds( + gateway: Gateway, tmp_path: Path +) -> None: + gate: Final = threading.Event() + frames: Final = (b'data: {"error":"rate limited"}\n\n', b"data: [DONE]\n\n") + + def respond(request: Request) -> Reply: + return Reply(status=429, content_type="text/event-stream", chunks=frames, gate_after_first=gate) + + path: Final = tmp_path / "gemini-stream-429.yaml" + with wire_server(respond) as wire: + _gemini_config(path, wire.url) + with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned: + candidate: Final = owned.gateway + with candidate.client.stream( + "POST", + _GEMINI_STREAM_PATH, + params={"alt": "sse"}, + json=_GENERATE_CONTENT, + headers=_gemini_headers(candidate), + timeout=httpx.Timeout(3, connect=5), + ) as response: + assert response.status_code == 429, response.text + iterator: Final = response.iter_bytes() + first: Final = next(iterator) + assert first.startswith(b'data: {"error":"rate limited"}'), first + gate.set() + rest: Final = b"".join(iterator) + assert first + rest == b"".join(frames), first + rest + warning: Final = _upstream_warning(owned.log) + assert "rate limited" in warning, warning + error_information: Final = _spend_error_information(response.headers["x-litellm-call-id"]) + assert error_information["error_code"] == "429", error_information + + +def test_gemini_passthrough_quota_wording_in_upstream_body_keeps_passthrough_normalized_error( + gateway: Gateway, tmp_path: Path +) -> None: + body: Final[dict[str, JsonValue]] = { + "error": { + "message": "You exceeded your current quota, please check your plan and billing details. " + f"Budget for Key={_LEAKED_UPSTREAM_KEY} is spent", + "type": "insufficient_quota", + } + } + + def respond(request: Request) -> Reply: + return Reply(status=429, body=json.dumps(body).encode()) + + path: Final = tmp_path / "gemini-quota-429.yaml" + with wire_server(respond) as wire: + _gemini_config(path, wire.url) + with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned: + candidate: Final = owned.gateway + response: Final = candidate.request( + "POST", _GEMINI_MODEL_PATH, _GENERATE_CONTENT, headers=_gemini_headers(candidate) + ) + assert response.status_code == 429, response.text + assert response.json() == body, response.text + error_information: Final = _spend_error_information(response.headers["x-litellm-call-id"]) + assert error_information["normalized_error"] == "500_UPSTREAM_PASSTHROUGH", error_information + assert error_information["error_code"] == "429", error_information + assert "exceeded your current quota" in str(error_information["error_message"]), error_information + assert _LEAKED_UPSTREAM_KEY not in str(error_information["error_message"]), error_information + warning: Final = _upstream_warning(owned.log) + assert "exceeded your current quota" in warning, warning + assert _LEAKED_UPSTREAM_KEY not in warning, warning diff --git a/tests/test_litellm/litellm_core_utils/test_error_normalization.py b/tests/test_litellm/litellm_core_utils/test_error_normalization.py index 87b9463cb01..750a4fff29a 100644 --- a/tests/test_litellm/litellm_core_utils/test_error_normalization.py +++ b/tests/test_litellm/litellm_core_utils/test_error_normalization.py @@ -177,6 +177,17 @@ def test_normalize_error_passthrough_prefix_wins_over_upstream_body_text() -> No assert normalize_error(exc, "400", message) == "500_UPSTREAM_PASSTHROUGH", message +def test_normalize_error_passthrough_prefix_wins_over_quota_and_budget_wording() -> None: + from fastapi import HTTPException + + detail = ( + 'Upstream passthrough request failed with status 429: {"error": {"message": ' + '"You exceeded your current quota, please check your plan and billing details. Budget spent"}}' + ) + exc = HTTPException(status_code=429, detail=detail) + assert normalize_error(exc, "429", f"429: {detail}") == "500_UPSTREAM_PASSTHROUGH" + + def test_router_no_healthy_deployment_wording_clusters_as_no_healthy_deployments() -> None: for message in (RouterErrors.no_healthy_deployments.value, "No healthy deployments found."): exc = litellm.BadRequestError(message, llm_provider="openai", model="gpt-4o") diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 7a64d5f2218..cbf6f486dd4 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -4569,6 +4569,80 @@ async def test_pass_through_request_streaming_upstream_error_single_large_chunk_ ) +class _GatedUpstreamErrorBodyStream(httpx.AsyncByteStream): + def __init__(self, first: bytes, second: bytes) -> None: + self._first: Final = first + self._second: Final = second + self.gate: Final = asyncio.Event() + + async def __aiter__(self): + yield self._first + await self.gate.wait() + yield self._second + + +@pytest.mark.asyncio +async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_before_preview_completes(): + """ + Regression: while the upstream holds a streaming error response open, the + client must receive the first chunk immediately; the log preview report + must fire exactly once, after the relay ends, with the full body. + """ + first_chunk: Final = b'data: {"error":"rate limited"}\n\n' + second_chunk: Final = b"data: [DONE]\n\n" + body_stream: Final = _GatedUpstreamErrorBodyStream(first_chunk, second_chunk) + upstream_response: Final = httpx.Response( + status_code=429, + headers={"content-type": "text/event-stream"}, + stream=body_stream, + request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"), + ) + + with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client" + ) as mock_get_client: + with patch( + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler" + ) as mock_success_handler: + mock_proxy_logging.pre_call_hook = AsyncMock(return_value={}) + mock_proxy_logging.post_call_failure_hook = AsyncMock() + mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None) + mock_success_handler.return_value = None + + async_client: Final = MagicMock() + async_client.build_request = MagicMock(return_value=MagicMock()) + async_client.send = AsyncMock(return_value=upstream_response) + mock_get_client.return_value = MagicMock(client=async_client) + + response: Final = await pass_through_request( + request=_upstream_error_request(), + target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent", + custom_headers={}, + user_api_key_dict=MagicMock(), + stream=True, + ) + + assert isinstance(response, StreamingResponse) + assert response.status_code == 429 + iterator: Final = response.body_iterator.__aiter__() + first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5) + assert not body_stream.gate.is_set() + body_stream.gate.set() + rest: Final = [chunk async for chunk in iterator] + relayed: Final = b"".join( + chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in (first, *rest) + ) + assert relayed == first_chunk + second_chunk + await upstream_response.aclose() + + mock_proxy_logging.post_call_failure_hook.assert_called_once() + detail: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail + assert ( + detail == 'Upstream passthrough request failed with status 429: data: {"error":"rate limited"} data: [DONE]' + ), detail + + class _UpstreamErrorBodyStreamDropping(httpx.AsyncByteStream): async def __aiter__(self): yield b'{"error": "half'