fix(proxy): bound client-supplied keepalive_seconds + drain cancellation

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 <noreply@anthropic.com>
This commit is contained in:
Arun Mittal 2026-06-09 17:51:08 -04:00
parent 699a873200
commit 67f785dd6c
2 changed files with 160 additions and 1 deletions

View file

@ -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:

View file

@ -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}"
)