From f3d1b9ca2d2d2f9f2f3665d6b0c3e683ca41ebba Mon Sep 17 00:00:00 2001 From: flyzstu Date: Wed, 8 Jul 2026 09:23:55 +0800 Subject: [PATCH] fix(mcp): defer stale stack cleanup during reconnect --- nanobot/agent/loop.py | 9 +++++++- nanobot/agent/tools/mcp.py | 20 ++++++++++------ tests/agent/test_mcp_connection.py | 37 ++++++++++++++++++++++++++++-- 3 files changed, 56 insertions(+), 10 deletions(-) diff --git a/nanobot/agent/loop.py b/nanobot/agent/loop.py index 392ec900..47f29287 100644 --- a/nanobot/agent/loop.py +++ b/nanobot/agent/loop.py @@ -361,6 +361,7 @@ class AgentLoop: self._running = False self._mcp_servers = mcp_servers or {} self._mcp_stacks: dict[str, AsyncExitStack] = {} + self._mcp_retired_stacks: list[tuple[str, AsyncExitStack]] = [] self._mcp_connecting = False self._active_tasks: dict[str, list[asyncio.Task]] = {} # session_key -> tasks self._background_tasks: list[asyncio.Task] = [] @@ -1177,9 +1178,15 @@ class AgentLoop: if self._background_tasks: await asyncio.gather(*self._background_tasks, return_exceptions=True) self._background_tasks.clear() - for name, stack in self._mcp_stacks.items(): + stacks = [*self._mcp_retired_stacks, *self._mcp_stacks.items()] + self._mcp_retired_stacks.clear() + for name, stack in stacks: try: await stack.aclose() + except asyncio.CancelledError as exc: + if not str(exc).startswith("Cancelled via cancel scope"): + raise + logger.debug("MCP server '{}' cleanup cancelled by SDK (can be ignored)", name) except (RuntimeError, BaseExceptionGroup): logger.debug("MCP server '{}' cleanup error (can be ignored)", name) self._mcp_stacks.clear() diff --git a/nanobot/agent/tools/mcp.py b/nanobot/agent/tools/mcp.py index e4dfad55..625127e4 100644 --- a/nanobot/agent/tools/mcp.py +++ b/nanobot/agent/tools/mcp.py @@ -1177,7 +1177,7 @@ async def reload_servers(state: Any, registry: ToolRegistry) -> dict[str, Any]: tools_removed = 0 for name in [*removed, *changed]: tools_removed += _unregister_server_tools(state, registry, name) - await _close_server(state, name) + _retire_server_stack(state, name) state._mcp_servers = next_servers retry_missing = sorted( @@ -1341,7 +1341,7 @@ async def _refresh_terminated_server( logger.warning("MCP server '{}' session terminated; refreshing connection", server_name) _unregister_server_tools(state, registry, server_name) - await _close_server(state, server_name) + _retire_server_stack(state, server_name) connected = await connect_mcp_servers({server_name: cfg}, registry) state._mcp_stacks.update(connected) @@ -1378,11 +1378,17 @@ def _unregister_server_tools(state: Any, registry: ToolRegistry, server_name: st return removed -async def _close_server(state: Any, server_name: str) -> None: +def _retire_server_stack(state: Any, server_name: str) -> None: + """Remove a stale MCP stack from active use without closing it mid-turn. + + MCP stream transports use AnyIO cancel scopes. Closing a stack from the + reconnecting dispatch task can inject ``CancelledError`` into the task that + originally opened it (often ``AgentLoop.run``), which crashes the gateway. + Retired stacks are closed later by ``AgentLoop.close_mcp`` during shutdown. + """ stack = state._mcp_stacks.pop(server_name, None) if stack is None: return - try: - await stack.aclose() - except (RuntimeError, BaseExceptionGroup): - logger.debug("MCP server '{}' cleanup error (can be ignored)", server_name) + retired = getattr(state, "_mcp_retired_stacks", None) + if retired is not None: + retired.append((server_name, stack)) diff --git a/tests/agent/test_mcp_connection.py b/tests/agent/test_mcp_connection.py index d5de1343..bf8893f3 100644 --- a/tests/agent/test_mcp_connection.py +++ b/tests/agent/test_mcp_connection.py @@ -74,6 +74,15 @@ class _FakeMcpTool(Tool): return "ok" +class _CancelScopeStack: + def __init__(self) -> None: + self.closed = False + + async def aclose(self) -> None: + self.closed = True + raise asyncio.CancelledError("Cancelled via cancel scope test") + + def _make_loop(tmp_path, *, mcp_servers: dict | None = None) -> AgentLoop: bus = MessageBus() provider = MagicMock() @@ -115,6 +124,18 @@ async def test_mcp_read_filter_drops_progress_notifications_without_progress_tok assert forwarded == [tool_change, valid_progress] +@pytest.mark.asyncio +async def test_close_mcp_swallows_sdk_cancel_scope_cleanup(tmp_path): + loop = _make_loop(tmp_path) + stack = _CancelScopeStack() + loop._mcp_stacks["remote"] = stack # type: ignore[assignment] + + await loop.close_mcp() + + assert stack.closed is True + assert loop._mcp_stacks == {} + + @pytest.mark.asyncio async def test_connect_mcp_retries_when_no_servers_connect(tmp_path, monkeypatch: pytest.MonkeyPatch): loop = _make_loop(tmp_path) @@ -220,6 +241,9 @@ async def test_reload_mcp_servers_adds_and_removes_tools_without_restart( assert removed["removed"] == ["browserbase"] assert not loop.tools.has("mcp_browserbase_navigate") assert "browserbase" not in loop._mcp_stacks + + assert closed == [] + await loop.close_mcp() assert closed == ["browserbase"] @@ -281,6 +305,9 @@ async def test_request_mcp_reload_reaches_runtime_control_without_restart( assert result["removed"] == ["browserbase"] assert result["requires_restart"] is False assert not loop.tools.has("mcp_browserbase_navigate") + + assert closed == [] + await loop.close_mcp() assert closed == ["browserbase"] @@ -376,12 +403,15 @@ async def test_mcp_tool_reconnects_after_session_terminated( assert output == "recovered" assert connect_count == 2 - assert closed == ["remote"] + assert closed == [] assert sessions[0].call_count == 1 assert sessions[1].call_count == 1 assert "remote" in loop._mcp_stacks assert loop.tools.get("mcp_remote_quote") is not old_tool + await loop.close_mcp() + assert closed == ["remote", "remote"] + @pytest.mark.asyncio async def test_mcp_reconnect_handler_uses_sanitized_server_prefix( @@ -491,4 +521,7 @@ async def test_concurrent_mcp_reconnect_reuses_fresh_session( assert outputs == ["fresh:alpha", "fresh:beta"] assert connect_count == 2 - assert closed == ["remote"] + assert closed == [] + + await loop.close_mcp() + assert closed == ["remote", "remote"]