From 955c996ad0ce19693bfcb9520864ebe1de7070ee Mon Sep 17 00:00:00 2001 From: Vonbai <107612985+vonbai@users.noreply.github.com> Date: Tue, 14 Apr 2026 00:11:11 +0800 Subject: [PATCH] Harden forced client disconnect cleanup (#3885) * Harden forced client disconnect cleanup * Handle cancelled force-close waits Generated with Codex. --------- Co-authored-by: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> --- src/fastmcp/client/client.py | 13 +++++- tests/client/client/test_client.py | 65 ++++++++++++++++++++++++++++++ 2 files changed, 76 insertions(+), 2 deletions(-) diff --git a/src/fastmcp/client/client.py b/src/fastmcp/client/client.py index 17fb7be90..4fa066a43 100644 --- a/src/fastmcp/client/client.py +++ b/src/fastmcp/client/client.py @@ -642,10 +642,19 @@ class Client( # stop the active session if self._session_state.session_task is None: return + session_task = self._session_state.session_task self._session_state.stop_event.set() # wait for session to finish to ensure state has been reset - await self._session_state.session_task - self._session_state.session_task = None + try: + if force: + with anyio.CancelScope(shield=True): + with anyio.move_on_after(self._disconnect_timeout): + with suppress(asyncio.CancelledError): + await session_task + else: + await session_task + finally: + self._session_state.session_task = None async def _session_runner(self): """ diff --git a/tests/client/client/test_client.py b/tests/client/client/test_client.py index b513b5367..f408f4748 100644 --- a/tests/client/client/test_client.py +++ b/tests/client/client/test_client.py @@ -439,6 +439,32 @@ class _DelayedConnectTransport(ClientTransport): await self._inner.close() +class _DelayedDisconnectTransport(ClientTransport): + def __init__( + self, + inner: ClientTransport, + disconnect_started: anyio.Event, + allow_disconnect: anyio.Event, + ) -> None: + self._inner = inner + self._disconnect_started = disconnect_started + self._allow_disconnect = allow_disconnect + + @contextlib.asynccontextmanager + async def connect_session( + self, **session_kwargs: Any + ) -> AsyncIterator[ClientSession]: + async with self._inner.connect_session(**session_kwargs) as session: + try: + yield session + finally: + self._disconnect_started.set() + await self._allow_disconnect.wait() + + async def close(self) -> None: + await self._inner.close() + + async def test_client_nested_context_manager(fastmcp_server): """Test that the client connects and disconnects once in nested context manager.""" @@ -552,6 +578,45 @@ async def test_cancelled_context_entry_waiter_does_not_close_active_session( assert await a == 3 +async def test_force_close_cancelled_wait_starts_fresh_session(fastmcp_server): + disconnect_started = anyio.Event() + allow_disconnect = anyio.Event() + client = Client( + transport=_DelayedDisconnectTransport( + FastMCPTransport(fastmcp_server), + disconnect_started=disconnect_started, + allow_disconnect=allow_disconnect, + ) + ) + + await client._connect() + original_session_task = client._session_state.session_task + assert original_session_task is not None + + close_task = asyncio.create_task(client.close()) + await disconnect_started.wait() + + close_task.cancel() + + async def reconnect_and_count_tools() -> int: + async with client: + assert client._session_state.session_task is not original_session_task + tools = await client.list_tools() + return len(tools) + + reconnect_task = asyncio.create_task(reconnect_and_count_tools()) + await asyncio.sleep(0) + assert not reconnect_task.done() + + allow_disconnect.set() + + with contextlib.suppress(asyncio.CancelledError): + await close_task + + assert await reconnect_task == 3 + assert original_session_task.done() + + async def test_concurrent_client_context_managers(): """ Test that concurrent client usage doesn't cause cross-task cancel scope issues.