fix(mcp): preserve original response chain on follow-up

This commit is contained in:
Devin AI 2026-07-16 12:11:25 +00:00
parent 2f03789927
commit 98432de3f8
3 changed files with 43 additions and 37 deletions

View file

@ -356,7 +356,7 @@ async def aresponses_api_with_mcp(
follow_up_input=follow_up_input,
model=model,
all_tools=all_tools,
response_id=response.id,
previous_response_id=previous_response_id,
**follow_up_call_params,
)

View file

@ -1005,7 +1005,7 @@ class LiteLLM_Proxy_MCP_Handler:
follow_up_input: List[Any],
model: str,
all_tools: Optional[List[Any]],
response_id: str,
previous_response_id: Optional[str],
**call_params: Any,
) -> Union[ResponsesAPIResponse, BaseResponsesAPIStreamingIterator]:
"""Make follow-up response API call with tool results."""
@ -1013,7 +1013,7 @@ class LiteLLM_Proxy_MCP_Handler:
input=follow_up_input,
model=model,
tools=all_tools, # Keep tools for potential future calls
previous_response_id=response_id, # Link to previous response
previous_response_id=previous_response_id,
**call_params,
)

View file

@ -103,9 +103,7 @@ def test_extract_tool_calls_from_chat_response_handles_tool_calls():
object="chat.completion",
)
tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response(
response
)
tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response(response)
assert len(tool_calls) == 1
assert tool_calls[0]["function"]["name"] == "foo"
@ -175,9 +173,7 @@ def test_transform_mcp_tools_to_openai_uses_chat_format(monkeypatch):
fake_transform_responses,
)
chat_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(
["tool"], target_format="chat"
)
chat_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(["tool"], target_format="chat")
resp_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(["tool"])
assert chat_tools == [{"chat": True}]
@ -216,6 +212,35 @@ def test_create_follow_up_input_handles_response_function_tool_call():
]
@pytest.mark.asyncio
@pytest.mark.parametrize("previous_response_id", [None, "resp_previous"])
async def test_make_follow_up_call_uses_original_previous_response_id(
monkeypatch: pytest.MonkeyPatch,
previous_response_id: str | None,
):
expected_response = MagicMock()
aresponses_mock = AsyncMock(return_value=expected_response)
handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler")
monkeypatch.setattr(handler_module, "aresponses", aresponses_mock)
response = await LiteLLM_Proxy_MCP_Handler._make_follow_up_call(
follow_up_input=[
{
"type": "function_call_output",
"call_id": "call-1",
"output": "done",
}
],
model="test-model",
all_tools=[],
previous_response_id=previous_response_id,
)
assert response is expected_response
assert aresponses_mock.await_args is not None
assert aresponses_mock.await_args.kwargs["previous_response_id"] == previous_response_id
@pytest.mark.asyncio
async def test_execute_tool_calls_strips_server_prefix(monkeypatch):
call_tool_mock = _setup_mcp_call_environment(monkeypatch)
@ -297,9 +322,7 @@ async def test_execute_tool_calls_strips_prefix_when_alias_differs_from_server_n
)
from litellm.proxy._experimental.mcp_server import mcp_server_manager as _msm
_msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock(
return_value=fake_server
)
_msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock(return_value=fake_server)
tool_name = "my_deepwiki-read_wiki_structure"
tool_calls = [
@ -371,18 +394,14 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey
"""
post_call_failure_hook = _setup_proxy_logging(monkeypatch)
fake_manager = types.SimpleNamespace(
call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom"))
)
fake_manager = types.SimpleNamespace(call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom")))
monkeypatch.setattr(
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
fake_manager,
)
tool_name = "deepwiki-read_wiki_structure"
tool_calls = [
{"id": "call-err", "function": {"name": tool_name, "arguments": "{}"}}
]
tool_calls = [{"id": "call-err", "function": {"name": tool_name, "arguments": "{}"}}]
user_auth = types.SimpleNamespace(api_key="test_key", user_id="test_user")
@ -400,10 +419,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey
post_call_failure_hook.assert_awaited_once()
assert post_call_failure_hook.await_args is not None
assert (
post_call_failure_hook.await_args.kwargs.get("route")
== "/responses/mcp/call_tool"
)
assert post_call_failure_hook.await_args.kwargs.get("route") == "/responses/mcp/call_tool"
@pytest.mark.asyncio
@ -426,9 +442,7 @@ async def test_execute_tool_calls_passes_litellm_call_id_and_trace_id_to_functio
# NOTE: Don't patch via dotted string path here because `litellm.responses`
# is a function attribute on the `litellm` package (shadowing the submodule),
# which breaks monkeypatch's importpath resolution.
handler_module = importlib.import_module(
"litellm.responses.mcp.litellm_proxy_mcp_handler"
)
handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler")
monkeypatch.setattr(handler_module, "function_setup", fake_function_setup)
tool_name = "deepwiki-read_wiki_structure"
@ -475,9 +489,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch
user_auth = types.SimpleNamespace(api_key="test_key", user_id="test_user")
tools, _server_names = await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager(
user_api_key_auth=user_auth,
mcp_tools_with_litellm_proxy=[
{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}
],
mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}],
)
assert tools == []
@ -488,9 +500,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch
def test_get_parent_request_tags_from_metadata():
tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(
{"metadata": {"tags": ["team-a", "prod"]}}
)
tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags({"metadata": {"tags": ["team-a", "prod"]}})
assert tags == ["team-a", "prod"]
@ -526,9 +536,7 @@ async def test_get_mcp_tools_from_manager_forwards_request_tags(monkeypatch):
await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager(
user_api_key_auth=types.SimpleNamespace(api_key="k", user_id="u"),
mcp_tools_with_litellm_proxy=[
{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}
],
mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}],
request_tags=["team-a"],
)
@ -545,9 +553,7 @@ async def test_execute_tool_calls_propagates_request_tags_to_function_setup(monk
captured.update(kwargs)
return None, None
handler_module = importlib.import_module(
"litellm.responses.mcp.litellm_proxy_mcp_handler"
)
handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler")
monkeypatch.setattr(handler_module, "function_setup", fake_function_setup)
tool_name = "deepwiki-read_wiki_structure"