diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 9923c3ce4bf..6f80cd9219a 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -2264,6 +2264,8 @@ class MCPServerManager: server: MCPServer, tool_name: str, arguments: Dict[str, Any], + mcp_auth_header: Optional[str] = None, + request_extra_headers: Optional[Dict[str, str]] = None, ) -> CallToolResult: """ Call an OpenAPI tool handler directly. @@ -2281,6 +2283,10 @@ class MCPServerManager: """ from mcp.types import TextContent + from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( + _request_auth_header, + _request_extra_headers, + ) from litellm.proxy._experimental.mcp_server.tool_registry import ( global_mcp_tool_registry, ) @@ -2297,9 +2303,24 @@ class MCPServerManager: ) try: + auth_header_value: Optional[str] = None + if mcp_auth_header: + if server.auth_type == MCPAuth.api_key: + auth_header_value = f"ApiKey {mcp_auth_header}" + elif server.auth_type == MCPAuth.basic: + auth_header_value = f"Basic {mcp_auth_header}" + else: + auth_header_value = f"Bearer {mcp_auth_header}" + + auth_token = _request_auth_header.set(auth_header_value) + extra_headers_token = _request_extra_headers.set(request_extra_headers) # Call the tool handler with the arguments # The handler is an async function that makes the HTTP request - handler_result = await tool.handler(**arguments) + try: + handler_result = await tool.handler(**arguments) + finally: + _request_extra_headers.reset(extra_headers_token) + _request_auth_header.reset(auth_token) # Convert the handler result (string response) to CallToolResult format result = CallToolResult( @@ -2474,6 +2495,50 @@ class MCPServerManager: ) ) + def _build_openapi_request_extra_headers( + self, + mcp_server: MCPServer, + oauth2_headers: Optional[Dict[str, str]], + raw_headers: Optional[Dict[str, str]], + hook_extra_headers: Optional[Dict[str, str]], + ) -> Optional[Dict[str, str]]: + """Build per-request headers for OpenAPI-generated MCP tool handlers.""" + extra_headers = oauth2_headers.copy() if oauth2_headers else None + + if mcp_server.extra_headers and raw_headers: + if extra_headers is None: + extra_headers = {} + + normalized_raw_headers = { + str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str) + } + for header in mcp_server.extra_headers: + if not isinstance(header, str): + continue + if ( + mcp_server.has_client_credentials + and header.lower() == "authorization" + ): + continue + header_value = normalized_raw_headers.get(header.lower()) + if header_value is None: + continue + extra_headers[header] = header_value + + if mcp_server.static_headers: + if extra_headers is None: + extra_headers = {} + extra_headers.update(mcp_server.static_headers) + + if hook_extra_headers: + if extra_headers is None: + extra_headers = {} + extra_headers.update(hook_extra_headers) + + if extra_headers is not None and len(extra_headers) == 0: + return None + return extra_headers + async def _call_regular_mcp_tool( # noqa: PLR0915 self, mcp_server: MCPServer, @@ -2746,17 +2811,21 @@ 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, - ) + request_extra_headers = self._build_openapi_request_extra_headers( + mcp_server=mcp_server, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + hook_extra_headers=hook_result.get("extra_headers"), + ) tasks.append( asyncio.create_task( - self._call_openapi_tool_handler(mcp_server, name, arguments) + self._call_openapi_tool_handler( + mcp_server, + name, + arguments, + mcp_auth_header=mcp_auth_header, + request_extra_headers=request_extra_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..fcc5a3a8ea4 100644 --- a/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py +++ b/litellm/proxy/_experimental/mcp_server/openapi_to_mcp_generator.py @@ -30,6 +30,9 @@ HEADERS: Dict[str, str] = {} _request_auth_header: contextvars.ContextVar[Optional[str]] = contextvars.ContextVar( "_request_auth_header", default=None ) +_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: @@ -273,6 +276,20 @@ def build_input_schema(operation: Dict[str, Any]) -> Dict[str, Any]: } +def _get_effective_headers(headers: Dict[str, str]) -> Dict[str, str]: + effective_headers = dict(headers) + + request_extra_headers = _request_extra_headers.get() + if request_extra_headers: + effective_headers.update(request_extra_headers) + + override_auth = _request_auth_header.get() + if override_auth: + effective_headers["Authorization"] = override_auth + + return effective_headers + + def create_tool_function( path: str, method: str, @@ -314,10 +331,7 @@ def create_tool_function( # The ContextVar holds the full Authorization header value, including the # correct prefix (Bearer / ApiKey / Basic) formatted by the caller in # server.py based on the server's configured auth_type. - effective_headers = dict(headers) - override_auth = _request_auth_header.get() - if override_auth: - effective_headers["Authorization"] = override_auth + effective_headers = _get_effective_headers(headers) # Build URL from base_url and path url = base_url + path diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index abb4b5cfa6f..e3a700fe67f 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -158,6 +158,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 ( @@ -1150,6 +1151,32 @@ if MCP_AVAILABLE: return server_auth_header, extra_headers + def _get_request_extra_headers_for_openapi_tool( + server: Optional[MCPServer], + oauth2_headers: Optional[Dict[str, str]], + raw_headers: Optional[Dict[str, str]], + ) -> Optional[Dict[str, str]]: + """Build per-request headers for a local OpenAPI-generated MCP tool.""" + extra_headers = oauth2_headers.copy() if oauth2_headers else None + + if server and server.extra_headers and raw_headers: + if extra_headers is None: + extra_headers = {} + + normalized_raw_headers = { + str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str) + } + + for header in server.extra_headers: + if not isinstance(header, str): + continue + header_value = normalized_raw_headers.get(header.lower()) + if header_value is None: + continue + extra_headers[header] = header_value + + return extra_headers + def _merge_gateway_initialize_instructions( allowed_mcp_servers: List[MCPServer], ) -> Optional[str]: @@ -2154,10 +2181,17 @@ if MCP_AVAILABLE: auth_header_value = f"Basic {mcp_auth_header}" else: auth_header_value = f"Bearer {mcp_auth_header}" + request_extra_headers = _get_request_extra_headers_for_openapi_tool( + server=mcp_server, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + ) _auth_token = _request_auth_header.set(auth_header_value) + _extra_headers_token = _request_extra_headers.set(request_extra_headers) try: local_content = await _handle_local_mcp_tool(name, arguments) finally: + _request_extra_headers.reset(_extra_headers_token) _request_auth_header.reset(_auth_token) response = CallToolResult(content=cast(Any, local_content), isError=False) 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 649a08e8744..a11a33b7bef 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 @@ -6,7 +6,7 @@ Validates that: 2. pre_call_tool_check returns hook-provided extra_headers AND modified arguments 3. call_tool flows hook headers and modified arguments downstream 4. Hook-provided headers take highest priority (merge after static_headers) -5. OpenAPI-backed servers log a warning and continue (skip injection) when hook headers are present +5. OpenAPI-backed servers merge hook headers into request-scoped generated-tool headers 6. JWT claims are propagated in both standard and virtual-key fast paths 7. Backward compatibility: hooks without extra_headers continue to work """ @@ -16,7 +16,6 @@ from typing import Any, Dict, Optional from unittest.mock import AsyncMock, MagicMock, patch import pytest -from fastapi import HTTPException from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager from litellm.proxy._types import UserAPIKeyAuth @@ -422,8 +421,8 @@ class TestCallToolFlowsHookHeaders: assert call_kwargs.kwargs.get("arguments") == modified_args @pytest.mark.asyncio - async def test_openapi_server_warns_and_continues_on_hook_headers(self): - """OpenAPI-backed servers log a warning and continue when hook injects headers.""" + async def test_openapi_server_forwards_hook_headers(self): + """OpenAPI-backed servers forward hook headers through request context.""" manager = MCPServerManager() server = MCPServer( server_id="test-id", @@ -454,24 +453,20 @@ class TestCallToolFlowsHookHeaders: "_call_openapi_tool_handler", new_callable=AsyncMock, return_value=MagicMock(), - ): - import litellm.proxy._experimental.mcp_server.mcp_server_manager as mgr_mod - + ) as mock_call: proxy_logging = MagicMock(spec=ProxyLogging) - with patch.object(mgr_mod, "verbose_logger") as mock_logger: - # Should NOT raise — just warn and proceed - await manager.call_tool( - server_name="openapi_server", - name="test_tool", - arguments={}, - proxy_logging_obj=proxy_logging, - ) - mock_logger.warning.assert_called_once() - assert ( - "header injection is not supported" - in mock_logger.warning.call_args[0][0] - ) + await manager.call_tool( + server_name="openapi_server", + name="test_tool", + arguments={}, + proxy_logging_obj=proxy_logging, + ) + + mock_call.assert_called_once() + assert mock_call.call_args.kwargs["request_extra_headers"] == { + "Authorization": "Bearer jwt" + } @pytest.mark.asyncio async def test_openapi_server_no_error_without_hook_headers(self): @@ -773,6 +768,28 @@ class TestHookHeaderMergePriority: assert "Authorization" not in headers assert headers.get("X-Custom") == "from-client" + def test_openapi_request_headers_merge_oauth_raw_and_hook_headers(self): + """OpenAPI tools receive the same runtime header sources as MCP transports.""" + manager = MCPServerManager() + server = self._make_server(extra_headers_config=["X-TOKEN", "X-Trace"]) + + headers = manager._build_openapi_request_extra_headers( + mcp_server=server, + oauth2_headers={"Authorization": "Bearer oauth-token"}, + raw_headers={ + "x-token": "request-token", + "x-trace": "trace-from-request", + "x-ignored": "not-forwarded", + }, + hook_extra_headers={"Authorization": "Bearer hook-token"}, + ) + + assert headers == { + "Authorization": "Bearer hook-token", + "X-TOKEN": "request-token", + "X-Trace": "trace-from-request", + } + class TestUserAPIKeyAuthJwtClaims: """Tests that UserAPIKeyAuth correctly carries jwt_claims.""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py index efe100a11dc..99b95daf403 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_openapi_to_mcp_generator.py @@ -17,6 +17,8 @@ import pytest from litellm.proxy._experimental.mcp_server.openapi_to_mcp_generator import ( _resolve_param_list, _resolve_ref, + _request_auth_header, + _request_extra_headers, build_input_schema, create_tool_function, extract_parameters, @@ -77,6 +79,67 @@ class TestCreateToolFunction: call_args[0][0] ) + @pytest.mark.asyncio + async def test_request_extra_headers_are_forwarded(self): + """OpenAPI tools should merge per-request header passthrough.""" + operation = {} + func = create_tool_function( + path="/protected", + method="get", + operation=operation, + base_url="https://api.example.com", + headers={"X-Static": "static"}, + ) + + with patch(GET_ASYNC_CLIENT_TARGET) as mock_client: + async_client = _create_mock_client("get", "ok") + mock_client.return_value = async_client + + token = _request_extra_headers.set({"X-TOKEN": "request-token"}) + try: + result = await func() + finally: + _request_extra_headers.reset(token) + + assert result == "ok" + call_args = async_client.get.call_args + assert call_args.kwargs["headers"] == { + "X-Static": "static", + "X-TOKEN": "request-token", + } + + @pytest.mark.asyncio + async def test_auth_context_overrides_request_extra_authorization_header(self): + """BYOK auth must keep highest precedence over generic request headers.""" + operation = {} + func = create_tool_function( + path="/protected", + method="get", + operation=operation, + base_url="https://api.example.com", + ) + + with patch(GET_ASYNC_CLIENT_TARGET) as mock_client: + async_client = _create_mock_client("get", "ok") + mock_client.return_value = async_client + + extra_token = _request_extra_headers.set( + {"Authorization": "Bearer request-token", "X-TOKEN": "request-token"} + ) + auth_token = _request_auth_header.set("Bearer byok-token") + try: + result = await func() + finally: + _request_auth_header.reset(auth_token) + _request_extra_headers.reset(extra_token) + + assert result == "ok" + call_args = async_client.get.call_args + assert call_args.kwargs["headers"] == { + "Authorization": "Bearer byok-token", + "X-TOKEN": "request-token", + } + @pytest.mark.asyncio async def test_leading_digit_parameter(self): """Test function with parameter starting with digit (e.g., 2fa-code)."""