diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 9078b7c134f..2858cf3f0b5 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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() 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 bf62e8cfddd..77fdd55a055 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 @@ -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