diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index c81d0ed5d6b..9078b7c134f 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, Iterable, Mapping, Sequence +from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Coroutine, Iterable, Mapping, Sequence from dataclasses import dataclass from datetime import datetime from itertools import count, groupby @@ -875,12 +875,15 @@ class _PreviewReportingStream(httpx.AsyncByteStream): upstream: httpx.Response, report: _ReportPreview, log_warning: Callable[..., None], + enqueue: Callable[[Coroutine[None, None, None]], None], ) -> None: self._upstream: Final = upstream self._report: Final = report 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 async def _report_once(self) -> None: if self._reported: @@ -888,6 +891,12 @@ class _PreviewReportingStream(httpx.AsyncByteStream): 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()) + async def __aiter__(self) -> AsyncIterator[bytes]: total = 0 # rebind-ok: running byte count against the preview budget try: @@ -898,17 +907,19 @@ class _PreviewReportingStream(httpx.AsyncByteStream): if total > PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS: await self._report_once() yield chunk + await self._report_once() 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__, ) - finally: await self._report_once() + finally: + self._enqueue_pending_report() async def aclose(self) -> None: - await self._report_once() + self._enqueue_pending_report() await self._upstream.aclose() @@ -990,7 +1001,12 @@ async def _log_passthrough_upstream_failure( return httpx.Response( status_code=response.status_code, headers=_headers_without_body_framing(response.headers), - stream=_PreviewReportingStream(upstream=response, report=report, log_warning=log_warning), + stream=_PreviewReportingStream( + upstream=response, + report=report, + log_warning=log_warning, + enqueue=GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue, + ), request=response.request, extensions=response.extensions, ) diff --git a/tests/integration/observability/test_passthrough_upstream_error_visibility.py b/tests/integration/observability/test_passthrough_upstream_error_visibility.py index 103474ba985..389aeddb81d 100644 --- a/tests/integration/observability/test_passthrough_upstream_error_visibility.py +++ b/tests/integration/observability/test_passthrough_upstream_error_visibility.py @@ -100,6 +100,18 @@ def _spend_error_information(call_id: str) -> dict[str, JsonValue]: return object_value(parsed["error_information"]) +def _spend_error_information_or_none(call_id: str, seconds: float = 20) -> dict[str, JsonValue] | None: + rows: Final = eventually( + lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)), + lambda values: len(values) == 1, + seconds=seconds, + ) + metadata: Final = rows[0]["metadata"] + parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata) + error_information: Final = parsed.get("error_information") + return None if error_information is None else object_value(error_information) + + def _spend_status(call_id: str) -> str: rows: Final = eventually( lambda: read_rows('SELECT status FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)), @@ -671,3 +683,40 @@ def test_gemini_passthrough_quota_wording_in_upstream_body_keeps_passthrough_nor warning: Final = _upstream_warning(owned.log) assert "exceeded your current quota" in warning, warning assert _LEAKED_UPSTREAM_KEY not in warning, warning + + +def test_gemini_passthrough_streaming_429_client_disconnect_still_logs_failure( + gateway: Gateway, tmp_path: Path +) -> None: + gate: Final = threading.Event() + frames: Final = (b'data: {"error":"rate limited"}\n\n', b"data: [DONE]\n\n") + + def respond(request: Request) -> Reply: + return Reply(status=429, content_type="text/event-stream", chunks=frames, gate_after_first=gate) + + path: Final = tmp_path / "gemini-stream-429-disconnect.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 + try: + with candidate.client.stream( + "POST", + _GEMINI_STREAM_PATH, + params={"alt": "sse"}, + json=_GENERATE_CONTENT, + headers=_gemini_headers(candidate), + timeout=httpx.Timeout(3, connect=5), + ) as response: + assert response.status_code == 429, response.text + first: Final = next(response.iter_bytes()) + assert first.startswith(b'data: {"error":"rate limited"}'), first + call_id: Final = response.headers["x-litellm-call-id"] + finally: + gate.set() + error_information: Final = _spend_error_information_or_none(call_id) + assert error_information is not None, f"no spend row for {call_id} after client disconnect" + assert error_information["error_code"] == "429", error_information + assert error_information["normalized_error"] == "500_UPSTREAM_PASSTHROUGH", error_information + warnings: Final = _upstream_warnings(owned.log, "returned 429") + assert len(warnings) == 1, warnings 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 cbf6f486dd4..bf62e8cfddd 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 @@ -1,11 +1,12 @@ import asyncio +import gc import gzip import json import logging import os import sys import zlib -from collections.abc import Callable +from collections.abc import Callable, Coroutine from contextlib import ExitStack, contextmanager from io import BytesIO from types import SimpleNamespace @@ -4643,6 +4644,74 @@ async def test_pass_through_request_streaming_upstream_error_relays_first_chunk_ ), detail +@pytest.mark.asyncio +async def test_pass_through_request_streaming_upstream_error_client_disconnect_enqueues_failure_report(): + """ + Regression: when the client disconnects mid-relay the response task is + cancelled, so the preview report cannot be awaited inline; it must be + handed to the logging worker, which then fires the failure hook once + with the chunks already relayed. + """ + first_chunk: Final = b'data: {"error":"rate limited"}\n\n' + second_chunk: Final = b"data: [DONE]\n\n" + body_stream: Final = _GatedUpstreamErrorBodyStream(first_chunk, second_chunk) + upstream_response: Final = httpx.Response( + status_code=429, + headers={"content-type": "text/event-stream"}, + stream=body_stream, + 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" + ) 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() + 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) + iterator = response.body_iterator.__aiter__() + first: Final = await asyncio.wait_for(iterator.__anext__(), timeout=5) + assert first_chunk in (first if isinstance(first, bytes) else first.encode()) + await iterator.aclose() + del iterator + gc.collect() + for _ in range(40): + if enqueued: + break + await asyncio.sleep(0.05) + + 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 detail == 'Upstream passthrough request failed with status 429: data: {"error":"rate limited"}', detail + + class _UpstreamErrorBodyStreamDropping(httpx.AsyncByteStream): async def __aiter__(self): yield b'{"error": "half'