Clean up cancelled connection startup (#2614)

This commit is contained in:
Shawn Thapa 2025-12-14 05:45:05 -08:00 committed by GitHub
commit d26b04f80e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 142 additions and 3 deletions

View file

@ -6,7 +6,7 @@ import datetime
import secrets
import uuid
import weakref
from contextlib import AsyncExitStack, asynccontextmanager
from contextlib import AsyncExitStack, asynccontextmanager, suppress
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any, Generic, Literal, TypeVar, cast, overload
@ -502,7 +502,35 @@ class Client(Generic[ClientTransportT]):
self._session_state.session_task = asyncio.create_task(
self._session_runner()
)
await self._session_state.ready_event.wait()
try:
await self._session_state.ready_event.wait()
except asyncio.CancelledError:
# Cancellation during initial connection startup can leave the
# background session task running because __aexit__ is never invoked
# when __aenter__ is cancelled. Since we hold the session lock here
# and we know we started the session task, it's safe to tear it down
# without impacting other active contexts.
session_task = self._session_state.session_task
if session_task is not None:
# Request a graceful stop if the runner has already reached
# its stop_event wait.
self._session_state.stop_event.set()
session_task.cancel()
# Preserve the original cancellation.
with suppress(asyncio.CancelledError, Exception):
await asyncio.shield(session_task)
# Reset session state so future callers can reconnect cleanly.
self._session_state.session_task = None
self._session_state.session = None
self._session_state.initialize_result = None
self._session_state.nesting_counter = 0
# Preserve the original cancellation.
with suppress(asyncio.CancelledError, Exception):
await asyncio.shield(self.transport.close())
raise
if self._session_state.session_task.done():
exception = self._session_state.session_task.exception()

View file

@ -1,9 +1,12 @@
import asyncio
import contextlib
import sys
from collections.abc import AsyncIterator
from typing import Any, cast
import anyio
import pytest
from mcp import McpError
from mcp import ClientSession, McpError
from mcp.client.auth import OAuthClientProvider
from pydantic import AnyUrl
@ -11,6 +14,7 @@ import fastmcp
from fastmcp.client import Client
from fastmcp.client.auth.bearer import BearerAuth
from fastmcp.client.transports import (
ClientTransport,
FastMCPTransport,
MCPConfigTransport,
SSETransport,
@ -479,6 +483,30 @@ async def test_server_info_custom_version():
assert result.serverInfo.version == fastmcp.__version__
class _DelayedConnectTransport(ClientTransport):
def __init__(
self,
inner: ClientTransport,
connect_started: anyio.Event,
allow_connect: anyio.Event,
) -> None:
self._inner = inner
self._connect_started = connect_started
self._allow_connect = allow_connect
@contextlib.asynccontextmanager
async def connect_session(
self, **session_kwargs: Any
) -> AsyncIterator[ClientSession]:
self._connect_started.set()
await self._allow_connect.wait()
async with self._inner.connect_session(**session_kwargs) as session:
yield session
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."""
@ -509,6 +537,89 @@ async def test_client_nested_context_manager(fastmcp_server):
assert client._session_state.session is None
async def test_client_context_entry_cancelled_starter_cleans_up(fastmcp_server):
connect_started = anyio.Event()
allow_connect = anyio.Event()
client = Client(
transport=_DelayedConnectTransport(
FastMCPTransport(fastmcp_server),
connect_started=connect_started,
allow_connect=allow_connect,
)
)
async def enter_and_never_reach_body() -> None:
async with client:
pytest.fail(
"Context body should not be reached when __aenter__ is cancelled"
)
task = asyncio.create_task(enter_and_never_reach_body())
await connect_started.wait()
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
# Connection startup was cancelled; session state should be fully reset.
assert client._session_state.session_task is None
assert client._session_state.session is None
assert client._session_state.nesting_counter == 0
# A future connection attempt should work normally.
allow_connect.set()
async with client:
tools = await client.list_tools()
assert len(tools) == 3
async def test_cancelled_context_entry_waiter_does_not_close_active_session(
fastmcp_server,
):
connect_started = anyio.Event()
allow_connect = anyio.Event()
client = Client(
transport=_DelayedConnectTransport(
FastMCPTransport(fastmcp_server),
connect_started=connect_started,
allow_connect=allow_connect,
)
)
b_done = asyncio.Event()
b_started = asyncio.Event()
async def task_a() -> int:
async with client:
await b_done.wait()
tools = await client.list_tools()
return len(tools)
async def task_b() -> None:
b_started.set()
async with client:
pytest.fail("This context should never be entered due to cancellation")
a = asyncio.create_task(task_a())
await connect_started.wait()
b = asyncio.create_task(task_b())
await b_started.wait()
await asyncio.sleep(0) # let task_b attempt to acquire the client lock
b.cancel()
allow_connect.set()
with pytest.raises(asyncio.CancelledError):
await b
# task_b is fully cancelled; allow task_a to exercise the connected session.
b_done.set()
assert await a == 3
async def test_concurrent_client_context_managers():
"""
Test that concurrent client usage doesn't cause cross-task cancel scope issues.