From 60c3a418f8a089a7f92ad41c4d0ec054c4df3f69 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 16 Jul 2026 02:11:39 +0000 Subject: [PATCH] fix(mcp): forward OpenAPI extra headers on call_tool --- .../mcp_server/mcp_server_manager.py | 47 ++++++++++++++- .../mcp_server/test_mcp_hook_extra_headers.py | 60 +++++++++++++++++++ 2 files changed, 106 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index b5832533b17..22f6aec792b 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -241,6 +241,34 @@ def _without_authorization( return filtered or None +def _openapi_forwarded_extra_headers( + mcp_server: MCPServer, + raw_headers: Optional[dict[str, str]], + user_api_key_auth: Optional[UserAPIKeyAuth], +) -> Optional[dict[str, str]]: + if not mcp_server.extra_headers or not raw_headers: + return None + + normalized_raw_headers = { + str(header_name).lower(): header_value + for header_name, header_value in raw_headers.items() + if isinstance(header_name, str) + } + strip_caller_authorization = _should_strip_caller_authorization( + mcp_server=mcp_server, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ) + forwarded_headers = { + header_name: normalized_raw_headers[header_name.lower()] + for header_name in mcp_server.extra_headers + if isinstance(header_name, str) + and not (strip_caller_authorization and header_name.lower() == "authorization") + and header_name.lower() in normalized_raw_headers + } + return forwarded_headers or None + + def _extract_upstream_auth_failure( exc: BaseException, ) -> Optional[Tuple[int, Optional[str]]]: @@ -3523,7 +3551,24 @@ class MCPServerManager: "transport to enable hook header injection.", server_name, ) - tasks.append(asyncio.create_task(self._call_openapi_tool_handler(mcp_server, name, arguments))) + forwarded_headers = _openapi_forwarded_extra_headers( + mcp_server=mcp_server, + raw_headers=raw_headers, + user_api_key_auth=user_api_key_auth, + ) + + async def _call_openapi_via_handler() -> CallToolResult: + from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( + _request_extra_headers, + ) + + extra_headers_token = _request_extra_headers.set(forwarded_headers) + try: + return await self._call_openapi_tool_handler(mcp_server, name, arguments) + finally: + _request_extra_headers.reset(extra_headers_token) + + tasks.append(asyncio.create_task(_call_openapi_via_handler())) else: return await self._call_regular_mcp_tool( mcp_server=mcp_server, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py index 363948ff4e6..0e9a60c1586 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_hook_extra_headers.py @@ -515,6 +515,66 @@ class TestCallToolFlowsHookHeaders: proxy_logging_obj=proxy_logging, ) + @pytest.mark.asyncio + async def test_openapi_server_forwards_allowlisted_client_headers(self): + from mcp.types import TextContent + + from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( + _request_extra_headers, + ) + from litellm.proxy._experimental.mcp_server.tool_registry import ( + global_mcp_tool_registry, + ) + + manager = MCPServerManager() + server = MCPServer( + server_id="test-id", + name="openapi_server", + server_name="openapi_server", + url="https://example.com", + transport=MCPTransport.http, + auth_type=MCPAuth.none, + spec_path="/path/to/spec.yaml", + extra_headers=["Authorization", "X-Tenant-ID"], + ) + manager.registry[server.server_id] = server + manager.tool_name_to_mcp_server_name_mapping["test_tool"] = server.name + manager.tool_name_to_mcp_server_name_mapping["openapi_server-test_tool"] = ( + server.name + ) + + async def capture_headers() -> Optional[Dict[str, str]]: + return _request_extra_headers.get() + + registered_name = "openapi_server-test_tool" + global_mcp_tool_registry.register_tool( + name=registered_name, + description="test", + input_schema={}, + handler=capture_headers, + ) + + try: + forwarded_result = await manager.call_tool( + server_name="openapi_server", + name="test_tool", + arguments={}, + raw_headers={ + "x-litellm-api-key": "sk-proxy", + "authorization": "Bearer user-token", + "x-tenant-id": "tenant-001", + "x-unlisted": "not-forwarded", + }, + ) + finally: + global_mcp_tool_registry.tools.pop(registered_name) + + assert forwarded_result.isError is False + assert isinstance(forwarded_result.content[0], TextContent) + assert forwarded_result.content[0].text == ( + "{'Authorization': 'Bearer user-token', 'X-Tenant-ID': 'tenant-001'}" + ) + class TestHookHeaderMergePriority: """Tests that hook-provided headers have highest priority in _call_regular_mcp_tool."""