fix(passthrough): hand the upstream error report to the logging worker instead of awaiting it in the relay

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
yucheng 2026-09-24 22:23:05 +00:00
parent b0d65741a5
commit e45275a1e6
2 changed files with 163 additions and 54 deletions

View file

@ -882,44 +882,37 @@ class _PreviewReportingStream(httpx.AsyncByteStream):
self._log_warning: Final = log_warning
self._enqueue: Final = enqueue
self._collected: Final[list[bytes]] = [] # mutable-ok: preview prefix accumulated while relaying
self._reported = False
self._enqueued = False
self._dispatched = False
async def _report_once(self) -> None:
if self._reported:
def _dispatch_report(self) -> None:
if self._dispatched:
return
self._reported = True
await self._report(b"".join(self._collected))
def _enqueue_pending_report(self) -> None:
if self._reported or self._enqueued:
return
self._enqueued = True
self._enqueue(self._report_once())
self._dispatched = True
self._enqueue(self._report(b"".join(self._collected)))
async def __aiter__(self) -> AsyncIterator[bytes]:
total = 0 # rebind-ok: running byte count against the preview budget
try:
async for chunk in self._upstream.aiter_bytes():
if not self._reported:
if not self._dispatched:
self._collected.append(chunk)
total += len(chunk)
if total > PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS:
await self._report_once()
self._dispatch_report()
yield chunk
await self._report_once()
self._dispatch_report()
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__,
)
await self._report_once()
self._dispatch_report()
finally:
self._enqueue_pending_report()
self._dispatch_report()
async def aclose(self) -> None:
self._enqueue_pending_report()
self._dispatch_report()
await self._upstream.aclose()

View file

