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