mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
8ae8dbe557
commit
8543b3cafd
4 changed files with 43 additions and 11 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 == []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue