From 56d7096f510a6b0cf2d963ce03718c3f511d02a2 Mon Sep 17 00:00:00 2001 From: onatozmenn Date: Tue, 4 Aug 2026 15:28:12 +0300 Subject: [PATCH] fix(mcp): put the caller's tag header back on the synthetic tool-call request The first pass carried the tags in the request body and built the header lookup out of new dict literals, which pushed LIT002 past its ceiling. The tool-call handler now restores the one header the tag merge actually reads onto the request it synthesizes, which is closer to the defect anyway: that request was dropping every caller header list_tools reuses the same header read, and the totals the type-discipline gate counts are back to the base --- .../proxy/_experimental/mcp_server/server.py | 44 +++++++++++++------ .../mcp_server/test_mcp_server.py | 17 ++++--- 2 files changed, 40 insertions(+), 21 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 818b05f1975..83d682526db 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -184,15 +184,32 @@ 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.""" +def _request_tags_header( + raw_headers: Mapping[str, str] | None, +) -> str | None: + """The caller's ``x-litellm-tags`` value, read case-insensitively like the other header + lookups in this module. ``None`` when the caller sent no tags.""" 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={}) + for key, value in raw_headers.items(): + if isinstance(key, str) and key.lower() == "x-litellm-tags": + return value or None + return None + + +def _request_tags_from_raw_headers( + raw_headers: Mapping[str, str] | None, +) -> Sequence[str] | None: + """The caller's tags, parsed by the same helper the LLM routes use so an MCP operation and a + chat completion attribute an identical header identically.""" + header_value = _request_tags_header(raw_headers) + if header_value is None: + return None + return LiteLLMProxyRequestSetup.add_request_tag_to_metadata( + llm_router=None, + headers={"x-litellm-tags": header_value}, # mutable-ok: the shared parser reads a plain dict + data={}, # mutable-ok: no request body to read tags from on this path + ) def _jsonrpc_text_has_top_level_method(text: str) -> bool: @@ -1045,24 +1062,23 @@ if MCP_AVAILABLE: host_progress_callback: Final = _capture_host_progress_callback(server) # Create a body date for logging - request_tags: Final = _request_tags_from_raw_headers(raw_headers) - body_data: Final = { - "name": name, - "arguments": arguments, - **({"tags": request_tags} if request_tags else {}), - } + body_data: Final = {"name": name, "arguments": arguments} # 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: body_data["litellm_trace_id"] = chain_id body_data["litellm_session_id"] = chain_id + tags_header: Final = _request_tags_header(raw_headers) + tags_scope_header: Final = ( + ((b"x-litellm-tags", tags_header.encode("latin-1")),) if tags_header is not None else () + ) request: Final = Request( scope={ "type": "http", "method": "POST", "path": "/mcp/tools/call", - "headers": [(b"content-type", b"application/json")], + "headers": [(b"content-type", b"application/json"), *tags_scope_header], } ) if user_api_key_auth is not None: 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 39c6632f669..098924cee85 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 @@ -118,15 +118,16 @@ async def test_mcp_server_tool_call_body_contains_request_data(): @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.""" + """The tool-call handler hands `add_litellm_data_to_request` a synthetic request that carried + only a content type, so the caller's `x-litellm-tags` never reached the tag merge running there + and a tools/call spend log had no tags. The synthetic request now carries the header, so the + shared parser resolves it exactly as it does on an LLM route.""" try: from litellm.proxy._experimental.mcp_server.server import ( mcp_server_tool_call, set_auth_context, ) + from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup except ImportError: pytest.skip("MCP server not available") @@ -135,10 +136,12 @@ async def test_mcp_server_tool_call_carries_x_litellm_tags_header_into_request_d raw_headers={"X-LiteLLM-Tags": "application:orders, service:checkout"}, ) - captured_data = {} + resolved_tags = {} async def mock_add_litellm_data_to_request(data, request, user_api_key_dict, proxy_config): - captured_data.update(data) + resolved_tags["tags"] = LiteLLMProxyRequestSetup.add_request_tag_to_metadata( + llm_router=None, headers=dict(request.headers), data={} + ) return data async def mock_call_mcp_tool(*args, **kwargs): @@ -155,7 +158,7 @@ async def test_mcp_server_tool_call_carries_x_litellm_tags_header_into_request_d 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"] + assert resolved_tags["tags"] == ["application:orders", "service:checkout"] @pytest.mark.asyncio