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>
This commit is contained in:
Vonbai 2026-04-14 00:11:11 +08:00 committed by William Easton
commit 955c996ad0
No known key found for this signature in database
2 changed files with 76 additions and 2 deletions

View file

@ -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):
"""

View file

@ -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.