diff --git a/src/fastmcp/client/client.py b/src/fastmcp/client/client.py index 641d34e45..b908fbbd7 100644 --- a/src/fastmcp/client/client.py +++ b/src/fastmcp/client/client.py @@ -1,3 +1,4 @@ +import asyncio import datetime from contextlib import AsyncExitStack, asynccontextmanager from pathlib import Path @@ -152,6 +153,10 @@ class Client(Generic[ClientTransportT]): self._session: ClientSession | None = None self._exit_stack: AsyncExitStack | None = None self._nesting_counter: int = 0 + self._context_lock = anyio.Lock() + self._session_task: asyncio.Task | None = None + self._ready_event = asyncio.Event() + self._stop_event = asyncio.Event() self._initialize_result: mcp.types.InitializeResult | None = None if log_handler is None: @@ -191,6 +196,7 @@ class Client(Generic[ClientTransportT]): self._session_kwargs["sampling_callback"] = create_sampling_callback( sampling_handler ) + # self._session_manager = self._context_manager() @property def session(self) -> ClientSession: @@ -242,34 +248,45 @@ class Client(Generic[ClientTransportT]): except TimeoutError: raise RuntimeError("Failed to initialize server session") finally: - self._exit_stack = None self._session = None self._initialize_result = None async def __aenter__(self): - if self._nesting_counter == 0: - # Create exit stack to manage both context managers - stack = AsyncExitStack() - await stack.__aenter__() - - await stack.enter_async_context(self._context_manager()) - - self._exit_stack = stack - - self._nesting_counter += 1 - + async with self._context_lock: + need_to_start = self._session_task is None or self._session_task.done() + if need_to_start: + self._stop_event = anyio.Event() + self._ready_event = anyio.Event() + self._session_task = asyncio.create_task(self._session_runner()) + await self._ready_event.wait() + self._nesting_counter += 1 return self async def __aexit__(self, exc_type, exc_val, exc_tb): - self._nesting_counter -= 1 + async with self._context_lock: + self._nesting_counter -= 1 + if self._nesting_counter != 0: + return + self._stop_event.set() + runner_task = self._session_task + self._session_task = None + if runner_task: + await runner_task + # Reset for future reconnects + self._stop_event = anyio.Event() + self._ready_event = anyio.Event() - if self._nesting_counter == 0: - # Exit the stack which will handle cleaning up the session - if self._exit_stack is not None: - try: - await self._exit_stack.__aexit__(exc_type, exc_val, exc_tb) - finally: - self._exit_stack = None + async def _session_runner(self): + async with AsyncExitStack() as stack: + try: + await stack.enter_async_context(self._context_manager()) + # Session/context is now ready + self._ready_event.set() + # Wait until disconnect/stop is requested + await self._stop_event.wait() + finally: + # On exit, ensure ready event is set (idempotent) + self._ready_event.set() async def close(self): await self.transport.close() diff --git a/tests/server/test_proxy.py b/tests/server/test_proxy.py index f310fa505..22fe3415b 100644 --- a/tests/server/test_proxy.py +++ b/tests/server/test_proxy.py @@ -3,6 +3,7 @@ from typing import Any import mcp.types import pytest +from anyio import create_task_group from dirty_equals import Contains from mcp import McpError @@ -242,3 +243,26 @@ class TestPrompts: assert result.messages[0].role == "user" assert isinstance(result.messages[0].content, mcp.types.TextContent) assert result.messages[0].content.text == "Welcome to FastMCP, Alice!" + + +async def test_proxy_handles_multiple_concurrent_tasks_correctly( + proxy_server: FastMCPProxy, +): + results = {} + + async def get_and_store(name, coro): + results[name] = await coro() + + async with create_task_group() as tg: + tg.start_soon(get_and_store, "prompts", proxy_server.get_prompts) + tg.start_soon(get_and_store, "resources", proxy_server.get_resources) + tg.start_soon(get_and_store, "tools", proxy_server.get_tools) + + assert list(results) == Contains("resources", "prompts", "tools") + assert list(results["prompts"]) == Contains("welcome") + assert [r.name for r in results["resources"].values()] == Contains( + "data://users", "resource://wave" + ) + assert list(results["tools"]) == Contains( + "greet", "add", "error_tool", "tool_without_description" + )