@ -4258,6 +4258,7 @@ async def test_pass_through_request_streaming_upstream_error_body_reaches_client
)
assert streamed_bytes == upstream_content
await _poll(lambda: mock_proxy_logging.post_call_failure_hook.called)
mock_proxy_logging.post_call_failure_hook.assert_called_once()
original_exception: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"]
assert "was not found or your project" in original_exception.detail
@ -4453,15 +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"),
)
served_at_warning: list[int] = []
real_warning: Final = verbose_proxy_logger.warning
enqueued: list[Coroutine[None, None, None]] = []
served_at_dispatch: list[int] = []
def _recording_warning(*args, **kwargs):
if args and args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s":
served_at_warning.append(body_stream.served)
return real_warning(*args, **kwargs)
def _recording_enqueue(coro):
served_at_dispatch.append(body_stream.served)
enqueued.append(coro)
with patch.object(verbose_proxy_logger, "warning", side_effect=_recording_warning):
with patch(
"litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
side_effect=_recording_enqueue,
):
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"
@ -4495,9 +4498,11 @@ async def test_pass_through_request_streaming_upstream_error_reads_only_preview_
)
assert streamed_bytes == upstream_content
assert served_at_warning == [5], (
"each raw chunk is yielded as-is; five 1024-byte chunks are the first point the preview budget is exceeded"
assert served_at_dispatch == [5], (
"the report is dispatched after the fifth 1024-byte chunk, the first point the preview budget is exceeded"
)
assert len(enqueued) == 1, enqueued
await enqueued[0]
expected_body: Final = f"{'x' * 4096}... (truncated at 4096 chars)"
assert (
mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
@ -4518,15 +4523,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"),
)
served_at_warning: list[int] = []
real_warning: Final = verbose_proxy_logger.warning
enqueued: list[Coroutine[None, None, None]] = []
served_at_dispatch: list[int] = []
def _recording_warning(*args, **kwargs):
if args and args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s":
served_at_warning.append(body_stream.served)
return real_warning(*args, **kwargs)
def _recording_enqueue(coro):
served_at_dispatch.append(body_stream.served)
enqueued.append(coro)
with patch.object(verbose_proxy_logger, "warning", side_effect=_recording_warning):
with patch(
"litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
side_effect=_recording_enqueue,
):
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"
@ -4560,9 +4567,11 @@ async def test_pass_through_request_streaming_upstream_error_single_large_chunk_
)
assert streamed_bytes == upstream_content
assert served_at_warning == [1], (
"the rechunked preview is served from the first raw chunk; the second must not be pulled before the warning"
assert served_at_dispatch == [1], (
"the report is dispatched after the first raw chunk crosses the preview budget; the second is not pulled first"
)
assert len(enqueued) == 1, enqueued
await enqueued[0]
expected_body: Final = f"{'x' * 4096}... (truncated at 4096 chars)"
assert (
mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
@ -4570,6 +4579,13 @@ async def test_pass_through_request_streaming_upstream_error_single_large_chunk_
)
async def _poll(condition: Callable[[], bool], seconds: float = 5) -> None:
deadline: Final = asyncio.get_running_loop().time() + seconds
while not condition():
assert asyncio.get_running_loop().time() < deadline, "condition not met in time"
await asyncio.sleep(0.01)
class _GatedUpstreamErrorBodyStream(httpx.AsyncByteStream):
def __init__(self, first: bytes, second: bytes) -> None:
self._first: Final = first
@ -4599,6 +4615,8 @@ 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]] = []
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"
@ -4606,29 +4624,34 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
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
with patch(
"litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
side_effect=lambda coro: enqueued.append(coro),
):
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)
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,
)
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__()
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()
mock_proxy_logging.post_call_failure_hook.assert_not_called()
body_stream.gate.set()
rest: Final = [chunk async for chunk in iterator]
relayed: Final = b"".join(
@ -4637,6 +4660,8 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_
assert relayed == first_chunk + second_chunk
await upstream_response.aclose()
assert len(enqueued) == 1, enqueued
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 (
@ -4712,6 +4737,81 @@ async def test_pass_through_request_streaming_upstream_error_client_disconnect_e
assert detail == 'Upstream passthrough request failed with status 429: data: {"error":"rate limited"}', detail
@pytest.mark.asyncio
async def test_pass_through_request_streaming_upstream_error_yields_over_budget_chunk_before_report_finishes():
"""
Regression: the failure report is never awaited inside the relay, so a
single chunk that crosses the preview budget reaches the client even
while the report coroutine is still running.
"""
first_chunk: Final = b"x" * 6144
release: Final = asyncio.Event()
body_stream: Final = _ChunkedUpstreamErrorBodyStream((first_chunk,))
upstream_response: Final = httpx.Response(
status_code=500,
headers={"content-type": "text/plain"},
stream=body_stream,
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
enqueued: list[Coroutine[None, None, None]] = []
async def _held_hook(**kwargs):
await release.wait()
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:
with patch(
"litellm.litellm_core_utils.logging_worker.GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue",
side_effect=lambda coro: enqueued.append(coro),
):
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock(side_effect=_held_hook)
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
iterator: Final = response.body_iterator.__aiter__()
first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5)
assert not release.is_set()
relayed: Final = b"".join(
[first]
+ [chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") async for chunk in iterator]
)
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
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 == f"Upstream passthrough request failed with status 500: {'x' * 4096}... (truncated at 4096 chars)"
), detail
class _UpstreamErrorBodyStreamDropping(httpx.AsyncByteStream):
async def __aiter__(self):
yield b'{"error": "half'
@ -4775,6 +4875,11 @@ async def test_pass_through_request_streaming_upstream_error_body_read_failure_k
assert streamed_bytes == b'{"error": "half'
await upstream_response.aclose()
await _poll(
lambda: any(
args and args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s" for args in recorded_warnings
)
)
rendered: Final = [str(args[0]) for args in recorded_warnings]
formats: Final = [args[0] for args in recorded_warnings]
assert any(
@ -4859,6 +4964,11 @@ async def test_pass_through_request_streaming_upstream_error_gzip_read_failure_r
assert streamed_bytes == plaintext
await upstream_response.aclose()
await _poll(
lambda: any(
args and args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s" for args in recorded_warnings
)
)
rendered: Final = [str(args[0]) for args in recorded_warnings]
assert any(
args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s" and plaintext.decode() in str(args[4])
@ -4910,6 +5020,12 @@ async def test_pass_through_request_streaming_upstream_error_gzip_body_decoded_f
)
assert streamed_bytes == upstream_content
await _poll(
lambda: any(
call.args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s"
for call in mock_warning.call_args_list
)
)
upstream_warnings: Final = [
call
for call in mock_warning.call_args_list