From 07d0b7556d96a74746df5505a5c6e882783d0d21 Mon Sep 17 00:00:00 2001 From: Rio Yu <52408936+rioyu123@users.noreply.github.com> Date: Mon, 14 Sep 2026 16:54:43 +0800 Subject: [PATCH] fix(mcp): cancel in-flight parallel server startup with its caller --- src/agents/mcp/manager.py | 25 ++++- tests/mcp/test_mcp_server_manager.py | 146 +++++++++++++++++++++++++++ 2 files changed, 167 insertions(+), 4 deletions(-) diff --git a/src/agents/mcp/manager.py b/src/agents/mcp/manager.py index 733e744363..b09d96d104 100644 --- a/src/agents/mcp/manager.py +++ b/src/agents/mcp/manager.py @@ -32,6 +32,7 @@ class _ServerCommand: action: str timeout_seconds: float | None future: asyncio.Future[None] + cancelled_by_caller: bool = False class _ServerWorker: @@ -40,6 +41,7 @@ def __init__(self, server: MCPServer) -> None: self._queue: asyncio.Queue[_ServerCommand] = asyncio.Queue() self._task = asyncio.create_task(self._run()) self._cleanup_future: asyncio.Future[None] | None = None + self._active_command: _ServerCommand | None = None @property def is_done(self) -> bool: @@ -85,15 +87,23 @@ async def cleanup(self, timeout_seconds: float | None) -> None: async def _submit(self, action: str, timeout_seconds: float | None) -> None: loop = asyncio.get_running_loop() future: asyncio.Future[None] = loop.create_future() - self._queue.put_nowait( - _ServerCommand(action=action, timeout_seconds=timeout_seconds, future=future) - ) - await future + command = _ServerCommand(action=action, timeout_seconds=timeout_seconds, future=future) + self._queue.put_nowait(command) + try: + await future + except asyncio.CancelledError: + # Interrupt only this startup, never another command or a cleanup owner. + if self._active_command is command and action == "connect": + command.cancelled_by_caller = self._task.cancel() + raise async def _run(self) -> None: while True: command = await self._queue.get() + if command.action == "connect" and command.future.cancelled(): + continue should_exit = command.action == "cleanup" + self._active_command = command try: if command.action == "connect": await _run_with_timeout_in_task(self._server.connect, command.timeout_seconds) @@ -106,6 +116,13 @@ async def _run(self) -> None: except BaseException as exc: if not command.future.cancelled(): command.future.set_exception(exc) + finally: + self._active_command = None + if command.cancelled_by_caller: + # Balance only the cancellation requested for this command. + uncancel = getattr(self._task, "uncancel", None) + if uncancel is not None: + uncancel() if should_exit: return diff --git a/tests/mcp/test_mcp_server_manager.py b/tests/mcp/test_mcp_server_manager.py index 40476867a6..681e5c3c60 100644 --- a/tests/mcp/test_mcp_server_manager.py +++ b/tests/mcp/test_mcp_server_manager.py @@ -103,6 +103,152 @@ async def cleanup(self) -> None: self.cleanup_finished.set() +@pytest.mark.asyncio +@pytest.mark.parametrize("suppress_cancelled_error", [False, True]) +@pytest.mark.parametrize("connect_timeout_seconds", [None, 10.0]) +async def test_manager_cancels_parallel_startup_with_its_caller( + suppress_cancelled_error: bool, + connect_timeout_seconds: float | None, +) -> None: + # Events preserve the active startup boundary without a model or network service. + started = asyncio.Event() + release = asyncio.Event() + cancelled = asyncio.Event() + + class StartingServer(TaskBoundServer): + async def connect(self) -> None: + await super().connect() + started.set() + try: + await release.wait() + except asyncio.CancelledError: + cancelled.set() + raise + + healthy = TaskBoundServer() + server = StartingServer() + manager = MCPServerManager( + [healthy, server], + connect_in_parallel=True, + connect_timeout_seconds=connect_timeout_seconds, + cleanup_timeout_seconds=TEST_TIMEOUT_SECONDS, + suppress_cancelled_error=suppress_cancelled_error, + ) + connecting = asyncio.create_task(manager.connect_all()) + try: + await asyncio.wait_for(started.wait(), TEST_TIMEOUT_SECONDS) + connecting.cancel() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(connecting, 3 * TEST_TIMEOUT_SECONDS) + assert cancelled.is_set() + assert server.cleaned + assert healthy.cleaned + assert manager.active_servers == [] + assert manager._workers == {} + assert server._connect_task is not None + cancelling = getattr(server._connect_task, "cancelling", None) + if cancelling is not None: + assert cancelling() == 0 + finally: + release.set() + await asyncio.gather(connecting, return_exceptions=True) + await manager.cleanup_all() + + +@pytest.mark.asyncio +async def test_manager_skips_parallel_startup_cancelled_before_worker_runs( + monkeypatch: pytest.MonkeyPatch, +) -> None: + worker_started = asyncio.Event() + allow_worker = asyncio.Event() + cleanup_queued = asyncio.Event() + original_run = manager_module._ServerWorker._run + original_cleanup = manager_module._ServerWorker.cleanup + + async def delayed_run(worker: manager_module._ServerWorker) -> None: + worker_started.set() + await allow_worker.wait() + await original_run(worker) + + async def observed_cleanup( + worker: manager_module._ServerWorker, timeout_seconds: float | None + ) -> None: + cleanup_queued.set() + await original_cleanup(worker, timeout_seconds) + + monkeypatch.setattr(manager_module._ServerWorker, "_run", delayed_run) + monkeypatch.setattr(manager_module._ServerWorker, "cleanup", observed_cleanup) + + class UnstartedServer(TaskBoundServer): + async def cleanup(self) -> None: + self.cleaned = True + + server = UnstartedServer() + manager = MCPServerManager([server], connect_in_parallel=True, suppress_cancelled_error=False) + connecting = asyncio.create_task(manager.connect_all()) + try: + await asyncio.wait_for(worker_started.wait(), TEST_TIMEOUT_SECONDS) + connecting.cancel() + await asyncio.wait_for(cleanup_queued.wait(), TEST_TIMEOUT_SECONDS) + allow_worker.set() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(connecting, TEST_TIMEOUT_SECONDS) + assert server._connect_task is None + assert server.cleaned + assert manager._workers == {} + finally: + allow_worker.set() + await asyncio.gather(connecting, return_exceptions=True) + await manager.cleanup_all() + + +@pytest.mark.asyncio +async def test_manager_keeps_startup_teardown_alive_after_repeated_caller_cancellation() -> None: + started = asyncio.Event() + teardown_started = asyncio.Event() + release_teardown = asyncio.Event() + teardown_finished = asyncio.Event() + + class TeardownServer(TaskBoundServer): + async def connect(self) -> None: + await super().connect() + started.set() + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + teardown_started.set() + await release_teardown.wait() + teardown_finished.set() + raise + + server = TeardownServer() + manager = MCPServerManager( + [server], + connect_in_parallel=True, + connect_timeout_seconds=None, + suppress_cancelled_error=False, + ) + connecting = asyncio.create_task(manager.connect_all()) + try: + await asyncio.wait_for(started.wait(), TEST_TIMEOUT_SECONDS) + connecting.cancel() + await asyncio.wait_for(teardown_started.wait(), TEST_TIMEOUT_SECONDS) + connecting.cancel() + release_teardown.set() + with pytest.raises(asyncio.CancelledError): + await asyncio.wait_for(connecting, TEST_TIMEOUT_SECONDS) + await manager.cleanup_all() + assert teardown_finished.is_set() + assert server.cleaned + assert manager._workers == {} + if server._connect_task is not None and hasattr(server._connect_task, "cancelling"): + assert server._connect_task.cancelling() == 0 + finally: + release_teardown.set() + await asyncio.gather(connecting, return_exceptions=True) + await manager.cleanup_all() + + class BlockingCleanupFailureServer(TaskBoundServer): def __init__(self) -> None: super().__init__()