diff --git a/litellm/constants.py b/litellm/constants.py index 7b40f432446..eb97420aaa1 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1520,6 +1520,12 @@ CLOUDZERO_EXPORT_INTERVAL_MINUTES: Final = int(os.getenv("CLOUDZERO_EXPORT_INTER MCP_TOOL_NAME_PREFIX: Final = "mcp_tool" MAXIMUM_TRACEBACK_LINES_TO_LOG: Final = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG", 100)) PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS: Final = 4096 +PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY: Final = int( + os.getenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY", "64") +) +PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS: Final = int( + os.getenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS", "10") +) # Headers to control callbacks X_LITELLM_DISABLE_CALLBACKS: Final = "x-litellm-disable-callbacks" diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 0b10ca4b05c..c17f4977ffc 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -5,7 +5,7 @@ import json import posixpath import traceback from base64 import b64encode -from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Coroutine, Iterable, Mapping, Sequence +from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterable, Mapping, Sequence from dataclasses import dataclass from datetime import datetime from itertools import count, groupby @@ -42,6 +42,8 @@ from litellm._uuid import uuid from litellm.constants import ( MAXIMUM_TRACEBACK_LINES_TO_LOG, PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS, + PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY, + PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS, REDACTED_BY_LITELLM, SESSION_ID_OMITTED_METADATA_KEY, WEBSOCKET_CLOSE_REASON_MAX_BYTES, @@ -875,7 +877,7 @@ class _PreviewReportingStream(httpx.AsyncByteStream): upstream: httpx.Response, report: _ReportPreview, log_warning: Callable[..., None], - spawn: Callable[[Coroutine[None, None, None]], asyncio.Future[None]], + spawn: Callable[[Awaitable[None]], asyncio.Future[None]], ) -> None: self._upstream: Final = upstream self._report: Final = report @@ -884,6 +886,7 @@ class _PreviewReportingStream(httpx.AsyncByteStream): self._collected: Final[list[bytes]] = [] # mutable-ok: preview prefix accumulated while relaying self._dispatched = False self._pending: asyncio.Future[None] | None = None + self._completed = False def _dispatch_report(self) -> None: if self._dispatched: @@ -896,6 +899,14 @@ class _PreviewReportingStream(httpx.AsyncByteStream): if pending is not None: await asyncio.shield(pending) + def _dispatch_disconnect_report(self) -> None: + if not self._completed and not self._dispatched: + 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), + ) + self._dispatch_report() + async def __aiter__(self) -> AsyncIterator[bytes]: total = 0 # rebind-ok: running byte count against the preview budget try: @@ -906,9 +917,11 @@ class _PreviewReportingStream(httpx.AsyncByteStream): if total > PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS: self._dispatch_report() yield chunk + self._completed = True self._dispatch_report() await self._drain_pending_report() except httpx.HTTPError as err: + dispatched_before: Final = self._dispatched self._log_warning( "pass_through_endpoint: upstream error body read failed after %d bytes: %s", sum(len(part) for part in self._collected), @@ -916,16 +929,38 @@ class _PreviewReportingStream(httpx.AsyncByteStream): ) self._dispatch_report() await self._drain_pending_report() + if dispatched_before: + raise finally: - self._dispatch_report() + self._dispatch_disconnect_report() async def aclose(self) -> None: - self._dispatch_report() + self._dispatch_disconnect_report() await self._upstream.aclose() -def _spawn_report_task(report: Coroutine[None, None, None]) -> asyncio.Future[None]: - return asyncio.ensure_future(report) +_REPORT_CONCURRENCY: Final = asyncio.Semaphore(PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY) +_REPORT_TASKS: Final[set[asyncio.Future[None]]] = set() # mutable-ok: in-flight report registry drained at shutdown + + +async def _bounded_report(report: Awaitable[None], limiter: asyncio.Semaphore) -> None: + async with limiter: + await report + + +def _spawn_report_task(report: Awaitable[None]) -> asyncio.Future[None]: + task: Final = asyncio.ensure_future(_bounded_report(report, _REPORT_CONCURRENCY)) + _REPORT_TASKS.add(task) + task.add_done_callback(_REPORT_TASKS.discard) + return task + + +async def drain_passthrough_upstream_error_reports( + timeout: float = PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS, +) -> None: + pending: Final = tuple(_REPORT_TASKS) + if pending: + await asyncio.wait(pending, timeout=timeout) def _headers_without_body_framing(headers: httpx.Headers) -> httpx.Headers: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index c0071aa7c81..ac9c8ea8374 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -730,6 +730,7 @@ from litellm.proxy.pass_through_endpoints.openai_passthrough_endpoints import ( router as openai_passthrough_router, ) from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + drain_passthrough_upstream_error_reports, initialize_pass_through_endpoints, ) from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( @@ -1584,6 +1585,11 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: await flush_spend_counters_on_shutdown() + try: + await drain_passthrough_upstream_error_reports() + except Exception as e: # noqa: BLE001 # shutdown must continue when a report drain fails + verbose_proxy_logger.error("Error draining passthrough upstream error reports: %s", e) + await _flush_spend_logs_queue_on_shutdown() await proxy_config.stop_config_sync_subscriber() diff --git a/tests/integration/observability/test_passthrough_upstream_error_visibility.py b/tests/integration/observability/test_passthrough_upstream_error_visibility.py index 6b0ddff36ac..c302d36ad88 100644 --- a/tests/integration/observability/test_passthrough_upstream_error_visibility.py +++ b/tests/integration/observability/test_passthrough_upstream_error_visibility.py @@ -1000,6 +1000,50 @@ def test_gemini_passthrough_streaming_429_upstream_abort_after_first_frame_still assert follow_up.status_code == 200, follow_up.text +def test_upstream_abort_after_preview_budget_reaches_client_as_truncated(gateway: Gateway, tmp_path: Path) -> None: + frames: Final = tuple(b"d" * 1000 for _ in range(5)) + (b"data: tail\n\n",) + + def respond(request: Request) -> Reply: + if "streamGenerateContent" in request.target: + return Reply(status=500, content_type="text/event-stream", chunks=frames, abort_after=5) + return Reply(status=200, body=json.dumps({"ok": True}).encode()) + + path: Final = tmp_path / "gemini-stream-500-abort-past-preview.yaml" + with wire_server(respond) as wire: + _gemini_config(path, wire.url) + with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned: + candidate: Final = owned.gateway + received: Final = bytearray() + + def consume_error_stream() -> None: + with candidate.client.stream( + "POST", + _GEMINI_STREAM_PATH, + params={"alt": "sse"}, + json=_GENERATE_CONTENT, + headers=_gemini_headers(candidate), + timeout=httpx.Timeout(15, connect=5), + ) as response: + assert response.status_code == 500, response.text + for chunk in response.iter_bytes(): + received.extend(chunk) + + with pytest.raises(httpx.HTTPError): + consume_error_stream() + assert bytes(received) == b"d" * 5000, bytes(received)[-64:] + eventually( + lambda: _upstream_warnings(owned.log), + lambda lines: ( + any("returned 500" in line for line in lines) and any("read failed" in line for line in lines) + ), + seconds=30, + ) + returned: Final = tuple(line for line in _upstream_warnings(owned.log) if "returned 500" in line) + read_failures: Final = tuple(line for line in _upstream_warnings(owned.log) if "read failed" in line) + assert len(returned) == 1, returned + assert len(read_failures) == 1, read_failures + + def test_gemini_passthrough_empty_streaming_429_still_logged(gateway: Gateway, tmp_path: Path) -> None: def respond(request: Request) -> Reply: return Reply(status=429, content_type="text/event-stream", chunks=()) 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 8a2e41b68bc..2824de19484 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 @@ -27,15 +27,20 @@ from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.proxy._types import ProxyException, UserAPIKeyAuth from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( + _REPORT_TASKS, DEFAULT_PASS_THROUGH_REQUEST_TIMEOUT_SECONDS, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, HttpPassThroughEndpointHelpers, InitPassThroughEndpointHelpers, + _bounded_report, + _PreviewReportingStream, _registered_pass_through_routes, + _spawn_report_task, _truncate_upstream_error_body, _with_trace_context, chat_completion_pass_through_endpoint, create_pass_through_route, + drain_passthrough_upstream_error_reports, initialize_pass_through_endpoints, pass_through_request, resolve_llm_passthrough_timeout, @@ -4935,6 +4940,211 @@ async def test_pass_through_request_streaming_upstream_error_body_read_failure_k ), rendered +class _UpstreamErrorBodyStreamHeld(httpx.AsyncByteStream): + def __init__(self, chunks: tuple[bytes, ...], hold: asyncio.Event) -> None: + self._chunks: Final = chunks + self._hold: Final = hold + + async def __aiter__(self): + for chunk in self._chunks: + yield chunk + await self._hold.wait() + + +class _UpstreamErrorBodyStreamAbortingAfter(httpx.AsyncByteStream): + def __init__(self, chunks: tuple[bytes, ...]) -> None: + self._chunks: Final = chunks + + async def __aiter__(self): + for chunk in self._chunks: + yield chunk + raise httpx.RemoteProtocolError("peer closed connection without sending complete message body") + + +@pytest.mark.asyncio +async def test_pass_through_request_streaming_upstream_abort_after_preview_budget_reraises_to_client(): + chunks: Final = tuple(b"d" * 1000 for _ in range(5)) + upstream_response: Final = httpx.Response( + status_code=500, + headers={"content-type": "text/event-stream"}, + stream=_UpstreamErrorBodyStreamAbortingAfter(chunks), + request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"), + ) + + enqueued: list[asyncio.Future[None]] = [] + + def _recording_spawn(coro): + enqueued.append(asyncio.ensure_future(coro)) + return enqueued[-1] + + with patch( + "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( + "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) + received: list[bytes] = [] + + async def consume_response() -> None: + async for chunk in response.body_iterator: + received.append(chunk if isinstance(chunk, bytes) else chunk.encode("utf-8")) + + with pytest.raises(httpx.RemoteProtocolError): + await consume_response() + assert b"".join(received) == b"d" * 5000 + + assert len(enqueued) == 1, enqueued + await enqueued[0] + mock_proxy_logging.post_call_failure_hook.assert_called_once() + expected_body: Final = f"{'d' * 4096}... (truncated at 4096 chars)" + assert ( + mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail + == f"Upstream passthrough request failed with status 500: {expected_body}" + ) + + +@pytest.mark.asyncio +async def test_shutdown_drain_finishes_report_after_consumer_cancelled(): + upstream_hold: Final = asyncio.Event() + chunks: Final = (b"first",) + upstream_response: Final = httpx.Response( + status_code=500, + headers={"content-type": "text/event-stream"}, + stream=_UpstreamErrorBodyStreamHeld(chunks, upstream_hold), + request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"), + ) + release: Final = asyncio.Event() + hook_done: Final = asyncio.Event() + log_warning: Final = MagicMock() + + async def report(preview: bytes) -> None: + await release.wait() + hook_done.set() + + relay: Final = _PreviewReportingStream( + upstream=upstream_response, + report=report, + log_warning=log_warning, + spawn=_spawn_report_task, + ) + + async def consume() -> None: + async for _ in relay.__aiter__(): + pass + + consumer: Final = asyncio.ensure_future(consume()) + for _ in range(20): + await asyncio.sleep(0) + consumer.cancel() + with pytest.raises(asyncio.CancelledError): + await consumer + + pending: Final = tuple(_REPORT_TASKS) + assert len(pending) == 1, pending + assert not pending[0].done() + assert not hook_done.is_set() + log_warning.assert_called_once_with( + "pass_through_endpoint: client disconnected after %d preview bytes of the upstream error body", 5 + ) + + release.set() + await drain_passthrough_upstream_error_reports(timeout=5) + assert hook_done.is_set() + assert pending[0].done() + assert not _REPORT_TASKS + + +@pytest.mark.asyncio +async def test_shutdown_drain_returns_with_report_still_pending_on_timeout(): + upstream_hold: Final = asyncio.Event() + chunks: Final = (b"first",) + upstream_response: Final = httpx.Response( + status_code=500, + headers={"content-type": "text/event-stream"}, + stream=_UpstreamErrorBodyStreamHeld(chunks, upstream_hold), + request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"), + ) + release: Final = asyncio.Event() + + async def report(preview: bytes) -> None: + await release.wait() + + relay: Final = _PreviewReportingStream( + upstream=upstream_response, + report=report, + log_warning=MagicMock(), + spawn=_spawn_report_task, + ) + + async def consume() -> None: + async for _ in relay.__aiter__(): + pass + + consumer: Final = asyncio.ensure_future(consume()) + for _ in range(20): + await asyncio.sleep(0) + consumer.cancel() + with pytest.raises(asyncio.CancelledError): + await consumer + + pending: Final = tuple(_REPORT_TASKS) + assert len(pending) == 1, pending + await drain_passthrough_upstream_error_reports(timeout=0.05) + assert not pending[0].done() + pending[0].cancel() + await asyncio.gather(pending[0], return_exceptions=True) + + +@pytest.mark.asyncio +async def test_report_concurrency_is_bounded_by_the_semaphore(): + limiter: Final = asyncio.Semaphore(2) + release: Final = asyncio.Event() + entered: list[int] = [] + finished: list[int] = [] + + def make_report(index: int): + async def report() -> None: + entered.append(index) + await release.wait() + finished.append(index) + + return report + + tasks: Final = tuple(asyncio.ensure_future(_bounded_report(make_report(index)(), limiter)) for index in range(4)) + for _ in range(20): + await asyncio.sleep(0) + assert len(entered) == 2, entered + + release.set() + await asyncio.gather(*tasks) + assert len(entered) == 4, entered + assert len(finished) == 4, finished + + class _UpstreamErrorGzipStreamDropping(httpx.AsyncByteStream): def __init__(self, flushed_prefix: bytes) -> None: self._flushed_prefix: Final = flushed_prefix