feat(proxy): per-deployment keepalive_seconds default via litellm_params

Operators can now set a per-model heartbeat default in the proxy config:

  model_list:
    - model_name: claude-opus-4
      litellm_params:
        model: bedrock/claude-opus-4
        keepalive_seconds: 15

Without this, every caller has to opt in per request - even for models
that are known to stall behind L7 inference proxies (ALB/nginx with 30s
idle limits) and produce filtered-out chunks (Anthropic ``ping`` events
that the OpenAI translation layer maps to empty chunks and drops).

Resolution order in _resolve_keepalive_seconds():
  1. request_data["keepalive_seconds"]              (explicit caller wins)
  2. routed deployment's litellm_params.keepalive_seconds
  3. 0  (disabled)

An explicit request 0 disables the heartbeat even when the deployment
sets a default - the request always overrides config.

Deployment lookup prefers the exact deployment that served the request
(response._hidden_params["model_id"] -> router.get_deployment, O(1)),
falling back to router.get_model_list(model_name=...) which honours
aliases, wildcards, and team-scoped models. Plain linear scans over
router.model_list miss those.

Clamping to [_KEEPALIVE_MIN_SECONDS, _KEEPALIVE_MAX_SECONDS] is applied
uniformly, so a deployment with a too-small default cannot DoS the event
loop either.

The inline ~18-line resolve+clamp block at the async_data_generator call
site is replaced with a single call to _resolve_keepalive_seconds().

Tests (extend the existing test_async_data_generator_keepalive.py):
- request value wins over deployment default
- O(1) fallback via model_id (_hidden_params present)
- model_name fallback when _hidden_params absent
- request 0 disables even with a deployment default
- llm_router is None -> 0
- deployment default is clamped too
- end-to-end: async_data_generator emits pings from a deployment-only default
This commit is contained in:
Arun Mittal 2026-06-10 15:05:24 -04:00
parent 67f785dd6c
commit f28c16ae58
2 changed files with 285 additions and 25 deletions

View file

@ -7080,6 +7080,72 @@ async def _iter_with_keepalive(aiter, keepalive_seconds: float):
pass
def _keepalive_from_deployment_config(request_data: dict, response: Any) -> Any:
"""Look up a deployment-level ``litellm_params.keepalive_seconds`` default.
Prefers the exact deployment that served this request (by ``model_id`` from
``response._hidden_params``, an O(1) router lookup). Falls back to resolving
by ``model`` name via ``get_model_list`` (alias/wildcard/team aware) when
the response doesn't carry a ``model_id`` (e.g. some streaming response
types). Returns the raw value (possibly ``None``) for the caller to coerce.
"""
if llm_router is None:
return None
hidden = getattr(response, "_hidden_params", None)
model_id = hidden.get("model_id") if isinstance(hidden, dict) else None
if model_id:
deployment = llm_router.get_deployment(model_id=model_id)
if deployment is not None:
# ``litellm_params`` is a pydantic model with ``extra="allow"``, so
# a custom field like ``keepalive_seconds`` is reached via getattr,
# not ``.get()``.
return getattr(deployment.litellm_params, "keepalive_seconds", None)
for deployment_dict in (
llm_router.get_model_list(model_name=request_data.get("model")) or []
):
raw = (deployment_dict.get("litellm_params") or {}).get("keepalive_seconds")
if raw is not None:
return raw
return None
def _resolve_keepalive_seconds(request_data: dict, response: Any = None) -> float:
"""Resolve the SSE keepalive interval for a streaming request.
Resolution order:
1. Explicit ``request_data["keepalive_seconds"]`` (if set).
2. The routed deployment's ``litellm_params.keepalive_seconds`` default.
3. ``0`` (disabled).
An explicit request ``0`` disables the heartbeat even when the deployment
sets a default — the request always overrides config. When enabled
(``> 0``) the result is clamped to
``[_KEEPALIVE_MIN_SECONDS, _KEEPALIVE_MAX_SECONDS]``; values outside the
band are clamped (not rejected) so existing callers don't break.
"""
raw = request_data.get("keepalive_seconds")
if raw is None:
raw = _keepalive_from_deployment_config(request_data, response)
try:
value = float(raw or 0)
except (TypeError, ValueError):
return 0.0
if value <= 0:
return 0.0 # disabled — never clamp up to the minimum
clamped = max(_KEEPALIVE_MIN_SECONDS, min(value, _KEEPALIVE_MAX_SECONDS))
if clamped != value:
verbose_proxy_logger.info(
"keepalive_seconds=%s clamped to %s [min=%s, max=%s]",
value,
clamped,
_KEEPALIVE_MIN_SECONDS,
_KEEPALIVE_MAX_SECONDS,
)
return clamped
async def async_data_generator( # noqa: PLR0915
response, user_api_key_dict: UserAPIKeyAuth, request_data: dict
):
@ -7115,36 +7181,19 @@ async def async_data_generator( # noqa: PLR0915
else:
stream_iterator = response
# Optional client-controlled SSE keepalive: when ``keepalive_seconds``
# > 0 is set on the request, emit an SSE comment (``: ping``) if no
# upstream chunk arrives within that interval. Useful when an
# Optional SSE keepalive: emit an SSE comment (``: ping``) if no
# upstream chunk arrives within the resolved interval. Useful when an
# intermediary proxy (e.g. an L7 inference proxy, ALB, nginx) cuts
# idle streams while the model is generating but producing
# filtered-out chunks (e.g. Anthropic ``ping`` events that the OpenAI
# translation layer maps to empty chunks and then drops). Absent or
# 0 -> no keepalive, no behaviour change vs. the upstream fast path.
try:
_ka_secs = float(request_data.get("keepalive_seconds") or 0)
except (TypeError, ValueError):
_ka_secs = 0.0
# translation layer maps to empty chunks and then drops). Resolves
# from the request (``keepalive_seconds``) first, then the routed
# deployment's ``litellm_params.keepalive_seconds``, otherwise 0
# (disabled) — see ``_resolve_keepalive_seconds`` for clamping rules.
_ka_secs = _resolve_keepalive_seconds(request_data, response)
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_clamped
stream_iterator.__aiter__(), _ka_secs
)
async for chunk in stream_iterator:

