From cfab4f62d3d9c4253682381f64e28977cadf9e72 Mon Sep 17 00:00:00 2001 From: onatozmenn Date: Tue, 4 Aug 2026 14:59:44 +0300 Subject: [PATCH] fix(mcp): honor x-litellm-tags on the MCP gateway's tools/list and tools/call LLM routes read the header in add_litellm_data_to_request, so caller tags land in LiteLLM_SpendLogs.request_tags. Nothing read it on the MCP gateway: the tool-call handler hands that same helper a synthetic request carrying only a content type, so the header never reached the tag merge, and the list_tools spend log has a request_tags parameter no caller populates Tool calls now carry the caller's tags in the body, which that helper already reads, and list_tools falls back to the header when no tags were passed in. Per-application attribution behind a gateway that stamps the header now works the same for MCP traffic as it does for chat completions --- .../proxy/_experimental/mcp_server/server.py | 21 +- .../mcp_server/test_mcp_server.py | 192 ++++++++++++++++++ 2 files changed, 211 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 49a1f1314f0..818b05f1975 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -184,6 +184,17 @@ def _mcp_session_id_from_headers( return None +def _request_tags_from_raw_headers( + raw_headers: dict[str, str] | None, +) -> list[str] | None: + """The caller's ``x-litellm-tags``, parsed by the same helper the LLM routes use so an + MCP operation and a chat completion attribute an identical header identically.""" + if not raw_headers: + return None + headers = {key.lower(): value for key, value in raw_headers.items() if isinstance(key, str)} + return LiteLLMProxyRequestSetup.add_request_tag_to_metadata(llm_router=None, headers=headers, data={}) + + def _jsonrpc_text_has_top_level_method(text: str) -> bool: """Whether a (possibly truncated) JSON-RPC envelope has a ``method`` key at the root object's top level. @@ -1034,7 +1045,12 @@ if MCP_AVAILABLE: host_progress_callback: Final = _capture_host_progress_callback(server) # Create a body date for logging - body_data: Final = {"name": name, "arguments": arguments} + request_tags: Final = _request_tags_from_raw_headers(raw_headers) + body_data: Final = { + "name": name, + "arguments": arguments, + **({"tags": request_tags} if request_tags else {}), + } # Set trace/session id from raw_headers so spend logs and logging_obj stay consistent (same as A2A) chain_id: Final = get_chain_id_from_headers(raw_headers) if chain_id: @@ -1890,6 +1906,7 @@ if MCP_AVAILABLE: list_tools_call_id: Final = str(uuid.uuid4()) # Derive trace_id from raw_headers when not explicitly passed (same as A2A / MCP call_tool) effective_litellm_trace_id: Final = litellm_trace_id or get_chain_id_from_headers(raw_headers) + effective_request_tags: Final = request_tags or _request_tags_from_raw_headers(raw_headers) spend_logs_metadata: Final[dict[str, object]] = { "mcp_operation": "list_tools", } @@ -1905,7 +1922,7 @@ if MCP_AVAILABLE: "litellm_trace_id": effective_litellm_trace_id, "metadata": { "spend_logs_metadata": spend_logs_metadata, - **({"tags": request_tags} if request_tags else {}), + **({"tags": effective_request_tags} if effective_request_tags else {}), }, # Provide a small input payload for standard logging "input": [ 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 850d01c6e34..39c6632f669 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 @@ -116,6 +116,48 @@ async def test_mcp_server_tool_call_body_contains_request_data(): assert body["arguments"] == tool_arguments +@pytest.mark.asyncio +async def test_mcp_server_tool_call_carries_x_litellm_tags_header_into_request_data(): + """The tool-call handler hands `add_litellm_data_to_request` a synthetic request that carries only + a content type, so the caller's `x-litellm-tags` never reached the tag merge that runs there and + the spend log for a tools/call had no tags. The tags travel in the body instead, which that same + helper already reads, so the header attributes MCP traffic exactly as it does an LLM route.""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + mcp_server_tool_call, + set_auth_context, + ) + except ImportError: + pytest.skip("MCP server not available") + + set_auth_context( + UserAPIKeyAuth(api_key="test_key", user_id="test_user"), + raw_headers={"X-LiteLLM-Tags": "application:orders, service:checkout"}, + ) + + captured_data = {} + + async def mock_add_litellm_data_to_request(data, request, user_api_key_dict, proxy_config): + captured_data.update(data) + return data + + async def mock_call_mcp_tool(*args, **kwargs): + return [{"type": "text", "text": "mocked response"}] + + with patch( + "litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request", + mock_add_litellm_data_to_request, + ): + with patch( + "litellm.proxy._experimental.mcp_server.server.call_mcp_tool", + mock_call_mcp_tool, + ): + with patch("litellm.proxy.proxy_server.proxy_config", MagicMock()): + await mcp_server_tool_call("test_tool", {"param1": "value1"}) + + assert captured_data["tags"] == ["application:orders", "service:checkout"] + + @pytest.mark.asyncio async def test_mcp_server_tool_call_relays_upstream_auth_error_as_iserror(): """The MCP session manager serializes handler exceptions as JSON-RPC errors, so a mid-session @@ -4436,6 +4478,156 @@ async def test_get_tools_from_mcp_servers_logs_list_tools_to_spendlogs_when_enab assert spend_meta["per_server_list_outcomes"] == {"server_a": {"status": "ok", "tool_count": 1}} +@pytest.mark.asyncio +async def test_get_tools_from_mcp_servers_takes_list_tools_tags_from_x_litellm_tags_header(): + """A gateway that stamps `x-litellm-tags` on proxied traffic gets per-application attribution on + LLM routes; list_tools must read the same header so MCP usage is not stuck under the shared key. + Nothing populated `request_tags`, so the header was the only source and it was being dropped.""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + _get_tools_from_mcp_servers, + ) + from litellm.proxy._types import UserAPIKeyAuth + from mcp.types import Tool as MCPTool + except ImportError: + pytest.skip("MCP server not available") + + user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user") + + server_a = MagicMock(name="server_a_obj") + server_a.name = "server_a" + server_a.alias = "server_a" + server_a.server_name = "server_a" + server_a.server_id = "a" + server_a.auth_type = None + server_a.extra_headers = None + + tool_1 = MCPTool(name="server_a-tool_1", description="test tool", inputSchema={"type": "object"}) + + dummy_logging_obj = MagicMock() + dummy_logging_obj.model_call_details = {"metadata": {"spend_logs_metadata": {}}} + dummy_logging_obj.async_success_handler = AsyncMock() + function_setup_kwargs = {} + + def _capture_function_setup(*_args, **kwargs): + function_setup_kwargs.update(kwargs) + return dummy_logging_obj, None + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server_a]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + return_value=(None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + ) as mock_manager, + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools", + side_effect=lambda tools, _server: tools, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions", + new=AsyncMock(side_effect=lambda tools, **_: tools), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.function_setup", + side_effect=_capture_function_setup, + ), + ): + mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1]) + + listing = await _get_tools_from_mcp_servers( + user_api_key_auth=user_auth, + mcp_auth_header=None, + mcp_servers=["server_a"], + mcp_server_auth_headers=None, + raw_headers={"X-LiteLLM-Tags": "application:orders, service:checkout"}, + log_list_tools_to_spendlogs=True, + list_tools_log_source="mcp_protocol", + ) + + assert listing.tools == [tool_1] + assert function_setup_kwargs["metadata"]["tags"] == ["application:orders", "service:checkout"] + + +@pytest.mark.asyncio +async def test_get_tools_from_mcp_servers_prefers_explicit_request_tags_over_the_header(): + """`request_tags` is the resolved value a caller passes in; a header must not override it.""" + try: + from litellm.proxy._experimental.mcp_server.server import ( + _get_tools_from_mcp_servers, + ) + from litellm.proxy._types import UserAPIKeyAuth + from mcp.types import Tool as MCPTool + except ImportError: + pytest.skip("MCP server not available") + + user_auth = UserAPIKeyAuth(api_key="test-key", user_id="test-user") + + server_a = MagicMock(name="server_a_obj") + server_a.name = "server_a" + server_a.alias = "server_a" + server_a.server_name = "server_a" + server_a.server_id = "a" + server_a.auth_type = None + server_a.extra_headers = None + + tool_1 = MCPTool(name="server_a-tool_1", description="test tool", inputSchema={"type": "object"}) + + dummy_logging_obj = MagicMock() + dummy_logging_obj.model_call_details = {"metadata": {"spend_logs_metadata": {}}} + dummy_logging_obj.async_success_handler = AsyncMock() + function_setup_kwargs = {} + + def _capture_function_setup(*_args, **kwargs): + function_setup_kwargs.update(kwargs) + return dummy_logging_obj, None + + with ( + patch( + "litellm.proxy._experimental.mcp_server.server._get_allowed_mcp_servers", + new=AsyncMock(return_value=[server_a]), + ), + patch( + "litellm.proxy._experimental.mcp_server.server._prepare_mcp_server_headers", + return_value=(None, None), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.global_mcp_server_manager", + ) as mock_manager, + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_allowed_tools", + side_effect=lambda tools, _server: tools, + ), + patch( + "litellm.proxy._experimental.mcp_server.server.filter_tools_by_key_team_permissions", + new=AsyncMock(side_effect=lambda tools, **_: tools), + ), + patch( + "litellm.proxy._experimental.mcp_server.server.function_setup", + side_effect=_capture_function_setup, + ), + ): + mock_manager._get_tools_from_server = AsyncMock(return_value=[tool_1]) + + await _get_tools_from_mcp_servers( + user_api_key_auth=user_auth, + mcp_auth_header=None, + mcp_servers=["server_a"], + mcp_server_auth_headers=None, + raw_headers={"x-litellm-tags": "from-header"}, + log_list_tools_to_spendlogs=True, + list_tools_log_source="mcp_protocol", + request_tags=["explicit"], + ) + + assert function_setup_kwargs["metadata"]["tags"] == ["explicit"] + + @pytest.mark.asyncio async def test_get_tools_from_mcp_servers_returns_tools_when_success_logging_fails(): """