fix(responses/mcp): keep reasoning order and caller previous_response_id on stateless follow-ups

This commit is contained in:
mateo-berri 2026-09-03 13:27:48 -07:00
parent de4c3e9006
commit 35d3478818
5 changed files with 81 additions and 41 deletions

View file

@ -352,7 +352,7 @@ async def aresponses_api_with_mcp(
follow_up_input=follow_up_input,
model=model,
all_tools=all_tools,
response_id=None if persistence_disabled else response.id,
response_id=previous_response_id if persistence_disabled else response.id,
**follow_up_call_params,
)

View file

@ -965,22 +965,9 @@ class LiteLLM_Proxy_MCP_Handler:
@staticmethod
def _is_persistence_disabled(call_params: Mapping[str, object]) -> bool:
"""Whether the caller opted out of server-side response persistence (store=false).
Zero data retention callers send store=false, so the provider never persisted the
first response and previous_response_id cannot be used to link the follow-up call.
"""
"""store=false means the provider kept nothing, so the follow-up call cannot chain on a response id."""
return call_params.get("store") is False
@staticmethod
def _extract_reasoning_items(response: ResponsesAPIResponse) -> tuple[Mapping[str, object], ...]:
"""Reasoning output items, kept whole so reasoning.encrypted_content survives replay."""
normalized: Final = tuple(
output_item if isinstance(output_item, dict) else output_item.model_dump(exclude_none=True)
for output_item in response.output
)
return tuple(item for item in normalized if item.get("type") == "reasoning")
@staticmethod
def _create_follow_up_input(
response: ResponsesAPIResponse,
@ -1002,11 +989,11 @@ class LiteLLM_Proxy_MCP_Handler:
# Add the assistant message with function calls
assistant_message_content: Final[list[object]] = []
function_calls: Final[list[dict[str, object]]] = []
turn_items: Final[list[Mapping[str, object]]] = []
for output_item in response.output:
if not isinstance(output_item, dict) and hasattr(output_item, "model_dump"):
output_item = output_item.model_dump()
output_item = output_item.model_dump(exclude_none=True)
if isinstance(output_item, dict):
if output_item.get("type") == "function_call":
@ -1016,7 +1003,7 @@ class LiteLLM_Proxy_MCP_Handler:
# Only add if we have required fields
if call_id and name:
function_calls.append(
turn_items.append(
{
"type": "function_call",
"call_id": call_id,
@ -1024,6 +1011,8 @@ class LiteLLM_Proxy_MCP_Handler:
"arguments": arguments,
}
)
elif output_item.get("type") == "reasoning" and preserve_reasoning:
turn_items.append(output_item)
elif output_item.get("type") == "message":
# Extract content from message
content = output_item.get("content", [])
@ -1044,12 +1033,7 @@ class LiteLLM_Proxy_MCP_Handler:
}
)
if preserve_reasoning:
follow_up_input.extend(LiteLLM_Proxy_MCP_Handler._extract_reasoning_items(response))
# Add function calls (these can come directly after user message for LLM)
for function_call in function_calls:
follow_up_input.append(function_call)
follow_up_input.extend(turn_items)
# Add tool results (function call outputs)
for tool_result in tool_results:
@ -1071,11 +1055,7 @@ class LiteLLM_Proxy_MCP_Handler:
response_id: str | None,
**call_params: Any,
) -> ResponsesAPIResponse | BaseResponsesAPIStreamingIterator:
"""Make follow-up response API call with tool results.
response_id is None for stateless (store=false) requests, where the whole prior
turn is replayed in follow_up_input instead of linked by previous_response_id.
"""
"""Make follow-up response API call with tool results."""
return await aresponses(
input=follow_up_input,
model=model,

View file

@ -800,8 +800,6 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
"stream": True,
}
)
if persistence_disabled:
follow_up_params.pop("previous_response_id", None)
else:
return
# Remove tool_choice to avoid forcing more tool calls

View file

