From 76f79610b3ee63474e9f886c15c1a24ada9fba53 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Mon, 12 Jan 2026 10:32:51 +0900 Subject: [PATCH] fix: forward MCP extra headers case-insensitively --- .../mcp_server/mcp_server_manager.py | 12 ++++- .../proxy/_experimental/mcp_server/server.py | 13 +++++- .../mcp_server/test_mcp_server.py | 27 +++++++++++ .../mcp_server/test_mcp_server_manager.py | 45 +++++++++++++++++++ 4 files changed, 93 insertions(+), 4 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 1029f2241a1..0b81bd7aff7 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1825,9 +1825,17 @@ class MCPServerManager: if mcp_server.extra_headers and raw_headers: if extra_headers is None: extra_headers = {} + + normalized_raw_headers = { + str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str) + } for header in mcp_server.extra_headers: - if isinstance(header, str) and header in raw_headers: - extra_headers[header] = raw_headers[header] + if not isinstance(header, str): + continue + header_value = normalized_raw_headers.get(header.lower()) + if header_value is None: + continue + extra_headers[header] = header_value if mcp_server.static_headers: if extra_headers is None: diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 9c7001266f0..2adfe97c611 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -715,9 +715,18 @@ if MCP_AVAILABLE: if server.extra_headers and raw_headers: if extra_headers is None: extra_headers = {} + + normalized_raw_headers = { + str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str) + } + for header in server.extra_headers: - if header in raw_headers: - extra_headers[header] = raw_headers[header] + if not isinstance(header, str): + continue + header_value = normalized_raw_headers.get(header.lower()) + if header_value is None: + continue + extra_headers[header] = header_value if server_auth_header is None: server_auth_header = mcp_auth_header 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 f1558ac5791..7aa176aaac5 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 @@ -79,6 +79,33 @@ async def test_mcp_server_tool_call_body_contains_request_data(): assert body["arguments"] == tool_arguments +def test_prepare_mcp_server_headers_case_insensitive_extra_headers(): + try: + from litellm.proxy._experimental.mcp_server.server import ( + _prepare_mcp_server_headers, + ) + except ImportError: + pytest.skip("MCP server not available") + + server = MCPServer( + server_id="server-case", + name="server", + transport=MCPTransport.http, + extra_headers=["Authorization"], + ) + + server_auth_header, extra_headers = _prepare_mcp_server_headers( + server=server, + mcp_server_auth_headers=None, + mcp_auth_header=None, + oauth2_headers=None, + raw_headers={"authorization": "Bearer token"}, + ) + + assert server_auth_header is None + assert extra_headers == {"Authorization": "Bearer token"} + + @pytest.mark.asyncio async def test_get_prompts_from_mcp_servers_success(): try: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index d59b3f04ef5..9ddfbf2059c 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -12,6 +12,7 @@ sys.path.insert(0, "../../../../../") import httpx from mcp import ReadResourceResult, Resource from mcp.types import ( + CallToolResult, GetPromptResult, Prompt, ResourceTemplate, @@ -286,6 +287,50 @@ class TestMCPServerManager: assert len(result) == 1 assert result[0].name == "github_tool_1" + @pytest.mark.asyncio + async def test_call_regular_mcp_tool_case_insensitive_extra_headers(self): + """_call_regular_mcp_tool should forward headers regardless of original casing.""" + + manager = MCPServerManager() + server = MCPServer( + server_id="server-case-call", + name="case-call-server", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.authorization, + extra_headers=["Authorization"], + ) + + mock_client = AsyncMock() + mock_client.call_tool = AsyncMock( + return_value=CallToolResult(content=[], isError=False) + ) + captured_extra_headers = None + + def capture_create_mcp_client( + server, mcp_auth_header, extra_headers, stdio_env + ): # pragma: no cover - helper + nonlocal captured_extra_headers + captured_extra_headers = extra_headers + return mock_client + + manager._create_mcp_client = MagicMock(side_effect=capture_create_mcp_client) + + result = await manager._call_regular_mcp_tool( + mcp_server=server, + original_tool_name="tool", + arguments={}, + tasks=[], + mcp_auth_header=None, + mcp_server_auth_headers=None, + oauth2_headers=None, + raw_headers={"authorization": "Bearer token"}, + proxy_logging_obj=None, + ) + + assert captured_extra_headers == {"Authorization": "Bearer token"} + assert isinstance(result, CallToolResult) + @pytest.mark.asyncio async def test_get_prompts_from_server_success(self): """Ensure prompts are fetched and prefixed when requested."""