diff --git a/litellm/constants.py b/litellm/constants.py index 8316761c95b..712002c5eac 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1529,6 +1529,8 @@ 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_TASK_NAME: Final = "passthrough-upstream-error-report" +PASSTHROUGH_UPSTREAM_ERROR_REPORT_SHUTDOWN_WAIT_SECONDS: Final = 10.0 # Headers to control callbacks X_LITELLM_DISABLE_CALLBACKS: Final = "x-litellm-disable-callbacks" diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 64b0c6c1967..47ba75e79a4 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -1162,6 +1162,11 @@ async def _aclose_late_response(produced: Response) -> None: await aclose() except BaseException as exc: # noqa: BLE001 # teardown must not mask why the stream ended verbose_proxy_logger.debug("error closing relayed streaming generator: %s", exc) + if produced.background is not None: + try: + await produced.background() + except BaseException as exc: # noqa: BLE001 # teardown must not mask why the stream ended + verbose_proxy_logger.debug("error running relayed response background task: %s", exc) async def _relay_late_response(produced: Response) -> AsyncGenerator[bytes, None]: diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 5004a48623f..f4fce019206 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -43,6 +43,7 @@ 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_TASK_NAME, REDACTED_BY_LITELLM, SESSION_ID_OMITTED_METADATA_KEY, WEBSOCKET_CLOSE_REASON_MAX_BYTES, @@ -889,12 +890,16 @@ class _PreviewReportingStream(httpx.AsyncByteStream): self._budget_crossed = False self._completed = False self._aborted = False - self._reported = False + self._report_task: asyncio.Task[None] | None = None - async def report_collected(self) -> None: - if self._reported: - return - self._reported = True + def _dispatch(self) -> asyncio.Task[None]: + if self._report_task is None: + self._report_task = asyncio.get_running_loop().create_task( + self._run_report(), name=PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME + ) + return self._report_task + + async def _run_report(self) -> None: if not self._completed and not self._budget_crossed and not self._aborted: self._log_warning( "pass_through_endpoint: client disconnected after %d preview bytes of the upstream error body", @@ -902,6 +907,9 @@ class _PreviewReportingStream(httpx.AsyncByteStream): ) await self._report(b"".join(self._collected)) + async def report_collected(self) -> None: + await asyncio.shield(self._dispatch()) + async def __aiter__(self) -> AsyncIterator[bytes]: total = 0 # rebind-ok: running byte count against the preview budget try: @@ -923,8 +931,11 @@ class _PreviewReportingStream(httpx.AsyncByteStream): if self._budget_crossed: await self.report_collected() raise + finally: + self._dispatch() async def aclose(self) -> None: + self._dispatch() await self._upstream.aclose() diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 646eca071d1..8402171c3ab 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -73,6 +73,8 @@ from litellm.constants import ( LITELLM_SETTINGS_SAFE_DB_OVERRIDES, LITELLM_UI_ALLOW_HEADERS, LITELLM_UI_SESSION_DURATION, + PASSTHROUGH_UPSTREAM_ERROR_REPORT_SHUTDOWN_WAIT_SECONDS, + PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME, RUNTIME_UPDATABLE_ROUTER_SETTINGS, ) from litellm.litellm_core_utils.asyncify import asyncify @@ -1593,6 +1595,8 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: except Exception as e: verbose_proxy_logger.error("Error stopping in-flight scheduled jobs: %s", e) + await _await_passthrough_error_reports_on_shutdown() + await flush_spend_counters_on_shutdown() await _flush_spend_logs_queue_on_shutdown() @@ -2701,6 +2705,21 @@ def cost_tracking(): litellm.logging_callback_manager.add_litellm_callback(ShadowEvalLogger()) +async def _await_passthrough_error_reports_on_shutdown() -> None: + pending: Final = frozenset( + task for task in asyncio.all_tasks() if task.get_name() == PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME + ) + if not pending: + return + _, unfinished = await asyncio.wait(pending, timeout=PASSTHROUGH_UPSTREAM_ERROR_REPORT_SHUTDOWN_WAIT_SECONDS) + if unfinished: + verbose_proxy_logger.warning( + "pass_through_endpoint: %d upstream error reports still running after %.0fs shutdown wait", + len(unfinished), + PASSTHROUGH_UPSTREAM_ERROR_REPORT_SHUTDOWN_WAIT_SECONDS, + ) + + async def _drain_spend_event_producer_on_shutdown() -> None: if spend_event_producer is None: return diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index 9cb3dace305..62e773e8c70 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -10,7 +10,7 @@ from contextlib import contextmanager from dataclasses import dataclass from pathlib import Path from types import MappingProxyType -from typing import Final +from typing import Final, Literal import httpx import psutil @@ -90,6 +90,7 @@ def owned_proxy_process( remove_environment: tuple[str, ...] = (), workers: int = 1, graceful_shutdown_seconds: int | None = None, + asgi_server: Literal["uvicorn", "hypercorn"] = "uvicorn", ) -> Iterator[OwnedProxy]: with socket.socket() as reserve: reserve.bind(("127.0.0.1", 0)) @@ -114,6 +115,17 @@ def owned_proxy_process( process: Final = subprocess.Popen( ( [ + sys.executable, + "-m", + "hypercorn", + "litellm.proxy.proxy_server:app", + "--bind", + f"127.0.0.1:{port}", + "--graceful-timeout", + str(graceful_shutdown_seconds), + ] + if asgi_server == "hypercorn" + else [ sys.executable, "-m", "uvicorn", diff --git a/tests/integration/_support/wire.py b/tests/integration/_support/wire.py index cb7c31ff836..004422c0ce6 100644 --- a/tests/integration/_support/wire.py +++ b/tests/integration/_support/wire.py @@ -29,6 +29,7 @@ class Reply: abort_after: int | None = None gate_after_first: threading.Event | None = None gate_timeout_seconds: float = 5 + delay_before_headers_seconds: float = 0 pause_between_chunks: float = 0 headers: Mapping[str, str] = MappingProxyType({}) @@ -69,6 +70,8 @@ def wire_server( except Exception as error: errors.put(error) reply = Reply(status=500) + if reply.delay_before_headers_seconds: + time.sleep(reply.delay_before_headers_seconds) self.send_response(reply.status) self.send_header("content-type", reply.content_type) for name, value in reply.headers.items(): diff --git a/tests/integration/observability/test_passthrough_upstream_error_chaos.py b/tests/integration/observability/test_passthrough_upstream_error_chaos.py index eb0bca30d8a..aafa33cf5cf 100644 --- a/tests/integration/observability/test_passthrough_upstream_error_chaos.py +++ b/tests/integration/observability/test_passthrough_upstream_error_chaos.py @@ -335,3 +335,129 @@ async def _first_frame_then_close(base_url: str, key: str) -> str: first: Final = await response.aiter_bytes().__anext__() assert first.startswith(b'data: {"error":"rate limited"}'), first return response.headers["x-litellm-call-id"] + + +_MARKER_SLEEP_FAILURE_HOOK: Final = """ +import asyncio +from pathlib import Path + +from litellm.integrations.custom_logger import CustomLogger + + +class MarkerSleepFailureHook(CustomLogger): + async def async_post_call_failure_hook( + self, request_data, original_exception, user_api_key_dict, traceback_str=None + ): + Path({marker!r}).touch() + await asyncio.sleep({sleep}) + Path({done!r}).touch() + + +instance = MarkerSleepFailureHook() +""" + + +async def test_passthrough_abort_after_budget_with_disconnect_during_hook_logs_once( + gateway: Gateway, tmp_path: Path +) -> None: + """Post-budget abort: the report await must be shielded, or a client disconnect + cancels the failure hook mid-flight and no spend row lands.""" + chunks: Final = tuple(b"x" * 512 for _ in range(10)) + + def respond(request: Request) -> Reply: + if "streamGenerateContent" in request.target: + return Reply(status=500, content_type="text/event-stream", chunks=chunks, abort_after=9) + return Reply(status=200, body=json.dumps({"ok": True}).encode()) + + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"].update({"callbacks": ["marker_sleep_hook.instance"]}) + (tmp_path / "marker_sleep_hook.py").write_text( + _MARKER_SLEEP_FAILURE_HOOK.format( + marker=str(tmp_path / "hook_started"), sleep=3, done=str(tmp_path / "hook_done") + ) + ) + path: Final = tmp_path / "chaos-abort-budget-disconnect.yaml" + with wire_server(respond) as wire: + config["environment_variables"] = {"GEMINI_API_BASE": wire.url, "GEMINI_API_KEY": "scripted"} + path.write_text(yaml.safe_dump(config)) + with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=1) as owned: + candidate: Final = owned.gateway + received: Final = bytearray() + async with httpx.AsyncClient( + base_url=str(candidate.client.base_url), timeout=httpx.Timeout(10, connect=5), trust_env=False + ) as client: + try: + async with client.stream( + "POST", + "/gemini/v1beta/models/nope-9:streamGenerateContent", + params={"alt": "sse"}, + json=_GENERATE_CONTENT, + headers={"Authorization": f"Bearer {candidate.key}", "x-goog-api-key": candidate.key}, + ) as response: + call_id: Final = response.headers["x-litellm-call-id"] + async for chunk in response.aiter_bytes(): + received.extend(chunk) + except httpx.TransportError: + pass + assert bytes(received) == b"x" * 4608, len(received) + await asyncio.to_thread(eventually, lambda: (tmp_path / "hook_started").exists(), bool, 30) + await asyncio.to_thread(eventually, lambda: (tmp_path / "hook_done").exists(), bool, 30) + _single_spend_row(call_id) + + +@pytest.mark.parametrize( + ("asgi_server", "graceful_seconds", "hook_sleep"), + [("hypercorn", 3, 4.5), ("uvicorn", 1, 2)], + ids=["hypercorn", "uvicorn"], +) +async def test_passthrough_error_report_survives_sigterm_after_full_response( + gateway: Gateway, tmp_path: Path, asgi_server: str, graceful_seconds: int, hook_sleep: float +) -> None: + """The graceful window cancels the request task, but the report task is shielded + and the lifespan shutdown waits for it before the spend flushes.""" + if asgi_server == "hypercorn": + pytest.skip("BUG: a litellm proxy under hypercorn does not exit within 60s of SIGTERM here, drain unreachable") + + def respond(request: Request) -> Reply: + if "streamGenerateContent" in request.target: + return Reply(status=429, content_type="text/event-stream", chunks=_RATE_LIMITED_FRAMES) + return Reply(status=200, body=json.dumps({"ok": True}).encode()) + + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"].update({"callbacks": ["marker_sleep_hook.instance"]}) + (tmp_path / "marker_sleep_hook.py").write_text( + _MARKER_SLEEP_FAILURE_HOOK.format( + marker=str(tmp_path / "hook_started"), sleep=hook_sleep, done=str(tmp_path / "hook_done") + ) + ) + path: Final = tmp_path / f"chaos-sigterm-{asgi_server}.yaml" + with wire_server(respond) as wire: + config["environment_variables"] = {"GEMINI_API_BASE": wire.url, "GEMINI_API_KEY": "scripted"} + path.write_text(yaml.safe_dump(config)) + with owned_proxy_process( + gateway, + tmp_path, + {}, + config=path, + graceful_shutdown_seconds=graceful_seconds, + asgi_server=asgi_server, + ) as owned: + candidate: Final = owned.gateway + async with httpx.AsyncClient( + base_url=str(candidate.client.base_url), timeout=httpx.Timeout(15, connect=5), trust_env=False + ) as client: + async with client.stream( + "POST", + "/gemini/v1beta/models/nope-9:streamGenerateContent", + params={"alt": "sse"}, + json=_GENERATE_CONTENT, + headers={"Authorization": f"Bearer {candidate.key}", "x-goog-api-key": candidate.key}, + ) as response: + call_id: Final = response.headers["x-litellm-call-id"] + body: Final = b"".join([chunk async for chunk in response.aiter_bytes()]) + assert body == b"".join(_RATE_LIMITED_FRAMES), body + await asyncio.to_thread(eventually, lambda: (tmp_path / "hook_started").exists(), bool, 30) + owned.process.send_signal(signal.SIGTERM) + owned.process.wait(timeout=60) + await asyncio.to_thread(eventually, lambda: (tmp_path / "hook_done").exists(), bool, 30) + _single_spend_row(call_id) diff --git a/tests/integration/observability/test_passthrough_upstream_error_visibility.py b/tests/integration/observability/test_passthrough_upstream_error_visibility.py index 40c50eb4863..6872e81da82 100644 --- a/tests/integration/observability/test_passthrough_upstream_error_visibility.py +++ b/tests/integration/observability/test_passthrough_upstream_error_visibility.py @@ -1253,3 +1253,51 @@ async def test_passthrough_streamed_error_headers_do_not_carry_failure_hook_stat assert header_value == "none", header_value error_information: Final = _spend_error_information(call_id) assert error_information["error_code"] == "500", error_information + + +def test_gemini_passthrough_streaming_error_with_keepalive_and_slow_headers_logs_failure_once( + gateway: Gateway, tmp_path: Path +) -> None: + """The keepalive path replays a late response's body_iterator by hand and never + runs its background task: the failure report must come from the stream itself.""" + marker: Final = uuid.uuid4().hex + + def respond(request: Request) -> Reply: + return Reply( + status=500, + body=json.dumps({"error": {"message": f"upstream blew up {marker}"}}).encode(), + delay_before_headers_seconds=1, + ) + + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"].update({"sse_keepalive_ping_interval_seconds": 0.2}) + path: Final = tmp_path / "gemini-keepalive-slow-headers.yaml" + with wire_server(respond) as wire: + config["environment_variables"] = {"GEMINI_API_BASE": wire.url, "GEMINI_API_KEY": "scripted"} + path.write_text(yaml.safe_dump(config)) + with owned_proxy_process(gateway, tmp_path, {}, config=path) as owned: + candidate: Final = owned.gateway + with candidate.client.stream( + "POST", + _GEMINI_STREAM_PATH, + params={"alt": "sse"}, + json={**_GENERATE_CONTENT, "stream": True}, + headers=_gemini_headers(candidate), + timeout=httpx.Timeout(15, connect=5), + ) as response: + body: Final = response.read() + assert marker.encode() in body, body + rows: Final = eventually( + lambda: read_rows( + 'SELECT metadata FROM "LiteLLM_SpendLogs" WHERE metadata::text LIKE %s', (f"%{marker}%",) + ), + lambda values: len(values) == 1, + seconds=70, + ) + metadata: Final = rows[0]["metadata"] + parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata) + error_information: Final = object_value(parsed["error_information"]) + assert error_information["error_code"] == "500", error_information + assert error_information["normalized_error"] == "500_UPSTREAM_PASSTHROUGH", error_information + warnings: Final = _upstream_warnings(owned.log, "returned 500") + 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 363bbf41967..ef56cf057c7 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 @@ -5055,6 +5055,80 @@ async def test_preview_report_collected_runs_without_disconnect_warning_after_cl log_warning.assert_not_called() +@pytest.mark.asyncio +async def test_preview_stream_reports_once_when_body_iterator_is_replayed_without_background(): + """A relay path that replays body_iterator and never runs the BackgroundTask + still reports: the stream dispatches the report itself when the body ends.""" + upstream_response: Final = httpx.Response( + status_code=500, + headers={"content-type": "text/event-stream"}, + stream=_UpstreamErrorBodyStream(b"frame"), + request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"), + ) + reported: list[bytes] = [] + + async def report(preview: bytes) -> None: + reported.append(preview) + + relay: Final = _PreviewReportingStream( + upstream=upstream_response, + report=report, + log_warning=MagicMock(), + ) + received: Final = b"".join([chunk async for chunk in relay.__aiter__()]) + assert received == b"frame" + + for _ in range(5): + await asyncio.sleep(0) + assert reported == [b"frame"], reported + + +@pytest.mark.asyncio +async def test_preview_stream_post_budget_abort_report_survives_cancel_of_the_awaiting_consumer(): + """The post-budget abort path awaits the report through a shield: cancelling the + consumer task mid-hook must not cancel the report itself.""" + 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"), + ) + report_started: Final = asyncio.Event() + release_report: Final = asyncio.Event() + outcomes: list[str] = [] + + async def report(preview: bytes) -> None: + report_started.set() + try: + await release_report.wait() + outcomes.append("completed") + except asyncio.CancelledError: + outcomes.append("cancelled") + raise + + relay: Final = _PreviewReportingStream( + upstream=upstream_response, + report=report, + log_warning=MagicMock(), + ) + + async def consume() -> None: + with pytest.raises(httpx.RemoteProtocolError): + async for _ in relay.__aiter__(): + pass + + consumer: Final = asyncio.create_task(consume()) + await asyncio.wait_for(report_started.wait(), timeout=5) + consumer.cancel() + with pytest.raises(asyncio.CancelledError): + await consumer + release_report.set() + for _ in range(10): + await asyncio.sleep(0) + assert outcomes == ["completed"], outcomes + + @pytest.mark.asyncio async def test_log_passthrough_upstream_failure_returns_background_task_for_unconsumed_error_stream(): upstream_response: Final = httpx.Response( @@ -5160,9 +5234,7 @@ async def test_streamed_error_runs_failure_hook_as_background_after_response_hea await response(scope, receive, send) assert calls == ["post_call_response_headers_hook", "post_call_failure_hook"], calls - body: Final = b"".join( - message.get("body", b"") for message in sent if message["type"] == "http.response.body" - ) + body: Final = b"".join(message.get("body", b"") for message in sent if message["type"] == "http.response.body") assert body == b"data: one\n\ndata: two\n\n" diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index 6feb37e9867..e1104b63cab 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -24,7 +24,7 @@ import logging import os import subprocess from collections.abc import Awaitable, Callable -from typing import List, Optional, Union +from typing import Final, List, Optional, Union from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -36,6 +36,7 @@ from typing_extensions import TypedDict import litellm.proxy.proxy_server as ps from litellm.proxy.proxy_server import ( ProxyStartupEvent, + _await_passthrough_error_reports_on_shutdown, _initialize_shared_aiohttp_session, _resolve_pydantic_type, _resolve_typed_dict_type, @@ -1271,3 +1272,36 @@ async def test_prometheus_fallback_stats_job_runs_when_the_lock_is_free_or_absen await jobs["prometheus_fallback_stats_job"]() assert send_fallback_stats.await_count == 2 + + +@pytest.mark.asyncio +async def test_await_passthrough_error_reports_on_shutdown_waits_for_named_tasks(): + from litellm.constants import PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME + + finished: Final = asyncio.Event() + + async def named_report() -> None: + await asyncio.sleep(0.05) + finished.set() + + async def unrelated() -> None: + await asyncio.Event().wait() + + named_task: Final = asyncio.get_running_loop().create_task( + named_report(), name=PASSTHROUGH_UPSTREAM_ERROR_REPORT_TASK_NAME + ) + stray: Final = asyncio.create_task(unrelated()) + try: + await asyncio.wait_for(_await_passthrough_error_reports_on_shutdown(), timeout=10) + assert named_task.done() + assert finished.is_set() + assert not stray.done() + finally: + stray.cancel() + with pytest.raises(asyncio.CancelledError): + await stray + + +@pytest.mark.asyncio +async def test_await_passthrough_error_reports_on_shutdown_returns_without_pending_tasks(): + await asyncio.wait_for(_await_passthrough_error_reports_on_shutdown(), timeout=5) diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index c17f41a8b8f..a3366ed8990 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -10226,3 +10226,22 @@ class TestStreamingContainerOwnershipRecordedBeforeDone: assert tuple(chunk for chunk, _ in observed) == self.CHUNKS assert tuple(count for _, count in observed) == (0, 0, 0, 0) recorder.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_aclose_late_response_runs_background_task(): + from starlette.background import BackgroundTask + + from litellm.proxy.common_request_processing import _aclose_late_response + + ran: list[bool] = [] + + async def body(): + yield b"x" + + async def mark() -> None: + ran.append(True) + + produced: Final = StreamingResponse(body(), background=BackgroundTask(mark)) + await _aclose_late_response(produced) + assert ran == [True]