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.
This commit is contained in:
nvr 2026-04-09 04:56:24 -07:00
parent 97f722f558
commit 6df35fa1a6
2 changed files with 230 additions and 49 deletions

View file

@ -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

View file

@ -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):
"""