From 67f785dd6cff7bd24d08a277193091dcaef9c9eb Mon Sep 17 00:00:00 2001 From: Arun Mittal Date: Tue, 9 Jun 2026 17:51:08 -0400 Subject: [PATCH] fix(proxy): bound client-supplied keepalive_seconds + drain cancellation MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Address two Greptile findings on PR #30058: Medium — Unbounded keepalive interval (DoS): An authenticated client could send ``keepalive_seconds=1e-9`` on a request whose upstream stalls, forcing ``_iter_with_keepalive`` into a tight ``: ping`` busy-loop that consumes event-loop time and response bandwidth. Clamp the request value to ``[1.0s, 300.0s]`` via the new ``_KEEPALIVE_MIN_SECONDS`` / ``_KEEPALIVE_MAX_SECONDS`` constants and log at INFO when a clamp is applied so operators can see abuse. P2 — Cancelled task not awaited: ``pending.cancel()`` only schedules a ``CancelledError`` on the task; the loop must tick before the task transitions to ``cancelled``. On an abandoned generator some uvicorn/Python combos surface this as a "Task destroyed but it is pending" warning. Await the cancellation in the ``finally`` (allowed in async generators per PEP 525) and swallow the expected propagation so the task fully drains. Tests: - test_keepalive_seconds_below_minimum_is_clamped_up - test_keepalive_seconds_above_maximum_is_clamped_down - (existing tests now patch ``_KEEPALIVE_MIN_SECONDS`` to keep sub-second execution; the production floor is 1.0s) Co-Authored-By: Claude Opus 4.7 --- litellm/proxy/proxy_server.py | 39 +++++- .../test_async_data_generator_keepalive.py | 122 ++++++++++++++++++ 2 files changed, 160 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 017d8d976f8..eda816895e5 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -7018,6 +7018,16 @@ def _format_streaming_sse_chunk(chunk: Union[str, bytes]) -> Union[str, bytes]: # Must be a unique object (not a string) since real chunks can be strings. _STREAM_KEEPALIVE = object() +# Bounds for the client-supplied ``keepalive_seconds`` request field. A +# malicious or buggy caller could otherwise pass a tiny positive value (e.g. +# ``1e-9``) which would make ``_iter_with_keepalive`` emit ``: ping`` in a +# tight loop on any stalled stream, consuming event-loop time and bandwidth +# (denial-of-service). The maximum simply caps absurdly long intervals that +# would defeat the purpose of the heartbeat. Values outside the band are +# clamped (rather than rejected) so existing callers don't break. +_KEEPALIVE_MIN_SECONDS = 1.0 +_KEEPALIVE_MAX_SECONDS = 300.0 + async def _iter_with_keepalive(aiter, keepalive_seconds: float): """Yield items from ``aiter``, optionally emitting keepalive heartbeats. @@ -7056,6 +7066,18 @@ async def _iter_with_keepalive(aiter, keepalive_seconds: float): finally: if pending is not None and not pending.done(): pending.cancel() + # Drain the cancellation so the Task fully transitions to + # ``cancelled`` before this generator returns; otherwise the loop + # may log "Task was destroyed but it is pending" warnings on early + # client disconnect. ``await`` inside an async-generator ``finally`` + # is supported (PEP 525). Swallow the propagated ``CancelledError`` + # — it is the expected outcome of the cancel — and also swallow any + # ``StopAsyncIteration`` / upstream exception raised by the wrapped + # coroutine while it unwinds; we are already in cleanup. + try: + await pending + except BaseException: + pass async def async_data_generator( # noqa: PLR0915 @@ -7106,8 +7128,23 @@ async def async_data_generator( # noqa: PLR0915 except (TypeError, ValueError): _ka_secs = 0.0 if _ka_secs > 0: + # Clamp to ``[_KEEPALIVE_MIN_SECONDS, _KEEPALIVE_MAX_SECONDS]`` + # so a hostile/buggy caller cannot busy-loop heartbeats with a + # tiny interval, and cannot disable the heartbeat semantically + # via an unreasonably long interval. + _ka_clamped = max( + _KEEPALIVE_MIN_SECONDS, min(_ka_secs, _KEEPALIVE_MAX_SECONDS) + ) + if _ka_clamped != _ka_secs: + verbose_proxy_logger.info( + "keepalive_seconds=%s clamped to %s [min=%s, max=%s]", + _ka_secs, + _ka_clamped, + _KEEPALIVE_MIN_SECONDS, + _KEEPALIVE_MAX_SECONDS, + ) stream_iterator = _iter_with_keepalive( - stream_iterator.__aiter__(), _ka_secs + stream_iterator.__aiter__(), _ka_clamped ) async for chunk in stream_iterator: diff --git a/tests/test_litellm/proxy/test_async_data_generator_keepalive.py b/tests/test_litellm/proxy/test_async_data_generator_keepalive.py index 86ab5c2d2fe..a797e790ba0 100644 --- a/tests/test_litellm/proxy/test_async_data_generator_keepalive.py +++ b/tests/test_litellm/proxy/test_async_data_generator_keepalive.py @@ -87,6 +87,10 @@ def test_keepalive_emits_ping_when_upstream_stalls(): "_fire_deferred_stream_logging", return_value=None, ), + # Lower the server-side minimum so this test can run sub-second. + # Production deployments enforce a 1.0s floor (see + # ``_KEEPALIVE_MIN_SECONDS``). + patch.object(proxy_server_module, "_KEEPALIVE_MIN_SECONDS", 0.05), ): emitted = _run( _collect( @@ -314,3 +318,121 @@ def test_iter_with_keepalive_cancels_pending_task_on_early_close(): ) _run(_run_test()) + + +def test_keepalive_seconds_below_minimum_is_clamped_up(): + """A hostile/buggy client could send ``keepalive_seconds=1e-9`` to force + ``_iter_with_keepalive`` into a tight ``: ping`` busy-loop on any stalled + stream (denial-of-service). The server must clamp such values up to + ``_KEEPALIVE_MIN_SECONDS`` so the heartbeat rate stays bounded. + + Verified by passing a tiny ``keepalive_seconds``, stalling the upstream + for less than the clamped floor, and asserting NO pings are emitted — + if the un-clamped value had been honoured, hundreds of pings would + appear during the stall window.""" + from litellm.proxy import proxy_server as proxy_server_module + + # 0.05s stall, but we'll set the floor to 0.5s. If the unclamped 1e-9 + # interval were honoured we'd see ~50,000,000 pings; if clamping works, + # we see zero (no full keepalive interval elapses before the stall ends). + upstream = _slow_chunk_stream( + chunks=["first", "second"], + stall_before_index=1, + stall_seconds=0.05, + ) + request_data = _make_request_data(keepalive_seconds=1e-9) + user_api_key_dict = MagicMock(name="user_api_key_dict") + + fake_logging = MagicMock(name="proxy_logging_obj") + fake_logging.needs_iterator_wrap.return_value = False + fake_logging.needs_per_chunk_streaming_hook.return_value = False + + with ( + patch.object(proxy_server_module, "proxy_logging_obj", fake_logging), + patch.object( + proxy_server_module, + "_get_client_requested_model_for_streaming", + return_value=None, + ), + patch.object( + proxy_server_module.ProxyLogging, + "_fire_deferred_stream_logging", + return_value=None, + ), + # Set the floor explicitly so the test is self-contained and not + # coupled to whatever the production default happens to be. + patch.object(proxy_server_module, "_KEEPALIVE_MIN_SECONDS", 0.5), + ): + emitted = _run( + _collect( + proxy_server_module.async_data_generator( + response=upstream, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) + ) + ) + + pings = [e for e in emitted if e == ": ping\n\n"] + assert pings == [], ( + f"keepalive_seconds=1e-9 must be clamped up to the server-side " + f"minimum; un-clamped, the 0.05s stall would emit a flood of pings. " + f"got {len(pings)} pings: {emitted!r}" + ) + data_lines = [e for e in emitted if isinstance(e, str) and e.startswith("data: ")] + assert data_lines == ["data: first\n\n", "data: second\n\n", "data: [DONE]\n\n"] + + +def test_keepalive_seconds_above_maximum_is_clamped_down(): + """An interval longer than ``_KEEPALIVE_MAX_SECONDS`` would defeat the + heartbeat (the intermediary proxy times out before our first ping). + Verify that the server clamps such values down — exercising the + other side of the clamp expression for branch coverage.""" + from litellm.proxy import proxy_server as proxy_server_module + + # Force the upper bound to a tiny value so the stall (0.15s) exceeds it + # and we get a measurable number of pings if clamping worked. Without + # the clamp, the request's 999999s interval would suppress every ping. + upstream = _slow_chunk_stream( + chunks=["first", "second"], + stall_before_index=1, + stall_seconds=0.15, + ) + request_data = _make_request_data(keepalive_seconds=999999.0) + user_api_key_dict = MagicMock(name="user_api_key_dict") + + fake_logging = MagicMock(name="proxy_logging_obj") + fake_logging.needs_iterator_wrap.return_value = False + fake_logging.needs_per_chunk_streaming_hook.return_value = False + + with ( + patch.object(proxy_server_module, "proxy_logging_obj", fake_logging), + patch.object( + proxy_server_module, + "_get_client_requested_model_for_streaming", + return_value=None, + ), + patch.object( + proxy_server_module.ProxyLogging, + "_fire_deferred_stream_logging", + return_value=None, + ), + patch.object(proxy_server_module, "_KEEPALIVE_MIN_SECONDS", 0.0), + patch.object(proxy_server_module, "_KEEPALIVE_MAX_SECONDS", 0.05), + ): + emitted = _run( + _collect( + proxy_server_module.async_data_generator( + response=upstream, + user_api_key_dict=user_api_key_dict, + request_data=request_data, + ) + ) + ) + + pings = [e for e in emitted if e == ": ping\n\n"] + assert len(pings) >= 1, ( + f"keepalive_seconds=999999 must be clamped down to the server-side " + f"maximum; un-clamped, the 0.15s stall would emit zero pings. " + f"got {len(pings)} pings: {emitted!r}" + )