From 6df35fa1a6be96464fa8514eb6b7540f10fbb13a Mon Sep 17 00:00:00 2001 From: nvr Date: Thu, 9 Apr 2026 04:56:24 -0700 Subject: [PATCH] Fix chat MCP follow-up auto-execution Allow non-streaming chat MCP requests to execute multiple tool rounds before returning a final answer. Drop tool_choice from follow-up calls so required-tool initial requests do not trap the model in repeated MCP-only turns. Add a regression test covering a resolve-then-query-docs style multi-round flow and aggregated MCP metadata on the final response. --- .../responses/mcp/chat_completions_handler.py | 109 ++++++----- .../mcp/test_chat_completions_handler.py | 170 ++++++++++++++++++ 2 files changed, 230 insertions(+), 49 deletions(-) diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index 24b5db28571..a8843a8ea9c 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -15,6 +15,8 @@ from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper +MAX_CHAT_MCP_AUTO_EXECUTION_ROUNDS = 6 + def _add_mcp_metadata_to_response( response: Union[ModelResponse, CustomStreamWrapper], @@ -602,69 +604,78 @@ async def acompletion_with_mcp( # noqa: PLR0915 return cast(CustomStreamWrapper, MCPStreamWrapper(initial_stream, iterator)) - # Non-streaming mode: use existing logic - initial_call_args = dict(base_call_args) - initial_call_args["stream"] = False + # Non-streaming mode: iterate tool rounds until the model produces a final answer. + call_args = dict(base_call_args) + call_args["stream"] = False if mock_tool_calls is not None: - initial_call_args["mock_tool_calls"] = mock_tool_calls + call_args["mock_tool_calls"] = mock_tool_calls - # Make initial call - initial_response = await litellm_acompletion(**initial_call_args) + response = await litellm_acompletion(**call_args) + current_messages = messages + aggregated_tool_calls: List[Any] = [] + aggregated_tool_results: List[Any] = [] - if not isinstance(initial_response, ModelResponse): - return initial_response + for _round_idx in range(MAX_CHAT_MCP_AUTO_EXECUTION_ROUNDS): + if not isinstance(response, ModelResponse): + return response - # Extract tool calls from response - tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response( - response=initial_response - ) - - if not tool_calls: - _add_mcp_metadata_to_response( - response=initial_response, - openai_tools=openai_tools, + tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response( + response=response ) - return initial_response - # Execute tool calls - tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( - tool_server_map=tool_server_map, - tool_calls=tool_calls, - user_api_key_auth=user_api_key_auth, - mcp_auth_header=mcp_auth_header, - mcp_server_auth_headers=mcp_server_auth_headers, - oauth2_headers=oauth2_headers, - raw_headers=raw_headers, - litellm_call_id=kwargs.get("litellm_call_id"), - litellm_trace_id=kwargs.get("litellm_trace_id"), - ) + if not tool_calls: + _add_mcp_metadata_to_response( + response=response, + openai_tools=openai_tools, + tool_calls=aggregated_tool_calls or None, + tool_results=aggregated_tool_results or None, + ) + return response - if not tool_results: - _add_mcp_metadata_to_response( - response=initial_response, - openai_tools=openai_tools, + aggregated_tool_calls.extend(tool_calls) + + tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map=tool_server_map, tool_calls=tool_calls, + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + litellm_call_id=kwargs.get("litellm_call_id"), + litellm_trace_id=kwargs.get("litellm_trace_id"), ) - return initial_response - # Create follow-up messages with tool results - follow_up_messages = LiteLLM_Proxy_MCP_Handler._create_follow_up_messages_for_chat( - original_messages=messages, - response=initial_response, - tool_results=tool_results, - ) + if not tool_results: + _add_mcp_metadata_to_response( + response=response, + openai_tools=openai_tools, + tool_calls=aggregated_tool_calls or None, + tool_results=aggregated_tool_results or None, + ) + return response - # Make follow-up call with original stream setting - follow_up_call_args = dict(base_call_args) - follow_up_call_args["messages"] = follow_up_messages - follow_up_call_args["stream"] = stream + aggregated_tool_results.extend(tool_results) + current_messages = LiteLLM_Proxy_MCP_Handler._create_follow_up_messages_for_chat( + original_messages=current_messages, + response=response, + tool_results=tool_results, + ) - response = await litellm_acompletion(**follow_up_call_args) - if isinstance(response, (ModelResponse, CustomStreamWrapper)): + follow_up_call_args = dict(base_call_args) + follow_up_call_args["messages"] = current_messages + follow_up_call_args["stream"] = False + # The follow-up provides tool results, so preserving tool_choice traps + # the model in extra MCP-only turns instead of allowing a final answer. + follow_up_call_args.pop("tool_choice", None) + + response = await litellm_acompletion(**follow_up_call_args) + + if isinstance(response, ModelResponse): _add_mcp_metadata_to_response( response=response, openai_tools=openai_tools, - tool_calls=tool_calls, - tool_results=tool_results, + tool_calls=aggregated_tool_calls or None, + tool_results=aggregated_tool_results or None, ) return response diff --git a/tests/test_litellm/responses/mcp/test_chat_completions_handler.py b/tests/test_litellm/responses/mcp/test_chat_completions_handler.py index 3ba41705733..e1c203ecd55 100644 --- a/tests/test_litellm/responses/mcp/test_chat_completions_handler.py +++ b/tests/test_litellm/responses/mcp/test_chat_completions_handler.py @@ -361,6 +361,176 @@ async def test_acompletion_with_mcp_auto_exec_performs_follow_up(monkeypatch): assert follow_up_call["stream"] is True +@pytest.mark.asyncio +async def test_acompletion_with_mcp_non_streaming_auto_exec_supports_multiple_rounds( + monkeypatch, +): + tools = [ + { + "type": "mcp", + "server_url": "litellm_proxy/mcp/local", + "require_approval": "never", + } + ] + openai_tools = [{"type": "function", "function": {"name": "local_search"}}] + + first_tool_call = { + "id": "call-1", + "type": "function", + "function": {"name": "local_search", "arguments": "{}"}, + } + second_tool_call = { + "id": "call-2", + "type": "function", + "function": {"name": "local_search", "arguments": '{"phase":"second"}'}, + } + + first_response = ModelResponse( + id="resp-1", + model="test-model", + created=1234567890, + object="chat.completion", + choices=[ + { + "index": 0, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": "", + "tool_calls": [first_tool_call], + }, + } + ], + ) + second_response = ModelResponse( + id="resp-2", + model="test-model", + created=1234567891, + object="chat.completion", + choices=[ + { + "index": 0, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": "", + "tool_calls": [second_tool_call], + }, + } + ], + ) + final_response = ModelResponse( + id="resp-3", + model="test-model", + created=1234567892, + object="chat.completion", + choices=[ + { + "index": 0, + "finish_reason": "stop", + "message": { + "role": "assistant", + "content": "final answer", + }, + } + ], + ) + + mock_acompletion = AsyncMock( + side_effect=[first_response, second_response, final_response] + ) + + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_should_use_litellm_mcp_gateway", + staticmethod(lambda tools: True), + ) + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_parse_mcp_tools", + staticmethod(lambda tools: (tools, [])), + ) + + async def mock_process(**_): + return (tools, {"local_search": "local"}) + + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_process_mcp_tools_without_openai_transform", + mock_process, + ) + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_transform_mcp_tools_to_openai", + staticmethod(lambda *_, **__: openai_tools), + ) + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_should_auto_execute_tools", + staticmethod(lambda **_: True), + ) + + async def mock_execute(**kwargs): + tool_calls = kwargs["tool_calls"] + tool_id = tool_calls[0]["id"] + return [ + { + "tool_call_id": tool_id, + "name": "local_search", + "result": f"result-for-{tool_id}", + } + ] + + monkeypatch.setattr( + LiteLLM_Proxy_MCP_Handler, + "_execute_tool_calls", + mock_execute, + ) + monkeypatch.setattr( + ResponsesAPIRequestUtils, + "extract_mcp_headers_from_request", + staticmethod(lambda **_: (None, None, None, None)), + ) + + with patch("litellm.acompletion", mock_acompletion), patch.object( + chat_completions_handler, + "litellm_acompletion", + mock_acompletion, + create=True, + ): + result = await acompletion_with_mcp( + model="gpt-4o-mini", + messages=[{"role": "user", "content": "hello"}], + tools=tools, + stream=False, + tool_choice="required", + ) + + assert isinstance(result, ModelResponse) + assert result.choices[0].message.content == "final answer" + assert mock_acompletion.await_count == 3 + + first_call = mock_acompletion.await_args_list[0].kwargs + second_call = mock_acompletion.await_args_list[1].kwargs + third_call = mock_acompletion.await_args_list[2].kwargs + + assert first_call["stream"] is False + assert first_call["tool_choice"] == "required" + assert second_call["stream"] is False + assert "tool_choice" not in second_call + assert third_call["stream"] is False + assert "tool_choice" not in third_call + assert any(msg.get("role") == "tool" for msg in second_call["messages"]) + assert any(msg.get("role") == "tool" for msg in third_call["messages"]) + + provider_fields = result.choices[0].message.provider_specific_fields + assert provider_fields is not None + assert len(provider_fields["mcp_tool_calls"]) == 2 + assert len(provider_fields["mcp_call_results"]) == 2 + assert provider_fields["mcp_call_results"][0]["tool_call_id"] == "call-1" + assert provider_fields["mcp_call_results"][1]["tool_call_id"] == "call-2" + + @pytest.mark.asyncio async def test_acompletion_with_mcp_adds_metadata_to_streaming(monkeypatch): """