diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index eda816895e5..cfa56944f23 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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: 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 a797e790ba0..e211af90264 100644 --- a/tests/test_litellm/proxy/test_async_data_generator_keepalive.py +++ b/tests/test_litellm/proxy/test_async_data_generator_keepalive.py @@ -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).