fix(mcp): forward OpenAPI extra headers on call_tool

This commit is contained in:
Devin AI 2026-07-16 02:11:39 +00:00
parent 7a4a68f022
commit 60c3a418f8
2 changed files with 106 additions and 1 deletions

View file

@ -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,

View file

@ -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."""