Merge pull request #36575 from BerriAI/litellm_mcp_stateless_follow_up_zdr

fix(responses/mcp): keep follow-up calls stateless when store=false
This commit is contained in:
Mateo Wang 2026-09-04 17:33:27 -07:00 • committed by GitHub
commit 6960a42008
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 313 additions and 9 deletions

View file

@ -327,8 +327,13 @@ async def aresponses_api_with_mcp(
)
if tool_results:
persistence_disabled: Final = LiteLLM_Proxy_MCP_Handler._is_persistence_disabled(call_params)
follow_up_input: Final = LiteLLM_Proxy_MCP_Handler._create_follow_up_input(
response=response, tool_results=tool_results, original_input=input
response=response,
tool_results=tool_results,
original_input=input,
preserve_reasoning=persistence_disabled,
)
# Prepare parameters for follow-up call (restores original stream setting)
@ -347,7 +352,7 @@ async def aresponses_api_with_mcp(
follow_up_input=follow_up_input,
model=model,
all_tools=all_tools,
response_id=response.id,
response_id=previous_response_id if persistence_disabled else response.id,
**follow_up_call_params,
)

View file

@ -963,11 +963,17 @@ class LiteLLM_Proxy_MCP_Handler:
return follow_up_messages
@staticmethod
def _is_persistence_disabled(call_params: Mapping[str, object]) -> bool:
"""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 _create_follow_up_input(
response: ResponsesAPIResponse,
tool_results: Sequence[Mapping[str, object]],
original_input: str | ResponseInputParam | None = None,
preserve_reasoning: bool = False,
) -> list[object]:
"""Create follow-up input with tool results in proper format."""
follow_up_input: Final[list[object]] = []
@ -983,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":
@ -997,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,
@ -1005,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", [])
@ -1025,9 +1033,7 @@ class LiteLLM_Proxy_MCP_Handler:
}
)
# 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:
@ -1046,7 +1052,7 @@ class LiteLLM_Proxy_MCP_Handler:
follow_up_input: list[Any],
model: str,
all_tools: Sequence[ResponsesToolParam] | None,
response_id: str,
response_id: str | None,
**call_params: Any,
) -> ResponsesAPIResponse | BaseResponsesAPIStreamingIterator:
"""Make follow-up response API call with tool results."""

View file

@ -781,10 +781,15 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator):
try:
# Create follow-up input
if self.collected_response is not None:
persistence_disabled: Final = LiteLLM_Proxy_MCP_Handler._is_persistence_disabled(
self.original_request_params
)
follow_up_input: Final = LiteLLM_Proxy_MCP_Handler._create_follow_up_input(
response=self.collected_response,
tool_results=self.tool_results,
original_input=self.original_request_params.get("input"),
preserve_reasoning=persistence_disabled,
)
# Make follow-up call with streaming

View file

@ -9,10 +9,13 @@ from fastapi import HTTPException
import importlib
from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing
from litellm.responses import main as responses_main
from litellm.responses.mcp import litellm_proxy_mcp_handler as mcp_handler_module
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
LiteLLM_Proxy_MCP_Handler,
)
from typing import Any, cast
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.utils import ModelResponse
from litellm.types.responses.main import OutputFunctionToolCall
@ -719,3 +722,210 @@ def test_extract_tool_call_details_still_prefers_openai_arguments():
assert name == "get_weather"
assert call_id == "call_123"
assert arguments == '{"city": "Paris"}'
def _response_with_reasoning_and_tool_call() -> Any:
"""A first-turn response as a reasoning model returns it: reasoning item, then a function call."""
return ResponsesAPIResponse(
id="resp_first",
created_at=1234567890,
model="gpt-5",
object="response",
status="completed",
output=[
{
"type": "reasoning",
"id": "rs_1",
"summary": [],
"encrypted_content": "gAAAAA-opaque-blob",
},
{
"type": "function_call",
"id": "fc_1",
"call_id": "call-1",
"name": "foo",
"arguments": "{}",
"status": "completed",
},
],
parallel_tool_calls=False,
tool_choice="auto",
tools=[],
)
def test_create_follow_up_input_preserves_reasoning_when_stateless():
"""
Regression test (LIT-5427): a store=false follow-up has to replay the reasoning
item, including reasoning.encrypted_content, since the provider kept no state.
"""
follow_up = LiteLLM_Proxy_MCP_Handler._create_follow_up_input(
response=_response_with_reasoning_and_tool_call(),
tool_results=[{"tool_call_id": "call-1", "name": "foo", "result": "done"}],
original_input="hi",
preserve_reasoning=True,
)
assert follow_up[1] == {
"type": "reasoning",
"id": "rs_1",
"summary": [],
"encrypted_content": "gAAAAA-opaque-blob",
}
assert follow_up[2] == {
"type": "function_call",
"call_id": "call-1",
"name": "foo",
"arguments": "{}",
}
assert follow_up[3] == {
"type": "function_call_output",
"call_id": "call-1",
"output": "done",
}
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(
response=_response_with_reasoning_and_tool_call(),
tool_results=[{"tool_call_id": "call-1", "name": "foo", "result": "done"}],
original_input="hi",
)
assert not [item for item in follow_up if isinstance(item, dict) and item.get("type") == "reasoning"]
@pytest.mark.parametrize(
"call_params, expected",
[
({"store": False}, True),
({"store": True}, False),
({"store": None}, False),
({}, False),
],
)
def test_is_persistence_disabled(call_params: dict[str, Any], expected: bool):
assert LiteLLM_Proxy_MCP_Handler._is_persistence_disabled(call_params) is expected
@pytest.mark.parametrize(
"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,
caller_previous_response_id: str | None,
expected_previous_response_id: str | None,
):
"""
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()
async def fake_aresponses(**kwargs: Any) -> ResponsesAPIResponse:
captured_calls.append(kwargs)
return first_response if len(captured_calls) == 1 else ResponsesAPIResponse(
id="resp_follow_up",
created_at=1234567891,
model="gpt-5",
object="response",
status="completed",
output=[],
parallel_tool_calls=False,
tool_choice="auto",
tools=[],
)
async def fake_process(**kwargs: Any) -> tuple[list[Any], dict[str, str]]:
return ([], {"foo": "litellm_proxy"})
async def fake_execute(**kwargs: Any) -> list[dict[str, Any]]:
return [{"tool_call_id": "call-1", "name": "foo", "result": "done"}]
monkeypatch.setattr(responses_main, "aresponses", fake_aresponses)
monkeypatch.setattr(mcp_handler_module, "aresponses", fake_aresponses)
monkeypatch.setattr(
LiteLLM_Proxy_MCP_Handler, "_process_mcp_tools_without_openai_transform", staticmethod(fake_process)
)
monkeypatch.setattr(LiteLLM_Proxy_MCP_Handler, "_execute_tool_calls", staticmethod(fake_execute))
await responses_main.aresponses_api_with_mcp(
input="hi",
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
follow_up_call = captured_calls[1]
assert follow_up_call["previous_response_id"] == expected_previous_response_id
reasoning_items = [
item for item in follow_up_call["input"] if isinstance(item, dict) and item.get("type") == "reasoning"
]
assert bool(reasoning_items) is (store is False)

View file

@ -258,3 +258,81 @@ async def test_initial_call_failure_is_stashed_for_eager_reraise(monkeypatch):
assert iterator._initial_creation_error is not None
assert "initial boom" in str(iterator._initial_creation_error)
def _reasoning_item(encrypted_content: str):
return {"type": "reasoning", "id": "rs_1", "summary": [], "encrypted_content": encrypted_content}
@pytest.mark.asyncio
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 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)
aresponses_mock = AsyncMock(side_effect=[_text_only_stream("done")])
monkeypatch.setattr(responses_main_module, "aresponses", aresponses_mock)
iterator = MCPEnhancedStreamingIterator(
base_iterator=_FakeAsyncStream(
[
_output_item_added_chunk(),
_completed_chunk([_reasoning_item("gAAAAA-opaque-blob"), _function_call("call_1", "read_wiki_contents")]),
]
),
mcp_events=[],
tool_server_map={"read_wiki_contents": "deepwiki"},
mcp_tools_with_litellm_proxy=[{"require_approval": "never"}],
user_api_key_auth=None,
original_request_params={
"model": "gpt-5",
"input": "what is berriai/litellm?",
"tools": [{"type": "mcp"}],
"store": False,
"previous_response_id": "resp_prev",
},
)
_ = [chunk async for chunk in iterator]
assert aresponses_mock.call_count == 1
follow_up_kwargs = aresponses_mock.call_args_list[0].kwargs
assert follow_up_kwargs["previous_response_id"] == "resp_prev"
assert _reasoning_item("gAAAAA-opaque-blob") in follow_up_kwargs["input"]
@pytest.mark.asyncio
async def test_streaming_follow_up_keeps_previous_response_id_when_stored(monkeypatch):
"""The stateful default is unchanged: previous_response_id still links the follow-up."""
_mock_mcp_environment(monkeypatch)
aresponses_mock = AsyncMock(side_effect=[_text_only_stream("done")])
monkeypatch.setattr(responses_main_module, "aresponses", aresponses_mock)
iterator = MCPEnhancedStreamingIterator(
base_iterator=_FakeAsyncStream(
[
_output_item_added_chunk(),
_completed_chunk([_reasoning_item("gAAAAA-opaque-blob"), _function_call("call_1", "read_wiki_contents")]),
]
),
mcp_events=[],
tool_server_map={"read_wiki_contents": "deepwiki"},
mcp_tools_with_litellm_proxy=[{"require_approval": "never"}],
user_api_key_auth=None,
original_request_params={
"model": "gpt-5",
"input": "what is berriai/litellm?",
"tools": [{"type": "mcp"}],
"previous_response_id": "resp_prev",
},
)
_ = [chunk async for chunk in iterator]
follow_up_kwargs = aresponses_mock.call_args_list[0].kwargs
assert follow_up_kwargs["previous_response_id"] == "resp_prev"
assert not [item for item in follow_up_kwargs["input"] if item.get("type") == "reasoning"]