From a8b38eb6ad1f064515a3c8a00f078ae7b5e42547 Mon Sep 17 00:00:00 2001 From: kalyanguru18 Date: Mon, 22 Jun 2026 19:59:13 +0530 Subject: [PATCH] fix: preserve MCP failure logging metadata --- .../proxy/_experimental/mcp_server/server.py | 30 +++- .../mcp_server/test_mcp_server.py | 168 ++++++++++++++++++ 2 files changed, 197 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 08e42e918e9..af28c2646ff 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -2770,6 +2770,21 @@ if MCP_AVAILABLE: litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get( "litellm_logging_obj", None ) + failure_request_data = dict(kwargs) + failure_request_data["call_type"] = CallTypes.call_mcp_tool.value + failure_request_data["model"] = f"MCP: {name}" + + if arguments is not None: + original_tool_name, server_name = split_server_prefix_from_name(name) + metadata = failure_request_data.get("metadata") + metadata = metadata.copy() if isinstance(metadata, dict) else {} + metadata["mcp_tool_call_metadata"] = _get_standard_logging_mcp_tool_call( + name=original_tool_name, + arguments=arguments, + server_name=server_name, + session_id=_mcp_session_id_from_headers(raw_headers), + ) + failure_request_data["metadata"] = metadata try: if arguments is None: @@ -2820,8 +2835,21 @@ if MCP_AVAILABLE: from litellm.proxy.proxy_server import proxy_logging_obj if proxy_logging_obj and user_api_key_auth: + if litellm_logging_obj is not None: + model_call_details = litellm_logging_obj.model_call_details + mcp_tool_call_metadata = model_call_details.get( + "mcp_tool_call_metadata" + ) + if mcp_tool_call_metadata is not None: + metadata = failure_request_data.get("metadata") + metadata = metadata.copy() if isinstance(metadata, dict) else {} + metadata["mcp_tool_call_metadata"] = mcp_tool_call_metadata + failure_request_data["metadata"] = metadata + failure_request_data["model"] = ( + model_call_details.get("model") or failure_request_data["model"] + ) await proxy_logging_obj.post_call_failure_hook( - request_data=kwargs, + request_data=failure_request_data, original_exception=e, user_api_key_dict=user_api_key_auth, route="/mcp/call_tool", diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index c86ae966f21..b61c5185123 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -4149,6 +4149,174 @@ async def test_call_mcp_tool_logs_failure_via_post_call_failure_hook(): ) +@pytest.mark.asyncio +async def test_call_mcp_tool_failure_hook_preserves_call_type_and_metadata(): + """ + Ensure unresolved local-registry tool failures log as MCP tool calls with attempted tool metadata. + """ + try: + from litellm.proxy._experimental.mcp_server.server import ( + call_mcp_tool, + global_mcp_server_manager, + ) + from litellm.proxy._types import MCPTransport, UserAPIKeyAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + except ImportError: + pytest.skip("MCP server not available") + + mock_server = MCPServer( + server_id="server-123", + name="test_server", + alias="test_server", + server_name="test_server", + url="https://test-server.com/mcp", + transport=MCPTransport.http, + mcp_info={"server_name": "test_server"}, + ) + + proxy_logging_mock = MagicMock() + proxy_logging_mock.post_call_failure_hook = AsyncMock() + + user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user") + + with ( + patch.object( + global_mcp_server_manager, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[mock_server.server_id], + ), + patch.object( + global_mcp_server_manager, + "get_mcp_server_by_id", + return_value=mock_server, + ), + patch.object( + global_mcp_server_manager, + "_get_mcp_server_from_tool_name", + return_value=None, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers_from_mcp_server_names", + new_callable=AsyncMock, + return_value=[mock_server], + ), + patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + proxy_logging_mock, + ), + ): + with pytest.raises(HTTPException) as exc_info: + await call_mcp_tool( + name="test_server-totally_made_up_tool", + arguments={"x": 1}, + user_api_key_auth=user_auth, + litellm_call_id="cid", + ) + + assert exc_info.value.status_code == 404 + proxy_logging_mock.post_call_failure_hook.assert_awaited_once() + request_data = proxy_logging_mock.post_call_failure_hook.await_args.kwargs[ + "request_data" + ] + assert request_data["call_type"] == "call_mcp_tool" + assert request_data["model"] == "MCP: test_server-totally_made_up_tool" + assert request_data["litellm_call_id"] == "cid" + assert request_data["metadata"]["mcp_tool_call_metadata"] == { + "name": "totally_made_up_tool", + "arguments": {"x": 1}, + "namespaced_tool_name": "test_server/totally_made_up_tool", + "mcp_session_id": None, + } + + +@pytest.mark.asyncio +async def test_call_mcp_tool_failure_hook_uses_logging_metadata_without_clearing_model(): + """ + Ensure logging-object metadata is preserved without replacing the fallback model with None. + """ + try: + from litellm.proxy._experimental.mcp_server.server import ( + call_mcp_tool, + global_mcp_server_manager, + ) + from litellm.proxy._types import MCPTransport, UserAPIKeyAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + except ImportError: + pytest.skip("MCP server not available") + + mock_server = MCPServer( + server_id="server-123", + name="test_server", + alias="test_server", + server_name="test_server", + url="https://test-server.com/mcp", + transport=MCPTransport.http, + mcp_info={"server_name": "test_server"}, + ) + logging_metadata = { + "name": "any_tool", + "arguments": {"x": 2}, + "namespaced_tool_name": "test_server/any_tool", + "mcp_session_id": None, + } + litellm_logging_obj = MagicMock() + litellm_logging_obj.model_call_details = { + "model": None, + "mcp_tool_call_metadata": logging_metadata, + } + litellm_logging_obj.async_failure_handler = AsyncMock() + litellm_logging_obj.failure_handler = MagicMock() + + proxy_logging_mock = MagicMock() + proxy_logging_mock.post_call_failure_hook = AsyncMock() + + user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user") + + with ( + patch.object( + global_mcp_server_manager, + "get_allowed_mcp_servers", + new_callable=AsyncMock, + return_value=[mock_server.server_id], + ), + patch.object( + global_mcp_server_manager, + "get_mcp_server_by_id", + return_value=mock_server, + ), + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers_from_mcp_server_names", + new_callable=AsyncMock, + return_value=[mock_server], + ), + patch( + "litellm.proxy._experimental.mcp_server.server.execute_mcp_tool", + new_callable=AsyncMock, + side_effect=Exception("boom"), + ), + patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + proxy_logging_mock, + ), + ): + with pytest.raises(Exception, match="boom"): + await call_mcp_tool( + name="test_server-any_tool", + arguments={"x": 2}, + user_api_key_auth=user_auth, + litellm_call_id="cid", + litellm_logging_obj=litellm_logging_obj, + ) + + proxy_logging_mock.post_call_failure_hook.assert_awaited_once() + request_data = proxy_logging_mock.post_call_failure_hook.await_args.kwargs[ + "request_data" + ] + assert request_data["model"] == "MCP: test_server-any_tool" + assert request_data["metadata"]["mcp_tool_call_metadata"] == logging_metadata + + @pytest.mark.asyncio async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enabled(): """