mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
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:
parent
295a36aa69
commit
22000ffecc
3 changed files with 60 additions and 12 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue