From cf5659f02520db7adda2b943d35e218974e30385 Mon Sep 17 00:00:00 2001 From: yucheng Date: Sat, 26 Sep 2026 16:50:09 +0000 Subject: [PATCH] fix(passthrough): close the upstream and log the right reason when the error relay is abandoned before the response Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../pass_through_endpoints.py | 29 ++++++--- .../test_pass_through_endpoints.py | 65 +++++++++++++++++++ 2 files changed, 85 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 9c2286a64a7..29d6fdb9c73 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -891,6 +891,7 @@ class _PreviewReportingStream(httpx.AsyncByteStream): self._budget_crossed = False self._completed = False self._aborted = False + self._abandoned = False self._report_task: asyncio.Task[None] | None = None def dispatch(self) -> asyncio.Task[None]: @@ -901,7 +902,12 @@ class _PreviewReportingStream(httpx.AsyncByteStream): return self._report_task async def _run_report(self) -> None: - if not self._completed and not self._budget_crossed and not self._aborted: + if self._abandoned: + self._log_warning( + "pass_through_endpoint: upstream error response abandoned before reaching the client after %d preview bytes", + sum(len(part) for part in self._collected), + ) + elif not self._completed and not self._budget_crossed and not self._aborted: self._log_warning( "pass_through_endpoint: client disconnected after %d preview bytes of the upstream error body", sum(len(part) for part in self._collected), @@ -911,6 +917,11 @@ class _PreviewReportingStream(httpx.AsyncByteStream): async def report_collected(self) -> None: await asyncio.shield(self.dispatch()) + async def abandon(self) -> None: + self._abandoned = True + self.dispatch() + await self._upstream.aclose() + async def __aiter__(self) -> AsyncIterator[bytes]: total = 0 # rebind-ok: running byte count against the preview budget try: @@ -1037,7 +1048,7 @@ def _passthrough_upstream_failure_reporter( class _UpstreamRelay: response: httpx.Response background: BackgroundTask | None - dispatch: Callable[[], object] | None + abandon: Callable[[], Awaitable[None]] | None async def _log_passthrough_upstream_failure( @@ -1047,7 +1058,7 @@ async def _log_passthrough_upstream_failure( logging_obj: LiteLLMLoggingObj, ) -> _UpstreamRelay: if response.status_code < 400: - return _UpstreamRelay(response=response, background=None, dispatch=None) + return _UpstreamRelay(response=response, background=None, abandon=None) from litellm.proxy.proxy_server import proxy_logging_obj log_warning: Final = verbose_proxy_logger.warning @@ -1056,7 +1067,7 @@ async def _log_passthrough_upstream_failure( ) if response.is_stream_consumed: await report(response.content) - return _UpstreamRelay(response=response, background=None, dispatch=None) + return _UpstreamRelay(response=response, background=None, abandon=None) stream: Final = _PreviewReportingStream( upstream=response, report=report, @@ -1071,7 +1082,7 @@ async def _log_passthrough_upstream_failure( extensions=response.extensions, ), background=BackgroundTask(stream.report_collected), - dispatch=stream.dispatch, + abandon=stream.abandon, ) @@ -1564,8 +1575,8 @@ async def pass_through_request( background=relay.background, ) except BaseException: # noqa: BLE001 # the relay owns the report until the response takes it over - if relay.dispatch is not None: - relay.dispatch() + if relay.abandon is not None: + await relay.abandon() raise if state_raw_body is not None: @@ -1662,8 +1673,8 @@ async def pass_through_request( background=detected_relay.background, ) except BaseException: # noqa: BLE001 # the relay owns the report until the response takes it over - if detected_relay.dispatch is not None: - detected_relay.dispatch() + if detected_relay.abandon is not None: + await detected_relay.abandon() raise if not _should_buffer_passthrough_response(response): 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 cf1ffa42f0e..a32de6a5d17 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 @@ -4172,6 +4172,15 @@ class _UpstreamErrorBodyStream(httpx.AsyncByteStream): yield self._body +class _UpstreamErrorBodyStreamCloseTracking(_UpstreamErrorBodyStream): + def __init__(self, body: bytes) -> None: + super().__init__(body) + self.closed = False + + async def aclose(self) -> None: + self.closed = True + + def _upstream_error_request() -> MagicMock: mock_request: Final = MagicMock(spec=Request) mock_request.method = "POST" @@ -5031,6 +5040,62 @@ async def test_preview_report_collected_reports_once_and_warns_on_disconnect(): log_warning.assert_called_once() +@pytest.mark.asyncio +async def test_preview_reporting_stream_abandon_closes_upstream_and_reports_empty_preview(): + """Abandoning the relay before the client ever reads it must return the + underlying httpx connection and still report, with an empty preview.""" + upstream_stream: Final = _UpstreamErrorBodyStreamCloseTracking(b"body") + upstream_response: Final = httpx.Response( + status_code=500, + headers={"content-type": "text/event-stream"}, + stream=upstream_stream, + request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"), + ) + reported: list[bytes] = [] + + async def report(preview: bytes) -> None: + reported.append(preview) + + relay: Final = _PreviewReportingStream( + upstream=upstream_response, + report=report, + log_warning=MagicMock(), + ) + + await relay.abandon() + await relay.dispatch() + assert upstream_stream.closed is True + assert reported == [b""], reported + + +@pytest.mark.asyncio +async def test_preview_reporting_stream_abandon_does_not_log_a_client_disconnect(): + """An abandoned stream reports its own warning, not the client-disconnect one.""" + upstream_response: Final = httpx.Response( + status_code=500, + headers={"content-type": "text/event-stream"}, + stream=_UpstreamErrorBodyStreamCloseTracking(b"body"), + request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"), + ) + log_warning: Final = MagicMock() + + async def report(preview: bytes) -> None: + pass + + relay: Final = _PreviewReportingStream( + upstream=upstream_response, + report=report, + log_warning=log_warning, + ) + + await relay.abandon() + await relay.dispatch() + warnings: Final = [call.args for call in log_warning.call_args_list] + assert len(warnings) == 1, warnings + assert "abandoned before reaching the client" in warnings[0][0], warnings + assert all("client disconnected" not in call[0] for call in warnings), warnings + + @pytest.mark.asyncio async def test_preview_report_collected_runs_without_disconnect_warning_after_clean_end(): upstream_response: Final = httpx.Response(