diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 251f271903b..dd816c8f468 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -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: diff --git a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py index 3b2fa097b70..f4b9fd8c779 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -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 diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index c2e998f01e5..f2f6f672edd 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -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)