From 8543b3cafd8dfca15c4d2e1882294b03a129bb1a Mon Sep 17 00:00:00 2001 From: Ayush Date: Tue, 15 Sep 2026 12:16:55 +0530 Subject: [PATCH] fix(router): make the mid-stream continuation marker unforgeable The deployment filter is process-global and read the continuation marker as a truthy value from request_kwargs. Since the proxy can forward arbitrary request-body fields into the router, a client could send `"_mid_stream_continuation": true` on a normal request to filter a mixed model group down to prefill-capable deployments and steer their prompt onto that provider. Set a private sentinel object internally and check it by type instead of truthiness: a JSON request body cannot construct one, so a forged flag is ignored. Adds a regression test that forged client values (true, "true", 1, a dict) leave the deployment list untouched. --- litellm/router.py | 3 ++- .../continuation_prefill_check.py | 25 +++++++++++++------ .../test_continuation_prefill_check.py | 23 +++++++++++++++-- tests/test_litellm/test_router.py | 3 ++- 4 files changed, 43 insertions(+), 11 deletions(-) diff --git a/litellm/router.py b/litellm/router.py index 226b5da7a83..6a43a6a05c0 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -2854,12 +2854,13 @@ class Router: if continue_after_content: from litellm.router_utils.pre_call_checks.continuation_prefill_check import ( MID_STREAM_CONTINUATION_KWARG, + MID_STREAM_CONTINUATION_MARKER, ) initial_kwargs["messages"] = self._build_completion_continuation_input( messages, e.generated_content ) - initial_kwargs[MID_STREAM_CONTINUATION_KWARG] = True + initial_kwargs[MID_STREAM_CONTINUATION_KWARG] = MID_STREAM_CONTINUATION_MARKER else: initial_kwargs["messages"] = messages self._update_kwargs_before_fallbacks(model=model_group, kwargs=initial_kwargs) diff --git a/litellm/router_utils/pre_call_checks/continuation_prefill_check.py b/litellm/router_utils/pre_call_checks/continuation_prefill_check.py index a4e05672879..f9750b4933b 100644 --- a/litellm/router_utils/pre_call_checks/continuation_prefill_check.py +++ b/litellm/router_utils/pre_call_checks/continuation_prefill_check.py @@ -1,9 +1,9 @@ """ Mid-stream fallback continuation: keep the fallback on a deployment that can -continue a prefilled assistant message. When a request carries -``MID_STREAM_CONTINUATION_KWARG``, deployments whose model does not support -assistant prefill are dropped, so the partial text is continued rather than -regenerated or rejected. Requests without the marker pass through untouched. +continue a prefilled assistant message. When the router marks a fallback +re-entry as a continuation, deployments whose model does not support assistant +prefill are dropped, so the partial text is continued rather than regenerated or +rejected. Requests without the marker pass through untouched. """ from collections.abc import Mapping, Sequence @@ -15,10 +15,20 @@ from litellm.integrations.custom_logger import CustomLogger, Span from litellm.types.llms.openai import AllMessageValues from litellm.utils import supports_assistant_prefill -# Marks a fallback re-entry as a mid-stream continuation. Router sets it (via a -# lazy import) and this filter reads it; kept here to avoid a module-level cycle. +# The router marks a continuation re-entry by placing MID_STREAM_CONTINUATION_MARKER +# under this key. Since the proxy can forward arbitrary request-body fields into the +# router, the marker is a private object checked by type rather than a truthy value: +# a JSON request body cannot construct one, so a client cannot forge the flag to steer +# deployment selection toward prefill-capable deployments. MID_STREAM_CONTINUATION_KWARG: Final = "_mid_stream_continuation" + +class _ContinuationMarker: + """Unforgeable sentinel; only the router can produce an instance.""" + + +MID_STREAM_CONTINUATION_MARKER: Final = _ContinuationMarker() + _STR_KEYED_DICT_ADAPTER: Final = TypeAdapter(dict[str, object]) @@ -41,7 +51,8 @@ class ContinuationPrefillDeploymentCheck(CustomLogger): request_kwargs: Mapping[str, object] | None = None, parent_otel_span: Span | None = None, ) -> list[dict[str, object]]: # mutable-ok: returns a mutable deployment list - if not (request_kwargs or {}).get(MID_STREAM_CONTINUATION_KWARG): + marker: Final = (request_kwargs or {}).get(MID_STREAM_CONTINUATION_KWARG) + if not isinstance(marker, _ContinuationMarker): return healthy_deployments eligible: Final = (deployment for deployment in healthy_deployments if _deployment_supports_prefill(deployment)) return list(eligible) # mutable-ok: downstream deployment selection consumes a mutable list diff --git a/tests/test_litellm/router_utils/pre_call_checks/test_continuation_prefill_check.py b/tests/test_litellm/router_utils/pre_call_checks/test_continuation_prefill_check.py index 28f220100fa..2e1c41bf6ab 100644 --- a/tests/test_litellm/router_utils/pre_call_checks/test_continuation_prefill_check.py +++ b/tests/test_litellm/router_utils/pre_call_checks/test_continuation_prefill_check.py @@ -2,6 +2,7 @@ import pytest from litellm.router_utils.pre_call_checks.continuation_prefill_check import ( MID_STREAM_CONTINUATION_KWARG, + MID_STREAM_CONTINUATION_MARKER, ContinuationPrefillDeploymentCheck, _deployment_supports_prefill, ) @@ -39,6 +40,24 @@ async def test_filter_is_noop_without_continuation_marker(): assert result == deployments +@pytest.mark.asyncio +async def test_filter_ignores_forged_client_flag(): + """A client cannot steer routing: a plain truthy value under the marker key + (which the proxy could forward from the request body) is not the internal + sentinel, so the filter leaves the deployment list untouched.""" + check = ContinuationPrefillDeploymentCheck() + deployments = [_deployment(PREFILL_MODEL, "a"), _deployment(NON_PREFILL_MODEL, "b")] + + for forged in (True, "true", 1, {"any": "json"}): + result = await check.async_filter_deployments( + model="group", + healthy_deployments=deployments, + messages=None, + request_kwargs={MID_STREAM_CONTINUATION_KWARG: forged}, + ) + assert result == deployments + + @pytest.mark.asyncio async def test_filter_keeps_only_prefill_capable_on_continuation(): check = ContinuationPrefillDeploymentCheck() @@ -48,7 +67,7 @@ async def test_filter_keeps_only_prefill_capable_on_continuation(): model="group", healthy_deployments=deployments, messages=None, - request_kwargs={MID_STREAM_CONTINUATION_KWARG: True}, + request_kwargs={MID_STREAM_CONTINUATION_KWARG: MID_STREAM_CONTINUATION_MARKER}, ) assert [d["model_info"]["id"] for d in result] == ["a"] @@ -64,6 +83,6 @@ async def test_filter_empties_group_when_no_prefill_capable_deployment(): model="group", healthy_deployments=deployments, messages=None, - request_kwargs={MID_STREAM_CONTINUATION_KWARG: True}, + request_kwargs={MID_STREAM_CONTINUATION_KWARG: MID_STREAM_CONTINUATION_MARKER}, ) assert result == [] diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 0c3bf43182d..c37a23df7d1 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -2229,6 +2229,7 @@ async def test_acompletion_streaming_iterator_continues_after_content_when_eligi from litellm.exceptions import MidStreamFallbackError from litellm.router_utils.pre_call_checks.continuation_prefill_check import ( MID_STREAM_CONTINUATION_KWARG, + MID_STREAM_CONTINUATION_MARKER, ) router = litellm.Router( @@ -2264,7 +2265,7 @@ async def test_acompletion_streaming_iterator_continues_after_content_when_eligi mock_fallback.assert_awaited_once() passed_kwargs = mock_fallback.await_args.kwargs["kwargs"] - assert passed_kwargs[MID_STREAM_CONTINUATION_KWARG] is True + assert passed_kwargs[MID_STREAM_CONTINUATION_KWARG] is MID_STREAM_CONTINUATION_MARKER assert passed_kwargs["messages"][-1] == {"role": "assistant", "content": "Hello", "prefix": True} assert fallback_chunk in collected