diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index a75b3768636..ec799cf9a70 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -459,7 +459,10 @@ async def acompletion_with_mcp( ) # Make follow-up call with streaming - follow_up_call_args: Final = dict(self.base_call_args) + follow_up_call_args: Final = LiteLLM_Proxy_MCP_Handler.prepare_follow_up_call_params( + self.base_call_args, + original_stream_setting=True, + ) follow_up_call_args["messages"] = follow_up_messages follow_up_call_args["stream"] = True # Ensure follow-up call doesn't trigger MCP handler again @@ -631,7 +634,10 @@ async def acompletion_with_mcp( ) # Make follow-up call with original stream setting - follow_up_call_args: Final = dict(base_call_args) + follow_up_call_args: Final = LiteLLM_Proxy_MCP_Handler.prepare_follow_up_call_params( + base_call_args, + original_stream_setting=stream, + ) follow_up_call_args["messages"] = follow_up_messages follow_up_call_args["stream"] = stream diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 197d0c02ba8..ad54796b102 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -71,6 +71,32 @@ class LiteLLM_Proxy_MCP_Handler: This handles when a user passes mcp server_url="litellm_proxy" in their tools. """ + @staticmethod + def prepare_chained_call_params( + params: Mapping[str, Any], + ) -> dict[str, Any]: # mutable-ok: returns a sanitized copy for the next provider call + """Copy request params without state owned by the previous LLM call. + + MCP auto-execution keeps the trace identifier so chained rounds remain + correlated, but each provider call must create its own logging object and + call identifier. Reusing either makes success dispatch one-shot and drops + spend rows for later rounds. + """ + follow_up_params = dict(params) + follow_up_params.pop("litellm_logging_obj", None) + follow_up_params.pop("litellm_call_id", None) + if follow_up_params.get("web_search_options") is None: + follow_up_params.pop("web_search_options", None) + + nested_params = follow_up_params.get("litellm_params") + if isinstance(nested_params, dict): + nested_params = dict(nested_params) + nested_params.pop("litellm_logging_obj", None) + nested_params.pop("litellm_call_id", None) + follow_up_params["litellm_params"] = nested_params + + return follow_up_params + @staticmethod def _get_parent_request_tags(kwargs: dict[str, Any] | None) -> list[str]: """Tags from the parent LLM request, using the same extraction logic as standard logging (incl. User-Agent).""" @@ -702,13 +728,19 @@ class LiteLLM_Proxy_MCP_Handler: }, } ] - tool_logging_call_id = litellm_call_id or str(uuid.uuid4()) + # SpendLogs.request_id is unique. The parent LLM call ID is + # therefore metadata, not the tool execution's call ID: reusing + # it causes all but the first MCP tool row to be skipped by the + # database's duplicate protection. + tool_logging_call_id = str(uuid.uuid4()) logging_metadata: dict[str, object] = { "tool_call_id": tool_call_id, "tool_name": sanitized_tool_name, "server_name": server_name, "headers": logging_safe_headers, } + if litellm_call_id: + logging_metadata["parent_litellm_call_id"] = litellm_call_id logging_request_data = { "model": f"MCP: {tool_name}", "metadata": logging_metadata, @@ -1229,14 +1261,14 @@ class LiteLLM_Proxy_MCP_Handler: return initial_params @staticmethod - def _prepare_follow_up_call_params(call_params: dict[str, Any], original_stream_setting: bool) -> dict[str, Any]: + def prepare_follow_up_call_params(call_params: dict[str, Any], original_stream_setting: bool) -> dict[str, Any]: """ Prepare call parameters for the follow-up LLM call after tool execution. Restores the original streaming setting and removes tool_choice since we're now providing tool results, not requesting tool calls. """ - follow_up_params: Final = call_params.copy() + follow_up_params: Final = LiteLLM_Proxy_MCP_Handler.prepare_chained_call_params(call_params) # Restore original streaming setting for follow-up call follow_up_params["stream"] = original_stream_setting @@ -1246,6 +1278,8 @@ class LiteLLM_Proxy_MCP_Handler: return follow_up_params + _prepare_follow_up_call_params = prepare_follow_up_call_params + @staticmethod def _add_mcp_output_elements_to_response( response: ResponsesAPIResponse, diff --git a/litellm/responses/mcp/mcp_streaming_iterator.py b/litellm/responses/mcp/mcp_streaming_iterator.py index 8f5dc926c68..c02240df438 100644 --- a/litellm/responses/mcp/mcp_streaming_iterator.py +++ b/litellm/responses/mcp/mcp_streaming_iterator.py @@ -788,7 +788,9 @@ class MCPEnhancedStreamingIterator(BaseResponsesAPIStreamingIterator): ) # Make follow-up call with streaming - follow_up_params: Final = self.original_request_params.copy() + follow_up_params: Final = LiteLLM_Proxy_MCP_Handler.prepare_chained_call_params( + self.original_request_params + ) follow_up_params.update( { "input": follow_up_input, diff --git a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py index b33ed3bb581..5006ea58087 100644 --- a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -1,20 +1,20 @@ +import importlib import subprocess import sys import textwrap import types +from typing import Any, cast from unittest.mock import AsyncMock, MagicMock import pytest from fastapi import HTTPException -import importlib - from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing +from litellm.types.responses.main import OutputFunctionToolCall +from litellm.types.utils import ModelResponse + from litellm.responses.mcp.litellm_proxy_mcp_handler import ( LiteLLM_Proxy_MCP_Handler, ) -from typing import Any, cast -from litellm.types.utils import ModelResponse -from litellm.types.responses.main import OutputFunctionToolCall class _DummyMCPResult: @@ -105,9 +105,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" @@ -177,9 +175,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}] @@ -299,9 +295,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 = [ @@ -375,7 +369,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey fake_manager = types.SimpleNamespace( get_registry=MagicMock(return_value={}), - call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom")) + 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", @@ -383,9 +377,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey ) 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") @@ -403,14 +395,11 @@ 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 -async def test_execute_tool_calls_passes_litellm_call_id_and_trace_id_to_function_setup( +async def test_execute_tool_calls_uses_unique_call_ids_and_preserves_parent_context( monkeypatch, ): """ @@ -420,22 +409,23 @@ async def test_execute_tool_calls_passes_litellm_call_id_and_trace_id_to_functio _setup_proxy_logging(monkeypatch) call_tool_mock = _setup_mcp_call_environment(monkeypatch) - captured = {} + captured = [] def fake_function_setup(*_args, **kwargs): - captured.update(kwargs) + captured.append(kwargs) return None, None # 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" - tool_calls = [{"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}}] + tool_calls = [ + {"id": "call-1", "function": {"name": tool_name, "arguments": "{}"}}, + {"id": "call-2", "function": {"name": tool_name, "arguments": "{}"}}, + ] await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( tool_server_map={tool_name: "deepwiki"}, @@ -446,10 +436,41 @@ async def test_execute_tool_calls_passes_litellm_call_id_and_trace_id_to_functio ) # Ensure the tool call was attempted (sanity) - assert call_tool_mock.await_count == 1 + assert call_tool_mock.await_count == 2 + assert len(captured) == 2 + assert captured[0]["litellm_call_id"] != captured[1]["litellm_call_id"] + assert all(item["litellm_call_id"] != "cid" for item in captured) + assert all(item["litellm_trace_id"] == "tid" for item in captured) + assert all(item["metadata"]["parent_litellm_call_id"] == "cid" for item in captured) - assert captured.get("litellm_call_id") == "cid" - assert captured.get("litellm_trace_id") == "tid" + +def test_prepare_follow_up_call_params_resets_per_call_logging_state(): + original = { + "model": "gpt-4", + "litellm_call_id": "parent-call", + "litellm_logging_obj": object(), + "litellm_trace_id": "trace-1", + "tool_choice": "auto", + "web_search_options": None, + "litellm_params": { + "litellm_call_id": "nested-parent-call", + "litellm_logging_obj": object(), + "metadata": {"team": "legal"}, + }, + } + + follow_up = LiteLLM_Proxy_MCP_Handler.prepare_follow_up_call_params(original, original_stream_setting=True) + + assert "litellm_call_id" not in follow_up + assert "litellm_logging_obj" not in follow_up + assert "litellm_call_id" not in follow_up["litellm_params"] + assert "litellm_logging_obj" not in follow_up["litellm_params"] + assert follow_up["litellm_trace_id"] == "trace-1" + assert follow_up["stream"] is True + assert "tool_choice" not in follow_up + assert "web_search_options" not in follow_up + assert follow_up["litellm_params"]["metadata"] == {"team": "legal"} + assert original["litellm_call_id"] == "parent-call" @pytest.mark.asyncio @@ -515,9 +536,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 == [] @@ -528,9 +547,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"] @@ -567,9 +584,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"], ) @@ -589,9 +604,7 @@ async def test_execute_tool_calls_exposes_sanitized_client_headers_to_logging(mo 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" @@ -617,9 +630,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"