From 34b500c7f5b24b6dd24c9c286081aa07c56b6d4b Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Mon, 22 Dec 2025 11:21:59 +0900 Subject: [PATCH] feat: support MCP stdio header env overrides --- docs/my-website/docs/mcp.md | 27 +++++- .../mcp_server/mcp_server_manager.py | 62 ++++++++++++- .../mcp_server/rest_endpoints.py | 58 ++++++++---- .../proxy/_experimental/mcp_server/server.py | 6 ++ tests/mcp_tests/test_mcp_server.py | 16 +++- .../mcp_server/test_mcp_server.py | 58 ++++++++++-- .../mcp_server/test_mcp_server_manager.py | 93 +++++++++++++++++-- .../mcp_server/test_rest_endpoints.py | 16 +++- 8 files changed, 295 insertions(+), 41 deletions(-) diff --git a/docs/my-website/docs/mcp.md b/docs/my-website/docs/mcp.md index f9c9cbb4562..a70e3d24188 100644 --- a/docs/my-website/docs/mcp.md +++ b/docs/my-website/docs/mcp.md @@ -746,8 +746,33 @@ curl --location 'http://localhost:4000/github_mcp/mcp' \ 3. **Header Forwarding**: LiteLLM automatically forwards matching headers to the backend MCP server 4. **Authentication**: The backend MCP server receives both the configured auth headers and the custom headers ---- +### Passing Request Headers to STDIO env Vars + +If your stdio MCP server needs per-request credentials, you can map HTTP headers from the client request directly into the environment for the launched stdio process. Reference the header name in the env value using the `${X-HEADER_NAME}` syntax. LiteLLM will read that header from the incoming request and set the env var before starting the command. + +```json title="Forward X-GITHUB_PERSONAL_ACCESS_TOKEN header to stdio env" showLineNumbers +{ + "mcpServers": { + "github": { + "command": "docker", + "args": [ + "run", + "-i", + "--rm", + "-e", + "GITHUB_PERSONAL_ACCESS_TOKEN", + "ghcr.io/github/github-mcp-server" + ], + "env": { + "GITHUB_PERSONAL_ACCESS_TOKEN": "${X-GITHUB_PERSONAL_ACCESS_TOKEN}" + } + } + } +} +``` + +In this example, when a client makes a request with the `X-GITHUB_PERSONAL_ACCESS_TOKEN` header, the proxy forwards that value into the stdio process as the `GITHUB_PERSONAL_ACCESS_TOKEN` environment variable. ## Using your MCP with client side credentials diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 8c9d8630457..2260649e8b2 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -84,6 +84,8 @@ def _deserialize_json_dict(data: Any) -> Optional[Dict[str, str]]: class MCPServerManager: + _STDIO_ENV_TEMPLATE_PATTERN = re.compile(r"^\$\{(X-[^}]+)\}$") + def __init__(self): self.registry: Dict[str, MCPServer] = {} self.config_mcp_servers: Dict[str, MCPServer] = {} @@ -671,11 +673,39 @@ class MCPServerManager: ######################################################### # Methods that call the upstream MCP servers ######################################################### + def _build_stdio_env( + self, + server: MCPServer, + raw_headers: Optional[Dict[str, str]] = None, + ) -> Optional[Dict[str, str]]: + """Resolve stdio env values, supporting header-driven placeholders.""" + + if server.transport != MCPTransport.stdio or not server.env: + return None + + resolved_env: Dict[str, str] = {} + normalized_headers = {k.lower(): v for k, v in (raw_headers or {}).items()} + + for env_key, env_value in server.env.items(): + stripped_value = env_value.strip() + match = self._STDIO_ENV_TEMPLATE_PATTERN.match(stripped_value) + if match: + header_name = match.group(1) + header_value = normalized_headers.get(header_name.lower()) + if header_value is None: + continue + resolved_env[env_key] = header_value + else: + resolved_env[env_key] = env_value + + return resolved_env + def _create_mcp_client( self, server: MCPServer, mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, extra_headers: Optional[Dict[str, str]] = None, + stdio_env: Optional[Dict[str, str]] = None, ) -> MCPClient: """ Create an MCPClient instance for the given server. @@ -692,10 +722,13 @@ class MCPServerManager: # Handle stdio transport if transport == MCPTransport.stdio: # For stdio, we need to get the stdio config from the server + resolved_env = stdio_env if stdio_env is not None else server.env or {} stdio_config: Optional[MCPStdioConfig] = None if server.command and server.args is not None: stdio_config = MCPStdioConfig( - command=server.command, args=server.args, env=server.env or {} + command=server.command, + args=server.args, + env=resolved_env, ) return MCPClient( @@ -725,6 +758,7 @@ class MCPServerManager: mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, extra_headers: Optional[Dict[str, str]] = None, add_prefix: bool = True, + raw_headers: Optional[Dict[str, str]] = None, ) -> List[MCPTool]: """ Helper method to get tools from a single MCP server with prefixed names. @@ -751,10 +785,13 @@ class MCPServerManager: extra_headers = {} extra_headers.update(server.static_headers) + stdio_env = self._build_stdio_env(server, raw_headers) + client = self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, + stdio_env=stdio_env, ) ## HANDLE OPENAPI TOOLS @@ -784,6 +821,7 @@ class MCPServerManager: mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, extra_headers: Optional[Dict[str, str]] = None, add_prefix: bool = True, + raw_headers: Optional[Dict[str, str]] = None, ) -> List[Prompt]: """ Helper method to get prompts from a single MCP server with prefixed names. @@ -807,10 +845,13 @@ class MCPServerManager: extra_headers = {} extra_headers.update(server.static_headers) + stdio_env = self._build_stdio_env(server, raw_headers) + client = self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, + stdio_env=stdio_env, ) prompts = await client.list_prompts() @@ -833,6 +874,7 @@ class MCPServerManager: mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, extra_headers: Optional[Dict[str, str]] = None, add_prefix: bool = True, + raw_headers: Optional[Dict[str, str]] = None, ) -> List[Resource]: """Fetch available resources from a single MCP server.""" @@ -847,10 +889,13 @@ class MCPServerManager: extra_headers = {} extra_headers.update(server.static_headers) + stdio_env = self._build_stdio_env(server, raw_headers) + client = self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, + stdio_env=stdio_env, ) resources = await client.list_resources() @@ -873,6 +918,7 @@ class MCPServerManager: mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, extra_headers: Optional[Dict[str, str]] = None, add_prefix: bool = True, + raw_headers: Optional[Dict[str, str]] = None, ) -> List[ResourceTemplate]: """Fetch available resource templates from a single MCP server.""" @@ -887,10 +933,13 @@ class MCPServerManager: extra_headers = {} extra_headers.update(server.static_headers) + stdio_env = self._build_stdio_env(server, raw_headers) + client = self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, + stdio_env=stdio_env, ) resource_templates = await client.list_resource_templates() @@ -913,6 +962,7 @@ class MCPServerManager: url: AnyUrl, mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, extra_headers: Optional[Dict[str, str]] = None, + raw_headers: Optional[Dict[str, str]] = None, ) -> ReadResourceResult: """Read resource contents from a specific MCP server.""" @@ -924,10 +974,13 @@ class MCPServerManager: extra_headers = {} extra_headers.update(server.static_headers) + stdio_env = self._build_stdio_env(server, raw_headers) + client = self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, + stdio_env=stdio_env, ) return await client.read_resource(url) @@ -939,6 +992,7 @@ class MCPServerManager: arguments: Optional[Dict[str, Any]] = None, mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, extra_headers: Optional[Dict[str, str]] = None, + raw_headers: Optional[Dict[str, str]] = None, ) -> GetPromptResult: """Fetch a specific prompt definition from a single MCP server.""" @@ -950,10 +1004,13 @@ class MCPServerManager: extra_headers = {} extra_headers.update(server.static_headers) + stdio_env = self._build_stdio_env(server, raw_headers) + client = self._create_mcp_client( server=server, mcp_auth_header=mcp_auth_header, extra_headers=extra_headers, + stdio_env=stdio_env, ) get_prompt_request_params = GetPromptRequestParams( @@ -1742,10 +1799,13 @@ class MCPServerManager: extra_headers = {} extra_headers.update(mcp_server.static_headers) + stdio_env = self._build_stdio_env(mcp_server, raw_headers) + client = self._create_mcp_client( server=mcp_server, mcp_auth_header=server_auth_header, extra_headers=extra_headers, + stdio_env=stdio_env, ) call_tool_params = MCPCallToolRequestParams( diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 032331ece02..891b52db7af 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -71,12 +71,17 @@ if MCP_AVAILABLE: for tool in tools ] - async def _get_tools_for_single_server(server, server_auth_header): + async def _get_tools_for_single_server( + server, + server_auth_header, + raw_headers: Optional[Dict[str, str]] = None, + ): """Helper function to get tools for a single server.""" tools = await global_mcp_server_manager._get_tools_from_server( server=server, mcp_auth_header=server_auth_header, add_prefix=False, + raw_headers=raw_headers, ) # Filter tools based on allowed_tools configuration @@ -122,6 +127,7 @@ if MCP_AVAILABLE: try: # Extract auth headers from request headers = request.headers + raw_headers_from_request = dict(headers) mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers( headers ) @@ -148,7 +154,7 @@ if MCP_AVAILABLE: try: list_tools_result = await _get_tools_for_single_server( - server, server_auth_header + server, server_auth_header, raw_headers_from_request ) except Exception as e: verbose_logger.exception( @@ -169,7 +175,7 @@ if MCP_AVAILABLE: try: tools_result = await _get_tools_for_single_server( - server, server_auth_header + server, server_auth_header, raw_headers_from_request ) list_tools_result.extend(tools_result) except Exception as e: @@ -232,13 +238,13 @@ if MCP_AVAILABLE: # but they weren't being extracted and passed to call_mcp_tool. # This fix ensures auth headers are properly extracted from the HTTP request # and passed through to the MCP server for authentication. + headers = request.headers + raw_headers_from_request = dict(headers) mcp_auth_header = MCPRequestHandler._get_mcp_auth_header_from_headers( - request.headers + headers ) mcp_server_auth_headers = ( - MCPRequestHandler._get_mcp_server_auth_headers_from_headers( - request.headers - ) + MCPRequestHandler._get_mcp_server_auth_headers_from_headers(headers) ) # Add extracted headers to data dict to pass to call_mcp_tool @@ -246,6 +252,7 @@ if MCP_AVAILABLE: data["mcp_auth_header"] = mcp_auth_header if mcp_server_auth_headers: data["mcp_server_auth_headers"] = mcp_server_auth_headers + data["raw_headers"] = raw_headers_from_request result = await call_mcp_tool(**data) return result @@ -300,6 +307,7 @@ if MCP_AVAILABLE: operation, mcp_auth_header: Optional[Union[str, Dict[str, str]]] = None, oauth2_headers: Optional[Dict[str, str]] = None, + raw_headers: Optional[Dict[str, str]] = None, ): """ Common helper to create MCP client, execute operation, and ensure proper cleanup. @@ -312,17 +320,27 @@ if MCP_AVAILABLE: Operation result or error response """ try: + server_model = MCPServer( + server_id=request.server_id or "", + name=request.alias or request.server_name or "", + url=request.url, + transport=request.transport, + auth_type=request.auth_type, + mcp_info=request.mcp_info, + command=request.command, + args=request.args, + env=request.env, + ) + + stdio_env = global_mcp_server_manager._build_stdio_env( + server_model, raw_headers + ) + client = global_mcp_server_manager._create_mcp_client( - server=MCPServer( - server_id=request.server_id or "", - name=request.alias or request.server_name or "", - url=request.url, - transport=request.transport, - auth_type=request.auth_type, - mcp_info=request.mcp_info, - ), + server=server_model, mcp_auth_header=mcp_auth_header, extra_headers=oauth2_headers, + stdio_env=stdio_env, ) return await operation(client) @@ -338,7 +356,8 @@ if MCP_AVAILABLE: @router.post("/test/connection") async def test_connection( - request: NewMCPServerRequest, + request: Request, + new_mcp_server_request: NewMCPServerRequest, ): """ Test if we can connect to the provided MCP server before adding it @@ -351,7 +370,11 @@ if MCP_AVAILABLE: await client.run_with_session(_noop) return {"status": "ok"} - return await _execute_with_mcp_client(request, _test_connection_operation) + return await _execute_with_mcp_client( + new_mcp_server_request, + _test_connection_operation, + raw_headers=dict(request.headers), + ) @router.post("/test/tools/list") async def test_tools_list( @@ -405,4 +428,5 @@ if MCP_AVAILABLE: _list_tools_operation, mcp_auth_header=mcp_auth_header, oauth2_headers=oauth2_headers, + raw_headers=dict(request.headers), ) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index bdff60c932b..e00fdbfb930 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -775,6 +775,7 @@ if MCP_AVAILABLE: mcp_auth_header=server_auth_header, extra_headers=extra_headers, add_prefix=add_prefix, + raw_headers=raw_headers, ) filtered_tools = filter_tools_by_allowed_tools(tools, server) @@ -854,6 +855,7 @@ if MCP_AVAILABLE: mcp_auth_header=server_auth_header, extra_headers=extra_headers, add_prefix=add_prefix, + raw_headers=raw_headers, ) all_prompts.extend(prompts) @@ -912,6 +914,7 @@ if MCP_AVAILABLE: mcp_auth_header=server_auth_header, extra_headers=extra_headers, add_prefix=add_prefix, + raw_headers=raw_headers, ) all_resources.extend(resources) @@ -969,6 +972,7 @@ if MCP_AVAILABLE: mcp_auth_header=server_auth_header, extra_headers=extra_headers, add_prefix=add_prefix, + raw_headers=raw_headers, ) ) all_resource_templates.extend(resource_templates) @@ -1392,6 +1396,7 @@ if MCP_AVAILABLE: arguments=arguments, mcp_auth_header=server_auth_header, extra_headers=extra_headers, + raw_headers=raw_headers, ) async def mcp_read_resource( @@ -1440,6 +1445,7 @@ if MCP_AVAILABLE: url=url, mcp_auth_header=server_auth_header, extra_headers=extra_headers, + raw_headers=raw_headers, ) def _get_standard_logging_mcp_tool_call( diff --git a/tests/mcp_tests/test_mcp_server.py b/tests/mcp_tests/test_mcp_server.py index d3112714a9c..9242dfc75f4 100644 --- a/tests/mcp_tests/test_mcp_server.py +++ b/tests/mcp_tests/test_mcp_server.py @@ -812,10 +812,19 @@ async def test_get_tools_from_mcp_servers(): return_value=["server1_id", "server2_id"] ) mock_manager_2.get_mcp_server_by_id = lambda server_id: mock_server_1 if server_id == "server1_id" else mock_server_2 + async def mock_get_tools_side_effect( + server, + mcp_auth_header=None, + extra_headers=None, + add_prefix=False, + raw_headers=None, + ): + if server.server_id == "server1_id": + return [mock_tool_1] + return [mock_tool_2] + mock_manager_2._get_tools_from_server = AsyncMock( - side_effect=lambda server, mcp_auth_header=None, extra_headers=None, add_prefix=False: ( - [mock_tool_1] if server.server_id == "server1_id" else [mock_tool_2] - ) + side_effect=mock_get_tools_side_effect ) with patch( @@ -1693,6 +1702,7 @@ async def test_get_tools_for_single_server(): server=mock_server, mcp_auth_header="Bearer test_token", add_prefix=False, + raw_headers=None, ) # Verify the result 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 4fc94000d61..a1fbddec586 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 @@ -294,6 +294,7 @@ async def test_mcp_get_prompt_success(): arguments={"foo": "bar"}, mcp_auth_header={"Authorization": "token"}, extra_headers={"X-Test": "1"}, + raw_headers=None, ) assert result is prompt_result @@ -349,6 +350,7 @@ async def test_mcp_read_resource_success(): url="https://example.com/resource", mcp_auth_header={"Authorization": "token"}, extra_headers={"X-Test": "1"}, + raw_headers=None, ) assert result is read_result @@ -428,7 +430,11 @@ async def test_get_tools_from_mcp_servers_continues_when_one_server_fails(): ) async def mock_get_tools_from_server( - server, mcp_auth_header=None, extra_headers=None, add_prefix=True + server, + mcp_auth_header=None, + extra_headers=None, + add_prefix=True, + raw_headers=None, ): if server.name == "working_server": # Working server returns tools @@ -524,7 +530,11 @@ async def test_get_tools_from_mcp_servers_handles_all_servers_failing(): ) async def mock_get_tools_from_server( - server, mcp_auth_header=None, extra_headers=None, add_prefix=True + server, + mcp_auth_header=None, + extra_headers=None, + add_prefix=True, + raw_headers=None, ): # All servers fail raise Exception(f"Server {server.name} connection failed") @@ -839,13 +849,19 @@ async def test_oauth2_headers_passed_to_mcp_client(): # This will capture the arguments passed to _create_mcp_client captured_client_args = {} - def mock_create_mcp_client(server, mcp_auth_header=None, extra_headers=None): + def mock_create_mcp_client( + server, + mcp_auth_header=None, + extra_headers=None, + stdio_env=None, + ): # Capture the arguments for verification captured_client_args.update( { "server": server, "mcp_auth_header": mcp_auth_header, "extra_headers": extra_headers, + "stdio_env": stdio_env, } ) # Return a mock client that doesn't actually connect @@ -934,7 +950,11 @@ async def test_list_tools_single_server_unprefixed_names(): mock_manager.get_mcp_server_by_id = MagicMock(return_value=server) async def mock_get_tools_from_server( - server, mcp_auth_header=None, extra_headers=None, add_prefix=False + server, + mcp_auth_header=None, + extra_headers=None, + add_prefix=False, + raw_headers=None, ): tool = MagicMock() tool.name = f"{server.alias}-toolA" if add_prefix else "toolA" @@ -1006,7 +1026,11 @@ async def test_list_tools_multiple_servers_prefixed_names(): ) async def mock_get_tools_from_server( - server, mcp_auth_header=None, extra_headers=None, add_prefix=True + server, + mcp_auth_header=None, + extra_headers=None, + add_prefix=True, + raw_headers=None, ): tool = MagicMock() # When multiple servers, add_prefix should be True -> prefixed names @@ -1147,7 +1171,11 @@ async def test_list_tools_filters_by_key_team_permissions(): mock_manager.get_mcp_server_by_id = lambda server_id: server async def mock_get_tools_from_server( - server, mcp_auth_header=None, extra_headers=None, add_prefix=False + server, + mcp_auth_header=None, + extra_headers=None, + add_prefix=False, + raw_headers=None, ): # Return 4 tools, but only 2 should be allowed tool1 = MagicMock() @@ -1248,7 +1276,11 @@ async def test_list_tools_with_team_tool_permissions_inheritance(): mock_manager.get_mcp_server_by_id = lambda server_id: server async def mock_get_tools_from_server( - server, mcp_auth_header=None, extra_headers=None, add_prefix=False + server, + mcp_auth_header=None, + extra_headers=None, + add_prefix=False, + raw_headers=None, ): # Return 4 tools tool1 = MagicMock() @@ -1334,7 +1366,11 @@ async def test_list_tools_with_no_tool_permissions_shows_all(): mock_manager.get_mcp_server_by_id = lambda server_id: server async def mock_get_tools_from_server( - server, mcp_auth_header=None, extra_headers=None, add_prefix=False + server, + mcp_auth_header=None, + extra_headers=None, + add_prefix=False, + raw_headers=None, ): # Return 3 tools tool1 = MagicMock() @@ -1423,7 +1459,11 @@ async def test_list_tools_strips_prefix_when_matching_permissions(): mock_manager.get_mcp_server_by_id = MagicMock(return_value=server) async def mock_get_tools_from_server( - server, mcp_auth_header=None, extra_headers=None, add_prefix=True + server, + mcp_auth_header=None, + extra_headers=None, + add_prefix=True, + raw_headers=None, ): # Return tools WITH prefix (as they come from MCP server) tool1 = MagicMock() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 7a6e5ad17f6..ff016a1a130 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -8,6 +8,7 @@ from fastapi import HTTPException # Add the parent directory to the path so we can import litellm sys.path.insert(0, "../../../../../") + import httpx from mcp import ReadResourceResult, Resource from mcp.types import ( @@ -99,6 +100,53 @@ class TestMCPServerManager: assert client.stdio_config["args"] == ["server.js"] assert client.stdio_config["env"] == {"NODE_ENV": "test"} + def test_build_stdio_env_only_accepts_x_prefixed_placeholders(self): + """Ensure only ${X-*} placeholders are substituted from headers.""" + manager = MCPServerManager() + server = MCPServer( + server_id="stdio-server-env", + name="stdio_env", + transport=MCPTransport.stdio, + command="node", + args=["server.js"], + env={ + "PASSTHROUGH": "${X-Test-Header}", + "STATIC": "value", + "IGNORED": "${Not-Allowed}", + }, + ) + + env = manager._build_stdio_env( + server, + raw_headers={ + "x-test-header": "resolved-value", + "x-not-used": "other", + }, + ) + + assert env == { + "PASSTHROUGH": "resolved-value", + "STATIC": "value", + "IGNORED": "${Not-Allowed}", + } + + def test_build_stdio_env_missing_header_skips_entry(self): + """Ensure missing headers drop the placeholder from the resolved env.""" + manager = MCPServerManager() + server = MCPServer( + server_id="stdio-server-env-miss", + name="stdio_env_miss", + transport=MCPTransport.stdio, + command="node", + args=["server.js"], + env={"EXPECTED": "${X-Missing}"}, + ) + + env = manager._build_stdio_env(server, raw_headers={}) + + # When the header isn't provided, the key is omitted entirely + assert env == {} + @pytest.mark.asyncio async def test_list_tools_with_server_specific_auth_headers(self): """Test list_tools method with server-specific auth headers""" @@ -123,7 +171,10 @@ class TestMCPServerManager: # Mock _get_tools_from_server to return different results async def mock_get_tools_from_server( - server, mcp_auth_header=None, mcp_protocol_version=None + server, + mcp_auth_header=None, + mcp_protocol_version=None, + raw_headers=None, ): if server.name == "github": tool1 = MagicMock() @@ -174,7 +225,10 @@ class TestMCPServerManager: # Mock _get_tools_from_server async def mock_get_tools_from_server( - server, mcp_auth_header=None, mcp_protocol_version=None + server, + mcp_auth_header=None, + mcp_protocol_version=None, + raw_headers=None, ): assert mcp_auth_header == "legacy-token" # Should use legacy header tool = MagicMock() @@ -209,7 +263,10 @@ class TestMCPServerManager: # Mock _get_tools_from_server async def mock_get_tools_from_server( - server, mcp_auth_header=None, mcp_protocol_version=None + server, + mcp_auth_header=None, + mcp_protocol_version=None, + raw_headers=None, ): assert ( mcp_auth_header == "server-specific-token" @@ -373,6 +430,7 @@ class TestMCPServerManager: server=server, mcp_auth_header="auth", extra_headers=None, + stdio_env=None, ) mock_client.list_resource_templates.assert_awaited_once() mock_prefix.assert_called_once_with(mock_templates, server, add_prefix=False) @@ -554,7 +612,10 @@ class TestMCPServerManager: # Mock _get_tools_from_server async def mock_get_tools_from_server( - server, mcp_auth_header=None, mcp_protocol_version=None + server, + mcp_auth_header=None, + mcp_protocol_version=None, + raw_headers=None, ): assert ( mcp_auth_header == "server-specific-token" @@ -587,7 +648,11 @@ class TestMCPServerManager: manager.get_mcp_server_by_id = MagicMock(return_value=server) # Mock successful _get_tools_from_server - async def mock_get_tools_from_server(server, mcp_auth_header=None): + async def mock_get_tools_from_server( + server, + mcp_auth_header=None, + raw_headers=None, + ): tool1 = MagicMock() tool1.name = "tool1" tool2 = MagicMock() @@ -621,7 +686,11 @@ class TestMCPServerManager: manager.get_mcp_server_by_id = MagicMock(return_value=server) # Mock failed _get_tools_from_server - async def mock_get_tools_from_server(server, mcp_auth_header=None): + async def mock_get_tools_from_server( + server, + mcp_auth_header=None, + raw_headers=None, + ): raise Exception("Connection timeout") manager._get_tools_from_server = mock_get_tools_from_server @@ -683,7 +752,11 @@ class TestMCPServerManager: manager.get_mcp_server_by_id = mock_get_server_by_id # Mock _get_tools_from_server with different results - async def mock_get_tools_from_server(server, mcp_auth_header=None): + async def mock_get_tools_from_server( + server, + mcp_auth_header=None, + raw_headers=None, + ): if server.server_id == "server1": tool = MagicMock() tool.name = "tool1" @@ -724,7 +797,11 @@ class TestMCPServerManager: manager.get_mcp_server_by_id = MagicMock(return_value=server) # Mock _get_tools_from_server to verify auth header is passed - async def mock_get_tools_from_server(server, mcp_auth_header=None): + async def mock_get_tools_from_server( + server, + mcp_auth_header=None, + raw_headers=None, + ): assert mcp_auth_header == "test-token" tool = MagicMock() tool.name = "tool1" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index a0c09663a88..ce38ee59e4d 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -37,7 +37,13 @@ async def test_test_tools_list_forwards_mcp_auth_header(monkeypatch): captured: dict = {} - async def fake_execute(request, operation, mcp_auth_header=None, oauth2_headers=None): + async def fake_execute( + request, + operation, + mcp_auth_header=None, + oauth2_headers=None, + raw_headers=None, + ): captured["mcp_auth_header"] = mcp_auth_header captured["oauth2_headers"] = oauth2_headers return { @@ -87,7 +93,13 @@ async def test_test_tools_list_extracts_oauth2_headers(monkeypatch): captured: dict = {} - async def fake_execute(request, operation, mcp_auth_header=None, oauth2_headers=None): + async def fake_execute( + request, + operation, + mcp_auth_header=None, + oauth2_headers=None, + raw_headers=None, + ): captured["mcp_auth_header"] = mcp_auth_header captured["oauth2_headers"] = oauth2_headers return {