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..4891e24bfae 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,23 @@ 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 _build_effective_headers(headers: Dict[str, str]) -> Dict[str, str]: + """Merge static OpenAPI headers with request-scoped MCP header overrides.""" + 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 _sanitize_path_parameter_value(param_value: Any, param_name: str) -> 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 = _build_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 ae6055217b8..0659ac43654 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -157,6 +157,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 ( @@ -1124,6 +1125,29 @@ if MCP_AVAILABLE: return server_auth_header, extra_headers + def _prepare_local_mcp_request_extra_headers( + server: Optional[MCPServer], + raw_headers: Optional[Dict[str, str]], + ) -> Optional[Dict[str, str]]: + """Build request-time extra headers for local OpenAPI-generated tools.""" + if not server or not server.extra_headers or not raw_headers: + return None + + normalized_raw_headers = { + str(k).lower(): v for k, v in raw_headers.items() if isinstance(k, str) + } + extra_headers: Dict[str, 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 not None: + extra_headers[header] = header_value + + return extra_headers or None + def _merge_gateway_initialize_instructions( allowed_mcp_servers: List[MCPServer], ) -> Optional[str]: @@ -2120,11 +2144,17 @@ if MCP_AVAILABLE: auth_header_value = f"Basic {mcp_auth_header}" else: auth_header_value = f"Bearer {mcp_auth_header}" + request_extra_headers = _prepare_local_mcp_request_extra_headers( + server=mcp_server, + 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_auth_header.reset(_auth_token) + _request_extra_headers.reset(_extra_headers_token) response = CallToolResult(content=cast(Any, local_content), isError=False) # Try managed MCP server tool (pass the full prefixed name) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 9df6408b0d7..5ecc9c004aa 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -135,6 +135,36 @@ def test_prepare_mcp_server_headers_case_insensitive_extra_headers(): assert extra_headers == {"Authorization": "Bearer token"} +def test_prepare_local_mcp_request_extra_headers_case_insensitive(): + try: + from litellm.proxy._experimental.mcp_server.server import ( + _prepare_local_mcp_request_extra_headers, + ) + except ImportError: + pytest.skip("MCP server not available") + + server = MCPServer( + server_id="server-case", + name="server", + transport=MCPTransport.http, + extra_headers=["X-TOKEN", "X-API-Key"], + ) + + extra_headers = _prepare_local_mcp_request_extra_headers( + server=server, + raw_headers={ + "x-token": "request-token", + "X-API-KEY": "request-api-key", + "x-litellm-api-key": "litellm-key", + }, + ) + + assert extra_headers == { + "X-TOKEN": "request-token", + "X-API-Key": "request-api-key", + } + + @pytest.mark.asyncio async def test_get_prompts_from_mcp_servers_success(): try: 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..68246474c9d 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, @@ -369,6 +371,69 @@ class TestCreateToolFunction: # Should have no exec() calls assert len(exec_calls) == 0, "create_tool_function should not use exec()" + @pytest.mark.asyncio + async def test_request_extra_headers_are_forwarded(self): + """Test request-time extra headers are forwarded to OpenAPI requests.""" + func = create_tool_function( + path="/protected", + method="get", + operation={}, + base_url="https://api.example.com", + headers={"X-Static": "static-value"}, + ) + + 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-value", + "X-TOKEN": "request-token", + } + + @pytest.mark.asyncio + async def test_auth_override_wins_over_request_extra_headers(self): + """Test x-mcp-auth Authorization override preserves existing precedence.""" + func = create_tool_function( + path="/protected", + method="get", + operation={}, + base_url="https://api.example.com", + headers={"Authorization": "Bearer static-token"}, + ) + + 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 mcp-auth-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 mcp-auth-token", + "X-TOKEN": "request-token", + } + class TestBuildInputSchema: """Test that build_input_schema preserves original parameter names."""