mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
commit
6960a42008
5 changed files with 313 additions and 9 deletions
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue