mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
974da4c2c9
commit
c75e3a69e4
3 changed files with 43 additions and 6 deletions
|
|
@ -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]
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue