fix(mcp): forward extra_headers for OpenAPI-backed MCP servers

When an MCP server is backed by an OpenAPI spec (spec_path),
extra_headers (header names to forward from caller requests to the
upstream API) were silently ignored. Only managed MCP servers forwarded
them correctly.

Root cause: The OpenAPI tool function closure bakes headers at
registration time and had no mechanism to receive per-request headers.

Fix: Add a _request_extra_headers ContextVar (analogous to the existing
_request_auth_header) that is set before invoking the tool handler and
merged into effective_headers inside the closure. Both the local
dispatch path in server.py and the managed dispatch path in
mcp_server_manager._call_openapi_tool_handler now build the extra
headers dict from raw_headers using the same normalization pattern as
_call_regular_mcp_tool, set the ContextVar, and reset it in a finally
block.

Fixes #26794

Made-with: Cursor
This commit is contained in:
Raj Nagulapalle 2026-04-29 14:00:08 -07:00
parent 295a36aa69
commit 22000ffecc
3 changed files with 60 additions and 12 deletions

View file

@ -2037,6 +2037,7 @@ class MCPServerManager:
server: MCPServer,
tool_name: str,
arguments: Dict[str, Any],
raw_headers: Optional[Dict[str, str]] = None,
) -> CallToolResult:
"""
Call an OpenAPI tool handler directly.
@ -2046,14 +2047,20 @@ class MCPServerManager:
HTTP requests to the API.
Args:
server: The MCPServer configuration.
tool_name: The full tool name (with prefix) to call
arguments: Tool arguments to pass to the handler
raw_headers: Optional raw headers from the inbound request, used to
extract headers listed in ``server.extra_headers``.
Returns:
CallToolResult with the response from the API
"""
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,
)
@ -2069,12 +2076,27 @@ class MCPServerManager:
isError=True,
)
# Build extra_headers dict from raw_headers using the same
# normalization pattern as _call_regular_mcp_tool.
extra_headers_dict: Optional[Dict[str, str]] = None
if server.extra_headers and raw_headers:
normalized = {
str(k).lower(): v
for k, v in raw_headers.items()
if isinstance(k, str)
}
extra_headers_dict = {}
for header in server.extra_headers:
if not isinstance(header, str):
continue
val = normalized.get(header.lower())
if val is not None:
extra_headers_dict[header] = val
_extra_token = _request_extra_headers.set(extra_headers_dict)
try:
# Call the tool handler with the arguments
# The handler is an async function that makes the HTTP request
handler_result = await tool.handler(**arguments)
# Convert the handler result (string response) to CallToolResult format
result = CallToolResult(
content=[TextContent(type="text", text=str(handler_result))],
isError=False,
@ -2089,6 +2111,8 @@ class MCPServerManager:
content=[TextContent(type="text", text=error_msg)],
isError=True,
)
finally:
_request_extra_headers.reset(_extra_token)
async def pre_call_tool_check(
self,
@ -2506,17 +2530,11 @@ class MCPServerManager:
verbose_logger.debug(
"Calling OpenAPI tool %s directly via HTTP handler", name
)
if hook_result.get("extra_headers"):
verbose_logger.warning(
"pre_mcp_call hook returned extra_headers for OpenAPI-backed "
"MCP server '%s' — header injection is not supported for "
"OpenAPI servers; headers will be ignored. Use SSE/HTTP "
"transport to enable hook header injection.",
server_name,
)
tasks.append(
asyncio.create_task(
self._call_openapi_tool_handler(mcp_server, name, arguments)
self._call_openapi_tool_handler(
mcp_server, name, arguments, raw_headers=raw_headers
)
)
)
else:

View file

@ -31,6 +31,13 @@ _request_auth_header: contextvars.ContextVar[Optional[str]] = contextvars.Contex
"_request_auth_header", default=None
)
# Per-request extra headers override for OpenAPI-backed MCP servers.
# Set this ContextVar before calling a local tool handler to inject headers
# listed in the server's ``extra_headers`` config into the upstream HTTP request.
_request_extra_headers: contextvars.ContextVar[Optional[Dict[str, str]]] = (
contextvars.ContextVar("_request_extra_headers", default=None)
)
def _sanitize_path_parameter_value(param_value: Any, param_name: str) -> str:
"""Ensure path params cannot introduce directory traversal."""
@ -315,6 +322,11 @@ def create_tool_function(
# correct prefix (Bearer / ApiKey / Basic) formatted by the caller in
# server.py based on the server's configured auth_type.
effective_headers = dict(headers)
request_extra = _request_extra_headers.get()
if request_extra:
effective_headers.update(request_extra)
override_auth = _request_auth_header.get()
if override_auth:
effective_headers["Authorization"] = override_auth

View file

@ -156,6 +156,7 @@ if MCP_AVAILABLE:
)
from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import (
_request_auth_header,
_request_extra_headers,
)
from litellm.proxy._experimental.mcp_server.sse_transport import SseServerTransport
from litellm.proxy._experimental.mcp_server.tool_registry import (
@ -2123,11 +2124,28 @@ if MCP_AVAILABLE:
auth_header_value = f"Basic {mcp_auth_header}"
else:
auth_header_value = f"Bearer {mcp_auth_header}"
extra_headers_dict: Optional[Dict[str, str]] = None
if mcp_server and mcp_server.extra_headers and raw_headers:
normalized = {
str(k).lower(): v
for k, v in raw_headers.items()
if isinstance(k, str)
}
extra_headers_dict = {}
for header in mcp_server.extra_headers:
if not isinstance(header, str):
continue
val = normalized.get(header.lower())
if val is not None:
extra_headers_dict[header] = val
_extra_token = _request_extra_headers.set(extra_headers_dict)
_auth_token = _request_auth_header.set(auth_header_value)
try:
local_content = await _handle_local_mcp_tool(name, arguments)
finally:
_request_auth_header.reset(_auth_token)
_request_extra_headers.reset(_extra_token)
response = CallToolResult(content=cast(Any, local_content), isError=False)
# Try managed MCP server tool (pass the full prefixed name)