View file

@ -383,6 +383,217 @@ def test_keepalive_seconds_below_minimum_is_clamped_up():
assert data_lines == ["data: first\n\n", "data: second\n\n", "data: [DONE]\n\n"]
def _make_deployment_obj(keepalive_seconds):
"""Stand-in for ``router.get_deployment(model_id=...)`` return value: an
object with a ``litellm_params`` attribute that itself exposes the custom
``keepalive_seconds`` field via ``getattr`` (matching the real pydantic
``extra="allow"`` shape)."""
params = MagicMock(name="litellm_params")
# ``getattr(params, "keepalive_seconds", None)`` on a MagicMock returns
# another MagicMock by default — force the attribute explicitly so the
# ``getattr`` lookup returns the value we want (incl. None).
params.keepalive_seconds = keepalive_seconds
deployment = MagicMock(name="deployment")
deployment.litellm_params = params
return deployment
def test_resolve_keepalive_request_value_wins_over_deployment_default():
"""When the request supplies ``keepalive_seconds``, the deployment-level
default is never consulted — request always overrides config."""
from litellm.proxy import proxy_server as proxy_server_module
fake_router = MagicMock(name="llm_router")
fake_router.get_deployment.return_value = _make_deployment_obj(60.0)
fake_router.get_model_list.return_value = [
{"litellm_params": {"keepalive_seconds": 60.0}}
]
with patch.object(proxy_server_module, "llm_router", fake_router):
# Request value (5.0) wins over both deployment paths (60.0).
resolved = proxy_server_module._resolve_keepalive_seconds(
request_data={"model": "gpt-3.5-turbo", "keepalive_seconds": 5.0},
response=MagicMock(_hidden_params={"model_id": "dep-id-1"}),
)
assert resolved == 5.0
# Deployment lookup must not have been consulted at all.
fake_router.get_deployment.assert_not_called()
fake_router.get_model_list.assert_not_called()
def test_resolve_keepalive_falls_back_to_deployment_via_model_id():
"""When the request has no ``keepalive_seconds`` and the response carries a
``model_id`` in ``_hidden_params``, the resolver looks up that exact
deployment via ``router.get_deployment(model_id=...)`` (the O(1) path)."""
from litellm.proxy import proxy_server as proxy_server_module
fake_router = MagicMock(name="llm_router")
fake_router.get_deployment.return_value = _make_deployment_obj(42.0)
# If this is touched we know the fallback path was taken when it
# shouldn't have been.
fake_router.get_model_list.return_value = [
{"litellm_params": {"keepalive_seconds": 999.0}}
]
with patch.object(proxy_server_module, "llm_router", fake_router):
resolved = proxy_server_module._resolve_keepalive_seconds(
request_data={"model": "gpt-3.5-turbo"},
response=MagicMock(_hidden_params={"model_id": "dep-id-1"}),
)
assert resolved == 42.0
fake_router.get_deployment.assert_called_once_with(model_id="dep-id-1")
# The model-name fallback path is only used when ``get_deployment`` fails
# to resolve.
fake_router.get_model_list.assert_not_called()
def test_resolve_keepalive_falls_back_to_model_name_when_no_model_id():
"""When the response has no ``_hidden_params["model_id"]`` (some streaming
response types don't carry it), the resolver falls back to
``router.get_model_list(model_name=...)`` — alias/wildcard/team aware."""
from litellm.proxy import proxy_server as proxy_server_module
fake_router = MagicMock(name="llm_router")
fake_router.get_model_list.return_value = [
{"litellm_params": {"keepalive_seconds": 30.0}}
]
with patch.object(proxy_server_module, "llm_router", fake_router):
resolved = proxy_server_module._resolve_keepalive_seconds(
request_data={"model": "gpt-3.5-turbo"},
response=MagicMock(spec=[]), # no _hidden_params attribute
)
assert resolved == 30.0
fake_router.get_model_list.assert_called_once_with(model_name="gpt-3.5-turbo")
def test_resolve_keepalive_zero_in_request_disables_even_with_deployment_default():
"""An explicit ``keepalive_seconds=0`` in the request disables the heartbeat
even when the deployment has a non-zero default — the request always wins,
including for the "disable" case."""
from litellm.proxy import proxy_server as proxy_server_module
fake_router = MagicMock(name="llm_router")
fake_router.get_deployment.return_value = _make_deployment_obj(60.0)
with patch.object(proxy_server_module, "llm_router", fake_router):
resolved = proxy_server_module._resolve_keepalive_seconds(
request_data={"model": "gpt-3.5-turbo", "keepalive_seconds": 0},
response=MagicMock(_hidden_params={"model_id": "dep-id-1"}),
)
assert resolved == 0.0
# Deployment must not be consulted — request 0 short-circuits.
fake_router.get_deployment.assert_not_called()
def test_resolve_keepalive_returns_zero_when_router_is_none():
"""No router (e.g. proxy started without a config / model_list) — no
fallback is possible, resolver returns 0 (disabled)."""
from litellm.proxy import proxy_server as proxy_server_module
with patch.object(proxy_server_module, "llm_router", None):
resolved = proxy_server_module._resolve_keepalive_seconds(
request_data={"model": "gpt-3.5-turbo"},
response=MagicMock(),
)
assert resolved == 0.0
def test_resolve_keepalive_clamps_deployment_default_too():
"""The clamp must apply regardless of where the value came from — a
deployment config with an out-of-band default still gets clamped."""
from litellm.proxy import proxy_server as proxy_server_module
fake_router = MagicMock(name="llm_router")
fake_router.get_deployment.return_value = _make_deployment_obj(999999.0)
with (
patch.object(proxy_server_module, "llm_router", fake_router),
patch.object(proxy_server_module, "_KEEPALIVE_MAX_SECONDS", 60.0),
):
resolved = proxy_server_module._resolve_keepalive_seconds(
request_data={"model": "gpt-3.5-turbo"},
response=MagicMock(_hidden_params={"model_id": "dep-id-1"}),
)
assert resolved == 60.0
def test_async_data_generator_uses_deployment_config_keepalive():
"""End-to-end: with no ``keepalive_seconds`` on the request but a
deployment-level default, ``async_data_generator`` emits ``: ping``
heartbeats when the upstream stalls. Proves the resolver is wired
into the streaming generator and not just unit-callable."""
from litellm.proxy import proxy_server as proxy_server_module
inner = _slow_chunk_stream(
chunks=["first", "second"],
stall_before_index=1,
stall_seconds=0.3,
)
# Wrap the async generator in a class that also carries ``_hidden_params``
# so the resolver's O(1) path (``router.get_deployment(model_id=...)``)
# is exercised — raw async generators don't allow attribute assignment.
class _UpstreamWithHiddenParams:
def __init__(self, gen, hidden_params):
self._gen = gen
self._hidden_params = hidden_params
def __aiter__(self):
return self._gen.__aiter__()
upstream = _UpstreamWithHiddenParams(inner, {"model_id": "dep-id-1"})
fake_router = MagicMock(name="llm_router")
fake_router.get_deployment.return_value = _make_deployment_obj(0.1)
request_data = _make_request_data(keepalive_seconds=None)
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, "llm_router", fake_router),
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,
),
# Lower the server-side minimum so this test can run sub-second
# (same trick as ``test_keepalive_emits_ping_when_upstream_stalls``).
patch.object(proxy_server_module, "_KEEPALIVE_MIN_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) >= 2, (
f"deployment-level keepalive_seconds=0.1 should trigger heartbeats "
f"during the 0.3s upstream stall; got {len(pings)} pings. full: {emitted!r}"
)
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).