From 90180498280cbf76aa4cc180a8031fd58d994a31 Mon Sep 17 00:00:00 2001 From: Sillocan Date: Fri, 30 May 2025 16:17:06 -0700 Subject: [PATCH 1/2] test: Demonstrate issue with concurrent proxy tasks --- tests/server/test_proxy.py | 24 ++++++++++++++++++++++++ 1 file changed, 24 insertions(+) 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" + ) From 90dd0c04f35f5cbbb4aaf6df32783dbfdb259529 Mon Sep 17 00:00:00 2001 From: Sillocan Date: Fri, 30 May 2025 18:45:59 -0700 Subject: [PATCH 2/2] fix: Use a background task for managing session state mimicing stdiotransport --- src/fastmcp/client/client.py | 57 +++++++++++++++++++++++------------- 1 file changed, 37 insertions(+), 20 deletions(-) diff --git a/src/fastmcp/client/client.py b/src/fastmcp/client/client.py index 03fb7fd5a..b982c0464 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 @@ -147,6 +148,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: @@ -186,6 +191,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: @@ -237,34 +243,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()