@ -785,6 +785,59 @@ def test_create_follow_up_input_preserves_reasoning_when_stateless():
}
def _response_with_interleaved_reasoning_and_tool_calls() -> Any:
"""A first-turn response that reasons before each of two function calls."""
return ResponsesAPIResponse(
id="resp_first",
created_at=1234567890,
model="gpt-5",
object="response",
status="completed",
output=[
{"type": "reasoning", "id": "rs_1", "summary": [], "encrypted_content": "blob-1"},
{"type": "function_call", "id": "fc_1", "call_id": "call-1", "name": "foo", "arguments": "{}"},
{"type": "reasoning", "id": "rs_2", "summary": [], "encrypted_content": "blob-2"},
{"type": "function_call", "id": "fc_2", "call_id": "call-2", "name": "bar", "arguments": "{}"},
],
parallel_tool_calls=False,
tool_choice="auto",
tools=[],
)
def test_create_follow_up_input_keeps_each_reasoning_item_before_its_function_call():
"""
Regression test (LIT-5427): the provider pairs a replayed reasoning item with the
item that follows it, so the replay has to keep the response's output order instead
of grouping every reasoning item ahead of every function call.
"""
follow_up = LiteLLM_Proxy_MCP_Handler._create_follow_up_input(
response=_response_with_interleaved_reasoning_and_tool_calls(),
tool_results=[
{"tool_call_id": "call-1", "name": "foo", "result": "one"},
{"tool_call_id": "call-2", "name": "bar", "result": "two"},
],
original_input="hi",
preserve_reasoning=True,
)
assert [cast(dict[str, Any], item)["type"] for item in follow_up] == [
"message",
"reasoning",
"function_call",
"reasoning",
"function_call",
"function_call_output",
"function_call_output",
]
assert [cast(dict[str, Any], item).get("id") or cast(dict[str, Any], item).get("call_id") for item in follow_up[1:5]] == [
"rs_1",
"call-1",
"rs_2",
"call-2",
]
def test_create_follow_up_input_omits_reasoning_when_stateful():
"""With store=true the provider still holds the reasoning item, so don't resend it."""
follow_up = LiteLLM_Proxy_MCP_Handler._create_follow_up_input(
@ -810,17 +863,25 @@ def test_is_persistence_disabled(call_params: dict[str, Any], expected: bool):
@pytest.mark.parametrize(
"store, expected_previous_response_id",
[(False, None), (True, "resp_first")],
"store, caller_previous_response_id, expected_previous_response_id",
[
(False, None, None),
(False, "resp_caller", "resp_caller"),
(True, None, "resp_first"),
(True, "resp_caller", "resp_first"),
],
)
@pytest.mark.asyncio
async def test_mcp_follow_up_call_is_stateless_when_store_is_false(
monkeypatch: pytest.MonkeyPatch, store: bool, expected_previous_response_id: str | None
monkeypatch: pytest.MonkeyPatch,
store: bool,
caller_previous_response_id: str | None,
expected_previous_response_id: str | None,
):
"""
Regression test (LIT-5427): linking the MCP follow-up call with
previous_response_id fails for zero data retention callers, because store=false
means the first response was never persisted.
Regression test (LIT-5427): linking the MCP follow-up call to the first response's id
fails for zero data retention callers, because store=false means it was never persisted.
The caller's own previous_response_id was valid for the first call, so it stays.
"""
captured_calls: list[dict[str, Any]] = []
first_response = _response_with_reasoning_and_tool_call()
@ -857,6 +918,7 @@ async def test_mcp_follow_up_call_is_stateless_when_store_is_false(
model="gpt-5",
tools=[{"type": "mcp", "server_url": "litellm_proxy", "require_approval": "never"}],
store=store,
previous_response_id=caller_previous_response_id,
)
assert len(captured_calls) == 2

View file

@ -265,11 +265,11 @@ def _reasoning_item(encrypted_content: str):
@pytest.mark.asyncio
async def test_streaming_follow_up_is_stateless_when_store_is_false(monkeypatch):
async def test_streaming_follow_up_replays_reasoning_when_store_is_false(monkeypatch):
"""
Regression test (LIT-5427): with store=false the provider persisted nothing, so the
streaming follow-up must drop previous_response_id and replay the reasoning item
(carrying reasoning.encrypted_content) instead of pointing at a response id.
streaming follow-up must replay the reasoning item (carrying reasoning.encrypted_content).
The caller's own previous_response_id was valid for the first call and stays on the follow-up.
"""
_mock_mcp_environment(monkeypatch)
@ -300,7 +300,7 @@ async def test_streaming_follow_up_is_stateless_when_store_is_false(monkeypatch)
assert aresponses_mock.call_count == 1
follow_up_kwargs = aresponses_mock.call_args_list[0].kwargs
assert "previous_response_id" not in follow_up_kwargs
assert follow_up_kwargs["previous_response_id"] == "resp_prev"
assert _reasoning_item("gAAAAA-opaque-blob") in follow_up_kwargs["input"]