fix(passthrough): await the shielded upstream error report at stream end so a delivered error always lands its row

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-25 00:36:43 +00:00
parent b6ba92a641
commit db834354bb
2 changed files with 82 additions and 29 deletions

View file

@ -875,20 +875,26 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
upstream: httpx.Response,
report: _ReportPreview,
log_warning: Callable[..., None],
enqueue: Callable[[Coroutine[None, None, None]], None],
spawn: Callable[[Coroutine[None, None, None]], asyncio.Future[None]],
) -> None:
self._upstream: Final = upstream
self._report: Final = report
self._log_warning: Final = log_warning
self._enqueue: Final = enqueue
self._spawn: Final = spawn
self._collected: Final[list[bytes]] = [] # mutable-ok: preview prefix accumulated while relaying
self._dispatched = False
self._pending: asyncio.Future[None] | None = None
def _dispatch_report(self) -> None:
if self._dispatched:
return
self._dispatched = True
self._enqueue(self._report(b"".join(self._collected)))
self._pending = self._spawn(self._report(b"".join(self._collected)))
async def _drain_pending_report(self) -> None:
pending: Final = self._pending
if pending is not None:
await asyncio.shield(pending)
async def __aiter__(self) -> AsyncIterator[bytes]:
total = 0 # rebind-ok: running byte count against the preview budget
@ -901,6 +907,7 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
self._dispatch_report()
yield chunk
self._dispatch_report()
await self._drain_pending_report()
except httpx.HTTPError as err:
self._log_warning(
"pass_through_endpoint: upstream error body read failed after %d bytes: %s",
@ -908,6 +915,7 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
type(err).__name__,
)
self._dispatch_report()
await self._drain_pending_report()
finally:
self._dispatch_report()
@ -916,6 +924,10 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
await self._upstream.aclose()
def _spawn_report_task(report: Coroutine[None, None, None]) -> asyncio.Future[None]:
return asyncio.ensure_future(report)
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")]
@ -998,7 +1010,7 @@ async def _log_passthrough_upstream_failure(
upstream=response,
report=report,
log_warning=log_warning,
enqueue=GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue,
spawn=_spawn_report_task,
),
request=response.request,
extensions=response.extensions,

View file

@ -6,7 +6,7 @@ import logging
import os
import sys
import zlib
from collections.abc import Callable, Coroutine
from collections.abc import Callable
from contextlib import ExitStack, contextmanager
from io import BytesIO
from types import SimpleNamespace
@ -4454,16 +4454,17 @@ async def test_pass_through_request_streaming_upstream_error_reads_only_preview_
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
enqueued: list[Coroutine[None, None, None]] = []
enqueued: list[asyncio.Future[None]] = []
served_at_dispatch: list[int] = []
def _recording_enqueue(coro):
def _recording_spawn(coro):
served_at_dispatch.append(body_stream.served)
enqueued.append(coro)
enqueued.append(asyncio.ensure_future(coro))
return enqueued[-1]
with patch(
"litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
side_effect=_recording_enqueue,
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
side_effect=_recording_spawn,
):
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
@ -4523,16 +4524,17 @@ async def test_pass_through_request_streaming_upstream_error_single_large_chunk_
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
enqueued: list[Coroutine[None, None, None]] = []
enqueued: list[asyncio.Future[None]] = []
served_at_dispatch: list[int] = []
def _recording_enqueue(coro):
def _recording_spawn(coro):
served_at_dispatch.append(body_stream.served)
enqueued.append(coro)
enqueued.append(asyncio.ensure_future(coro))
return enqueued[-1]
with patch(
"litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
side_effect=_recording_enqueue,
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
side_effect=_recording_spawn,
):
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
@ -4615,7 +4617,7 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
enqueued: list[Coroutine[None, None, None]] = []
enqueued: list[asyncio.Future[None]] = []
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
@ -4625,8 +4627,8 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
with patch(
"litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
side_effect=lambda coro: enqueued.append(coro),
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
side_effect=lambda coro: enqueued.append(asyncio.ensure_future(coro)) or enqueued[-1],
):
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
@ -4687,7 +4689,7 @@ async def test_pass_through_request_streaming_upstream_error_client_disconnect_e
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
enqueued: list[Coroutine[None, None, None]] = []
enqueued: list[asyncio.Future[None]] = []
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
@ -4697,8 +4699,8 @@ async def test_pass_through_request_streaming_upstream_error_client_disconnect_e
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
with patch(
"litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
side_effect=lambda coro: enqueued.append(coro),
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
side_effect=lambda coro: enqueued.append(asyncio.ensure_future(coro)) or enqueued[-1],
):
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
@ -4754,7 +4756,7 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
enqueued: list[Coroutine[None, None, None]] = []
enqueued: list[asyncio.Future[None]] = []
async def _held_hook(**kwargs):
await release.wait()
@ -4767,8 +4769,8 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
with patch(
"litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
side_effect=lambda coro: enqueued.append(coro),
"litellm.proxy.pass_through_endpoints.pass_through_endpoints._spawn_report_task",
side_effect=lambda coro: enqueued.append(asyncio.ensure_future(coro)) or enqueued[-1],
):
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock(side_effect=_held_hook)
@ -4793,6 +4795,8 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
iterator: Final = response.body_iterator.__aiter__()
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
assert not release.is_set()
assert not enqueued[0].done()
release.set()
relayed: Final = b"".join(
[first]
+ [chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") async for chunk in iterator]
@ -4800,11 +4804,7 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
assert relayed == first_chunk
assert len(enqueued) == 1, enqueued
report_task: Final = asyncio.ensure_future(enqueued[0])
await asyncio.sleep(0)
assert not report_task.done()
release.set()
await report_task
await enqueued[0]
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 (
@ -4812,6 +4812,47 @@ async def test_pass_through_request_streaming_upstream_error_yields_over_budget_
), detail
@pytest.mark.asyncio
async def test_pass_through_request_streaming_upstream_error_drains_the_report_before_stream_end():
upstream_response: Final = httpx.Response(
status_code=500,
headers={"content-type": "text/plain"},
stream=_ChunkedUpstreamErrorBodyStream((b"x" * 512,)),
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 == 500
drained: Final = [chunk async for chunk in response.body_iterator]
assert drained == [b"x" * 512]
mock_proxy_logging.post_call_failure_hook.assert_called_once()
class _UpstreamErrorBodyStreamDropping(httpx.AsyncByteStream):
async def __aiter__(self):
yield b'{"error": "half'