From 391d8f0f937833b0cc2562232277b17e0ac21028 Mon Sep 17 00:00:00 2001 From: yucheng Date: Fri, 25 Sep 2026 19:24:22 +0000 Subject: [PATCH] fix(passthrough): await own report without a shared limiter, register reports per loop and drain without a timeout by default Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 9 +- .../pass_through_endpoints.py | 37 +++--- .../test_passthrough_upstream_error_chaos.py | 50 +++++++ .../test_pass_through_endpoints.py | 123 ++++++++++++++---- tests/unit/test_constants.py | 19 ++- 5 files changed, 191 insertions(+), 47 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 972de4f8f16..89996a36f96 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1529,11 +1529,10 @@ 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 = max( - 1, 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") +PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS: Final[float | None] = ( + max(0.0, float(_raw_drain)) + if (_raw_drain := os.getenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS")) + else None ) # Headers to control callbacks diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 22cc99413c5..f41a19e12b3 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -4,6 +4,7 @@ import copy import json import posixpath import traceback +import weakref from base64 import b64encode from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, Iterable, Mapping, Sequence from dataclasses import dataclass @@ -42,7 +43,6 @@ 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, @@ -944,28 +944,35 @@ class _PreviewReportingStream(httpx.AsyncByteStream): await self._upstream.aclose() -_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 +_REPORT_TASKS: Final[ # mutable-ok: in-flight report registry per loop, drained at shutdown + weakref.WeakKeyDictionary[asyncio.AbstractEventLoop, set[asyncio.Future[None]]] +] = weakref.WeakKeyDictionary() 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) + task: Final = asyncio.ensure_future(report) + registry: Final = _REPORT_TASKS.setdefault( + asyncio.get_running_loop(), + set(), # mutable-ok: per-loop task set, tasks discard themselves on completion + ) + registry.add(task) + task.add_done_callback(registry.discard) return task async def drain_passthrough_upstream_error_reports( - timeout: float = PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS, + timeout: float | None = PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS, + log_warning: Callable[..., None] = verbose_proxy_logger.warning, ) -> None: - pending: Final = tuple(_REPORT_TASKS) - if pending: - await asyncio.wait(pending, timeout=timeout) + pending: Final = tuple(_REPORT_TASKS.get(asyncio.get_running_loop(), ())) + if not pending: + return + _, still_pending = await asyncio.wait(pending, timeout=timeout) + if still_pending: + log_warning( + "pass_through_endpoint: shutdown drain timed out with %d upstream error reports still pending", + len(still_pending), + ) def _headers_without_body_framing(headers: httpx.Headers) -> httpx.Headers: diff --git a/tests/integration/observability/test_passthrough_upstream_error_chaos.py b/tests/integration/observability/test_passthrough_upstream_error_chaos.py index 7619083df84..70035030ebc 100644 --- a/tests/integration/observability/test_passthrough_upstream_error_chaos.py +++ b/tests/integration/observability/test_passthrough_upstream_error_chaos.py @@ -193,6 +193,56 @@ async def test_passthrough_disconnect_burst_logs_every_failure_once(gateway: Gat assert follow_up.status_code == 200, follow_up.text +_SLOW_FAILURE_HOOK: Final = """ +import asyncio + +from litellm.integrations.custom_logger import CustomLogger + + +class SlowFailureHook(CustomLogger): + async def async_post_call_failure_hook( + self, request_data, original_exception, user_api_key_dict, traceback_str=None + ): + await asyncio.sleep(15) + + +instance = SlowFailureHook() +""" + + +async def test_passthrough_sigterm_drains_reports_parked_on_a_slow_failure_hook( + gateway: Gateway, tmp_path: Path +) -> None: + gate: Final = threading.Event() + + def respond(request: Request) -> Reply: + if "streamGenerateContent" in request.target: + return Reply( + status=429, content_type="text/event-stream", chunks=_RATE_LIMITED_FRAMES, gate_after_first=gate + ) + 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": ["slow_hook.instance"]}) + (tmp_path / "slow_hook.py").write_text(_SLOW_FAILURE_HOOK) + path: Final = tmp_path / "chaos-sigterm-drain.yaml" + path.write_text(yaml.safe_dump(config)) + 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 + call_ids: Final = await asyncio.gather( + *(_first_frame_then_close(str(candidate.client.base_url), candidate.key) for _ in range(20)) + ) + assert len(set(call_ids)) == 20, call_ids + gate.set() + owned.process.send_signal(signal.SIGTERM) + owned.process.wait(timeout=90) + for call_id in call_ids: + _single_spend_row(call_id) + + async def _first_frame_then_close(base_url: str, key: str) -> str: async with httpx.AsyncClient(base_url=base_url, timeout=httpx.Timeout(5, connect=5), trust_env=False) as client: async with client.stream( 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 9fc09d9a8ae..7ec0491c824 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 @@ -5,6 +5,7 @@ import json import logging import os import sys +import time import zlib from collections.abc import Callable, Mapping from contextlib import ExitStack, contextmanager @@ -33,7 +34,6 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, HttpPassThroughEndpointHelpers, InitPassThroughEndpointHelpers, - _bounded_report, _PreviewReportingStream, _registered_pass_through_routes, _spawn_report_task, @@ -5066,7 +5066,7 @@ async def test_shutdown_drain_finishes_report_after_consumer_cancelled(): with pytest.raises(asyncio.CancelledError): await consumer - pending: Final = tuple(_REPORT_TASKS) + pending: Final = tuple(_REPORT_TASKS.get(asyncio.get_running_loop(), ())) assert len(pending) == 1, pending assert not pending[0].done() assert not hook_done.is_set() @@ -5078,7 +5078,7 @@ async def test_shutdown_drain_finishes_report_after_consumer_cancelled(): await drain_passthrough_upstream_error_reports(timeout=5) assert hook_done.is_set() assert pending[0].done() - assert not _REPORT_TASKS + assert not _REPORT_TASKS.get(asyncio.get_running_loop()) @pytest.mark.asyncio @@ -5114,38 +5114,115 @@ async def test_shutdown_drain_returns_with_report_still_pending_on_timeout(): with pytest.raises(asyncio.CancelledError): await consumer - pending: Final = tuple(_REPORT_TASKS) + pending: Final = tuple(_REPORT_TASKS.get(asyncio.get_running_loop(), ())) assert len(pending) == 1, pending - await drain_passthrough_upstream_error_reports(timeout=0.05) + warnings: Final = MagicMock() + await drain_passthrough_upstream_error_reports(timeout=0.05, log_warning=warnings) + warnings.assert_called_once_with( + "pass_through_endpoint: shutdown drain timed out with %d upstream error reports still pending", 1 + ) 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] = [] +async def test_parked_reports_do_not_block_the_error_stream(): + """A pile of still-running reports must not sit on the response path: a stream whose + own report finishes immediately delivers its body without waiting on the others.""" + hold: Final = asyncio.Event() - def make_report(index: int): - async def report() -> None: - entered.append(index) + async def parked() -> None: + await hold.wait() + + spawned: Final = tuple(_spawn_report_task(parked()) for _ in range(70)) + assert len(_REPORT_TASKS.get(asyncio.get_running_loop(), ())) == 70 + + upstream_response: Final = httpx.Response( + status_code=500, + headers={"content-type": "text/event-stream"}, + stream=_UpstreamErrorBodyStream(b"d" * 5000), + request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"), + ) + hook_done: Final = asyncio.Event() + + async def report(preview: bytes) -> None: + hook_done.set() + + relay: Final = _PreviewReportingStream( + upstream=upstream_response, + report=report, + log_warning=MagicMock(), + spawn=_spawn_report_task, + ) + + started: Final = time.monotonic() + received: Final = b"".join([chunk async for chunk in relay.__aiter__()]) + elapsed: Final = time.monotonic() - started + assert received == b"d" * 5000 + assert hook_done.is_set() + assert elapsed < 0.5, elapsed + + for task in spawned: + task.cancel() + await asyncio.gather(*spawned, return_exceptions=True) + + +def test_report_registry_is_scoped_to_each_event_loop(): + """Two consecutive asyncio.run calls: each loop's reports register and drain on their + own loop, so a shared module-level primitive bound to the first loop never breaks the second.""" + + async def run_once() -> None: + loop: Final = asyncio.get_running_loop() + release: Final = asyncio.Event() + + async def parked() -> None: await release.wait() - finished.append(index) - return report + spawned: Final = tuple(_spawn_report_task(parked()) for _ in range(70)) + assert len(_REPORT_TASKS.get(loop, ())) == 70 - 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 drain_passthrough_upstream_error_reports() + assert all(task.done() for task in spawned) + assert not _REPORT_TASKS.get(loop) - release.set() - await asyncio.gather(*tasks) - assert len(entered) == 4, entered - assert len(finished) == 4, finished + asyncio.run(run_once()) + asyncio.run(run_once()) + assert all(len(pending) == 0 for pending in _REPORT_TASKS.values()) + + +@pytest.mark.asyncio +async def test_shutdown_drain_waits_without_a_timeout(): + finished: Final = asyncio.Event() + + async def report() -> None: + await asyncio.sleep(0.3) + finished.set() + + task: Final = _spawn_report_task(report()) + warnings: Final = MagicMock() + await drain_passthrough_upstream_error_reports(timeout=None, log_warning=warnings) + assert task.done() + assert finished.is_set() + warnings.assert_not_called() + + +@pytest.mark.asyncio +async def test_shutdown_drain_timeout_warns_with_the_pending_count(): + hold: Final = asyncio.Event() + + async def parked() -> None: + await hold.wait() + + task: Final = _spawn_report_task(parked()) + warnings: Final = MagicMock() + await drain_passthrough_upstream_error_reports(timeout=0.05, log_warning=warnings) + warnings.assert_called_once_with( + "pass_through_endpoint: shutdown drain timed out with %d upstream error reports still pending", 1 + ) + task.cancel() + await asyncio.gather(task, return_exceptions=True) class _UpstreamErrorGzipStreamDropping(httpx.AsyncByteStream): diff --git a/tests/unit/test_constants.py b/tests/unit/test_constants.py index 85cc63e1e54..d13cc2b0e88 100644 --- a/tests/unit/test_constants.py +++ b/tests/unit/test_constants.py @@ -70,12 +70,23 @@ def _build_constant_env_var_map() -> dict[str, str]: return env_var_map -def test_passthrough_error_report_concurrency_env_zero_clamps_to_one(monkeypatch): - """A zero/negative concurrency would deadlock every report; the constant clamps to 1.""" - monkeypatch.setenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY", "0") +@pytest.mark.parametrize( + "env_value, expected", + [ + (None, None), + ("2.5", 2.5), + ("-1", 0.0), + ], +) +def test_passthrough_error_report_drain_seconds_env_parsing(monkeypatch, env_value, expected): + """Unset waits for every report; a value clamps to >= 0 seconds.""" + if env_value is None: + monkeypatch.delenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS", raising=False) + else: + monkeypatch.setenv("PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS", env_value) try: reloaded = importlib.reload(constants) - assert reloaded.PASSTHROUGH_UPSTREAM_ERROR_REPORT_CONCURRENCY == 1 + assert reloaded.PASSTHROUGH_UPSTREAM_ERROR_REPORT_DRAIN_SECONDS == expected finally: monkeypatch.undo() importlib.reload(constants)