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.
This commit is contained in:
Ayush 2026-09-15 12:16:55 +05:30
parent 8ae8dbe557
commit 8543b3cafd
4 changed files with 43 additions and 11 deletions

View file

@ -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)

View file

@ -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

View file

@ -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 == []

View file

@ -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