fix(router): harden mid-stream continuation eligibility and prefill folding

Address review of the mid-stream continuation:
- keep response_format={"type": "text"} eligible; only json_object / json_schema
  and other structured formats decline, since "text" is the unconstrained default
- fold a new partial into an existing trailing assistant prefill instead of
  appending a second one, so a nested mid-stream break stays a single prefill
  turn even on providers that do not merge consecutive assistant messages
- split the deployment prefill-capability lookup into two readable steps
- cover the merge_reasoning decline branch, the text response_format pass, and
  the prefill-folding path with tests
This commit is contained in:
Ayush 2026-09-14 18:12:27 +05:30
parent 974da4c2c9
commit c75e3a69e4
3 changed files with 43 additions and 6 deletions

View file

@ -3067,7 +3067,9 @@ class Router:
return False
# Any structured-output request (response_format, or a forced tool call)
# produces a partial that cannot be resumed from an arbitrary cut point.
if request_kwargs.get("response_format") is not None:
# `{"type": "text"}` is the unconstrained default and stays eligible.
response_format: Final = request_kwargs.get("response_format")
if response_format is not None and response_format != {"type": "text"}:
return False
tool_choice: Final = request_kwargs.get("tool_choice")
if tool_choice == "required" or isinstance(tool_choice, Mapping):
@ -3087,7 +3089,16 @@ class Router:
regenerating text already delivered to the caller. The deployment filter
guarantees the target supports ``prefix: True`` (parity with
``_build_responses_continuation_input`` for the Responses-API path).
A nested mid-stream break can re-enter here with a prefill already
appended; the new partial is folded into that trailing assistant turn so
the request keeps a single prefill rather than two consecutive assistant
messages a non-merging provider would reject.
"""
if messages and messages[-1].get("role") == "assistant" and messages[-1].get("prefix"):
last: Final = messages[-1]
merged: dict[str, object] = {**last, "content": str(last.get("content") or "") + generated_content}
return [*messages[:-1], merged]
prefill: dict[str, object] = {"role": "assistant", "content": generated_content, "prefix": True}
return [*messages, prefill]

View file

@ -28,12 +28,11 @@ _STR_KEYED_DICT_ADAPTER: Final = TypeAdapter(dict[str, object])
def _deployment_supports_prefill(deployment: object) -> bool:
try:
params: Final = _STR_KEYED_DICT_ADAPTER.validate_python(
_STR_KEYED_DICT_ADAPTER.validate_python(deployment).get("litellm_params")
)
deployment_map: Final = _STR_KEYED_DICT_ADAPTER.validate_python(deployment)
litellm_params: Final = _STR_KEYED_DICT_ADAPTER.validate_python(deployment_map.get("litellm_params"))
except ValidationError:
return False
model: Final = params.get("model")
model: Final = litellm_params.get("model")
return isinstance(model, str) and bool(model) and supports_assistant_prefill(model=model)

View file

@ -2276,8 +2276,9 @@ async def test_acompletion_streaming_iterator_continues_after_content_when_eligi
({"emitted_disqualifying_content": True}, {}),
({"emitted_disqualifying_content": False}, {"response_format": {"type": "json_object"}}),
({"emitted_disqualifying_content": False}, {"tool_choice": "required"}),
({"emitted_disqualifying_content": False}, {"merge_reasoning_content_in_choices": True}),
],
ids=["tool_or_thinking_emitted", "json_mode", "forced_tool_choice"],
ids=["tool_or_thinking_emitted", "json_mode", "forced_tool_choice", "merged_reasoning"],
)
async def test_acompletion_streaming_iterator_declines_ineligible_after_content(error_kwargs, request_kwargs):
"""Flag on but the break is not continuation-safe: the router re-raises and
@ -2352,6 +2353,32 @@ def test_build_completion_continuation_input_appends_assistant_prefill():
assert built[-1] == {"role": "assistant", "content": "partial answer", "prefix": True}
def test_build_completion_continuation_input_folds_into_existing_prefill():
"""A nested break must not leave two trailing assistant turns: the new partial
folds into the prior prefill so a non-merging provider still gets one."""
once = litellm.Router._build_completion_continuation_input([{"role": "user", "content": "hi"}], "part one ")
twice = litellm.Router._build_completion_continuation_input(list(once), "part two")
assert [m["role"] for m in twice] == ["user", "assistant"]
assert twice[-1] == {"role": "assistant", "content": "part one part two", "prefix": True}
def test_mid_stream_continuation_eligible_allows_text_response_format():
"""response_format={"type": "text"} is the unconstrained default and must stay
eligible, unlike json_object / json_schema."""
from litellm.exceptions import MidStreamFallbackError
router = litellm.Router(
model_list=[{"model_name": "gpt-4", "litellm_params": {"model": "gpt-4", "api_key": "k"}}],
enable_mid_stream_fallback_continuation=True,
)
e = MidStreamFallbackError(
message="boom", model="gpt-4", llm_provider="openai", generated_content="Hello",
is_pre_first_chunk=False, emitted_disqualifying_content=False,
)
assert router._mid_stream_continuation_eligible(e=e, request_kwargs={"response_format": {"type": "text"}}) is True
assert router._mid_stream_continuation_eligible(e=e, request_kwargs={"response_format": {"type": "json_object"}}) is False
@pytest.mark.asyncio
async def test_acompletion_streaming_iterator_reraises_original_exception_when_available():
"""Async: when the mid-stream MidStreamFallbackError wraps a real provider