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>
This commit is contained in:
yucheng 2026-09-26 16:50:09 +00:00
parent 823c5c8bb8
commit cf5659f025
2 changed files with 85 additions and 9 deletions

View file

@ -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):

View file

@ -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(