mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
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:
parent
699a873200
commit
67f785dd6c
2 changed files with 160 additions and 1 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue