diff --git a/src/fastmcp/server/proxy.py b/src/fastmcp/server/proxy.py index 57ec5facf..8ce7dc5d2 100644 --- a/src/fastmcp/server/proxy.py +++ b/src/fastmcp/server/proxy.py @@ -1,5 +1,6 @@ from __future__ import annotations +import asyncio import inspect import warnings from collections.abc import Awaitable, Callable @@ -53,10 +54,44 @@ logger = get_logger(__name__) ClientFactoryT = Callable[[], Client] | Callable[[], Awaitable[Client]] +class _SynchronizedClientContext: + """Async context manager that synchronizes client initialization for stdio transports.""" + + def __init__(self, client: Client, semaphore: asyncio.Semaphore): + self.client = client + self.semaphore = semaphore + + async def __aenter__(self) -> Client: + # Acquire semaphore before client initialization + await self.semaphore.acquire() + try: + # Enter the client context + await self.client.__aenter__() + return self.client + except Exception: + # If client initialization fails, release semaphore + self.semaphore.release() + raise + + async def __aexit__(self, exc_type, exc_val, exc_tb): + try: + # Exit the client context + return await self.client.__aexit__(exc_type, exc_val, exc_tb) + finally: + # Always release semaphore + self.semaphore.release() + + class ProxyManagerMixin: """A mixin for proxy managers to provide a unified client retrieval method.""" client_factory: ClientFactoryT + _client_semaphore: asyncio.Semaphore | None = None + + def __init_subclass__(cls, **kwargs): + super().__init_subclass__(**kwargs) + # Initialize semaphore for stdio transport synchronization + cls._client_semaphore = asyncio.Semaphore(1) async def _get_client(self) -> Client: """Gets a client instance by calling the sync or async factory.""" @@ -65,6 +100,27 @@ class ProxyManagerMixin: client = await client return client + async def _get_client_with_context(self): + """Gets a client instance and returns an async context manager. + + For stdio transports, this synchronizes client initialization to prevent + race conditions when multiple parallel requests try to connect simultaneously. + """ + from fastmcp.client.transports import PythonStdioTransport, StdioTransport + + client = await self._get_client() + + # Check if this is a stdio transport that needs synchronization + transport = client.transport + needs_sync = isinstance(transport, PythonStdioTransport | StdioTransport) + + if needs_sync: + # Use semaphore to prevent concurrent stdio client initialization + return _SynchronizedClientContext(client, self._client_semaphore) + else: + # For non-stdio transports, use client directly + return client + class ProxyToolManager(ToolManager, ProxyManagerMixin): """A ToolManager that sources its tools from a remote client in addition to local and mounted tools.""" @@ -80,8 +136,8 @@ class ProxyToolManager(ToolManager, ProxyManagerMixin): # Then add proxy tools, but don't overwrite existing ones try: - client = await self._get_client() - async with client: + client_context = await self._get_client_with_context() + async with client_context as client: client_tools = await client.list_tools() for tool in client_tools: if tool.name not in all_tools: @@ -111,8 +167,8 @@ class ProxyToolManager(ToolManager, ProxyManagerMixin): return await super().call_tool(key, arguments) except NotFoundError: # If not found locally, try proxy - client = await self._get_client() - async with client: + client_context = await self._get_client_with_context() + async with client_context as client: result = await client.call_tool(key, arguments) return ToolResult( content=result.content, @@ -134,8 +190,8 @@ class ProxyResourceManager(ResourceManager, ProxyManagerMixin): # Then add proxy resources, but don't overwrite existing ones try: - client = await self._get_client() - async with client: + client_context = await self._get_client_with_context() + async with client_context as client: client_resources = await client.list_resources() for resource in client_resources: if str(resource.uri) not in all_resources: @@ -157,8 +213,8 @@ class ProxyResourceManager(ResourceManager, ProxyManagerMixin): # Then add proxy templates, but don't overwrite existing ones try: - client = await self._get_client() - async with client: + client_context = await self._get_client_with_context() + async with client_context as client: client_templates = await client.list_resource_templates() for template in client_templates: if template.uriTemplate not in all_templates: @@ -190,8 +246,8 @@ class ProxyResourceManager(ResourceManager, ProxyManagerMixin): return await super().read_resource(uri) except NotFoundError: # If not found locally, try proxy - client = await self._get_client() - async with client: + client_context = await self._get_client_with_context() + async with client_context as client: result = await client.read_resource(uri) if isinstance(result[0], TextResourceContents): return result[0].text @@ -215,8 +271,8 @@ class ProxyPromptManager(PromptManager, ProxyManagerMixin): # Then add proxy prompts, but don't overwrite existing ones try: - client = await self._get_client() - async with client: + client_context = await self._get_client_with_context() + async with client_context as client: client_prompts = await client.list_prompts() for prompt in client_prompts: if prompt.name not in all_prompts: @@ -247,8 +303,8 @@ class ProxyPromptManager(PromptManager, ProxyManagerMixin): return await super().render_prompt(name, arguments) except NotFoundError: # If not found locally, try proxy - client = await self._get_client() - async with client: + client_context = await self._get_client_with_context() + async with client_context as client: result = await client.get_prompt(name, arguments) return result diff --git a/tests/server/proxy/test_proxy_stdio_race_fix.py b/tests/server/proxy/test_proxy_stdio_race_fix.py new file mode 100644 index 000000000..0b6707b78 --- /dev/null +++ b/tests/server/proxy/test_proxy_stdio_race_fix.py @@ -0,0 +1,128 @@ +"""Test for stdio proxy race condition fix.""" + +import asyncio +import inspect +import tempfile +from pathlib import Path + +from fastmcp import FastMCP +from fastmcp.client import Client +from fastmcp.client.transports import PythonStdioTransport + + +class TestProxyStdioRaceFix: + """Test that the proxy stdio race condition is fixed.""" + + async def test_proxy_parallel_calls_no_race_condition(self): + """Test that parallel calls through stdio proxy don't fail with race conditions.""" + + # Create a temporary script for the backend server + server_script = inspect.cleandoc(""" + from fastmcp import FastMCP + + mcp = FastMCP() + + @mcp.tool + def add(a: int, b: int) -> int: + return a + b + + if __name__ == '__main__': + mcp.run() + """) + + with tempfile.TemporaryDirectory() as tmp_dir: + script_path = Path(tmp_dir) / "test.py" + script_path.write_text(server_script) + + # Set up the backend client (stdio transport) + backend_client = Client( + transport=PythonStdioTransport(script_path=script_path) + ) + + # Create proxy server + proxy = FastMCP.as_proxy(backend=backend_client, name="test_parallel_calls") + + # Create client that connects to the proxy + client = Client(transport=proxy) + + # Test with enough parallel calls to trigger race condition + count = 20 + + tasks = [client.list_tools() for _ in range(count)] + + async with backend_client, client: + results = await asyncio.gather(*tasks, return_exceptions=True) + + # All calls should succeed + exceptions = [result for result in results if isinstance(result, Exception)] + successes = [ + result for result in results if not isinstance(result, Exception) + ] + + assert len(exceptions) == 0, ( + f"Got {len(exceptions)} exceptions: {exceptions}" + ) + assert len(successes) == count + assert all( + len(result) == 1 for result in successes + ) # Each should have 1 tool (add) + + async def test_proxy_parallel_tool_calls_no_race_condition(self): + """Test that parallel tool calls through stdio proxy don't fail with race conditions.""" + + # Create a temporary script for the backend server + server_script = inspect.cleandoc(""" + from fastmcp import FastMCP + + mcp = FastMCP() + + @mcp.tool + def add(a: int, b: int) -> int: + return a + b + + if __name__ == '__main__': + mcp.run() + """) + + with tempfile.TemporaryDirectory() as tmp_dir: + script_path = Path(tmp_dir) / "test.py" + script_path.write_text(server_script) + + # Set up the backend client (stdio transport) + backend_client = Client( + transport=PythonStdioTransport(script_path=script_path) + ) + + # Create proxy server + proxy = FastMCP.as_proxy(backend=backend_client, name="test_parallel_calls") + + # Create client that connects to the proxy + client = Client(transport=proxy) + + # Test with parallel tool calls + count = 15 + + tasks = [ + client.call_tool("add", {"a": i, "b": i + 1}) for i in range(count) + ] + + async with backend_client, client: + results = await asyncio.gather(*tasks, return_exceptions=True) + + # All calls should succeed + exceptions = [result for result in results if isinstance(result, Exception)] + successes = [ + result for result in results if not isinstance(result, Exception) + ] + + assert len(exceptions) == 0, ( + f"Got {len(exceptions)} exceptions: {exceptions}" + ) + assert len(successes) == count + + # Verify results are correct + for i, result in enumerate(successes): + expected = i + (i + 1) # a + b where a=i, b=i+1 + assert result.data == expected, ( + f"Expected {expected}, got {result.data}" + )