fix(mcp): keep transport cleanup in owner tasks

This commit is contained in:
Xubin Ren
2026-07-11 11:45:20 +08:00
parent bd0dd85f44
commit edf78e7054
5 changed files with 187 additions and 100 deletions
+4
View File
@@ -44,6 +44,10 @@ async def connect_mcp(state: Any, tools: ToolRegistry) -> None:
await mcp_tools.connect_missing_servers(state, tools) await mcp_tools.connect_missing_servers(state, tools)
async def close_mcp(state: Any) -> None:
await mcp_tools.close_mcp_servers(state)
async def handle_runtime_control(state: Any, msg: InboundMessage, tools: ToolRegistry) -> bool: async def handle_runtime_control(state: Any, msg: InboundMessage, tools: ToolRegistry) -> bool:
return await mcp_tools.handle_runtime_control(state, msg, tools) return await mcp_tools.handle_runtime_control(state, msg, tools)
+4 -15
View File
@@ -7,7 +7,7 @@ import dataclasses
import os import os
import time import time
from collections.abc import Mapping from collections.abc import Mapping
from contextlib import AsyncExitStack, nullcontext, suppress from contextlib import nullcontext, suppress
from dataclasses import dataclass, field from dataclasses import dataclass, field
from enum import Enum, auto from enum import Enum, auto
from functools import partial from functools import partial
@@ -82,6 +82,7 @@ from nanobot.utils.runtime import (
) )
if TYPE_CHECKING: if TYPE_CHECKING:
from nanobot.agent.tools.mcp import MCPConnection
from nanobot.config.schema import ( from nanobot.config.schema import (
ChannelsConfig, ChannelsConfig,
ProviderConfig, ProviderConfig,
@@ -360,8 +361,7 @@ class AgentLoop:
self._unified_session = unified_session self._unified_session = unified_session
self._running = False self._running = False
self._mcp_servers = mcp_servers or {} self._mcp_servers = mcp_servers or {}
self._mcp_stacks: dict[str, AsyncExitStack] = {} self._mcp_stacks: dict[str, MCPConnection] = {}
self._mcp_retired_stacks: list[tuple[str, AsyncExitStack]] = []
self._mcp_connecting = False self._mcp_connecting = False
self._active_tasks: dict[str, list[asyncio.Task]] = {} # session_key -> tasks self._active_tasks: dict[str, list[asyncio.Task]] = {} # session_key -> tasks
self._background_tasks: list[asyncio.Task] = [] self._background_tasks: list[asyncio.Task] = []
@@ -1178,18 +1178,7 @@ class AgentLoop:
if self._background_tasks: if self._background_tasks:
await asyncio.gather(*self._background_tasks, return_exceptions=True) await asyncio.gather(*self._background_tasks, return_exceptions=True)
self._background_tasks.clear() self._background_tasks.clear()
stacks = [*self._mcp_retired_stacks, *self._mcp_stacks.items()] await agent_context.close_mcp(self)
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()
def _schedule_background(self, coro) -> None: def _schedule_background(self, coro) -> None:
"""Schedule a coroutine as a tracked background task (drained on shutdown).""" """Schedule a coroutine as a tracked background task (drained on shutdown)."""
+131 -41
View File
@@ -9,7 +9,7 @@ import shutil
import urllib.parse import urllib.parse
from collections.abc import Awaitable, Callable from collections.abc import Awaitable, Callable
from contextlib import AsyncExitStack, suppress from contextlib import AsyncExitStack, suppress
from typing import Any, Mapping from typing import Any, Mapping, Protocol
from weakref import WeakKeyDictionary from weakref import WeakKeyDictionary
import httpx import httpx
@@ -54,6 +54,26 @@ _RELOAD_LOCKS: WeakKeyDictionary[Any, asyncio.Lock] = WeakKeyDictionary()
_ReconnectCallback = Callable[[str, str, Tool], Awaitable[Tool | None]] _ReconnectCallback = Callable[[str, str, Tool], Awaitable[Tool | None]]
class MCPConnection(Protocol):
async def aclose(self) -> None: ...
class _OwnedMCPConnection:
"""Close an MCP transport from the task that originally opened it."""
def __init__(self, owner: asyncio.Task[None], close_requested: asyncio.Event) -> None:
self._owner = owner
self._close_requested = close_requested
async def aclose(self) -> None:
self._close_requested.set()
try:
await asyncio.shield(self._owner)
except asyncio.CancelledError:
if not self._owner.cancelled():
raise
def _is_malformed_mcp_progress_notification(message: Any) -> bool: def _is_malformed_mcp_progress_notification(message: Any) -> bool:
payload = _mcp_jsonrpc_payload(message) payload = _mcp_jsonrpc_payload(message)
if _payload_value(payload, "method") != "notifications/progress": if _payload_value(payload, "method") != "notifications/progress":
@@ -814,19 +834,19 @@ class MCPPromptWrapper(_MCPWrapperBase):
async def connect_mcp_servers( async def connect_mcp_servers(
mcp_servers: dict, registry: ToolRegistry mcp_servers: dict, registry: ToolRegistry
) -> dict[str, AsyncExitStack]: ) -> dict[str, MCPConnection]:
"""Connect to configured MCP servers and register their tools, resources, prompts. """Connect to configured MCP servers and register their tools, resources, prompts.
Returns a dict mapping server name -> its dedicated AsyncExitStack. Returns one connection handle per server. Each handle keeps the task that
Each server gets its own stack to prevent cancel scope conflicts entered the MCP SDK contexts alive so reconnect and shutdown can close
when multiple MCP servers are configured. AnyIO cancel scopes from their owning task.
""" """
from mcp import ClientSession, StdioServerParameters from mcp import ClientSession, StdioServerParameters
from mcp.client.sse import sse_client from mcp.client.sse import sse_client
from mcp.client.stdio import stdio_client from mcp.client.stdio import stdio_client
from mcp.client.streamable_http import streamable_http_client from mcp.client.streamable_http import streamable_http_client
async def connect_single_server(name: str, cfg) -> tuple[str, AsyncExitStack | None]: async def open_single_server(name: str, cfg) -> tuple[str, AsyncExitStack | None]:
server_stack = AsyncExitStack() server_stack = AsyncExitStack()
await server_stack.__aenter__() await server_stack.__aenter__()
@@ -1045,7 +1065,43 @@ async def connect_mcp_servers(
await server_stack.aclose() await server_stack.aclose()
return name, None return name, None
server_stacks: dict[str, AsyncExitStack] = {} async def connect_single_server(name: str, cfg) -> tuple[str, MCPConnection | None]:
loop = asyncio.get_running_loop()
ready: asyncio.Future[bool] = loop.create_future()
close_requested = asyncio.Event()
async def own_connection() -> None:
stack: AsyncExitStack | None = None
try:
_, stack = await open_single_server(name, cfg)
if not ready.done():
ready.set_result(stack is not None)
if stack is not None:
await close_requested.wait()
except BaseException as exc:
if not ready.done():
ready.set_exception(exc)
raise
finally:
if stack is not None:
await stack.aclose()
owner = asyncio.create_task(own_connection(), name=f"mcp:{name}")
connection = _OwnedMCPConnection(owner, close_requested)
try:
connected = await ready
except BaseException:
close_requested.set()
owner.cancel()
with suppress(BaseException):
await asyncio.shield(owner)
raise
if not connected:
await connection.aclose()
return name, None
return name, connection
server_stacks: dict[str, MCPConnection] = {}
for name, cfg in mcp_servers.items(): for name, cfg in mcp_servers.items():
try: try:
@@ -1124,31 +1180,44 @@ def runtime_lines(
async def connect_missing_servers(state: Any, registry: ToolRegistry) -> None: async def connect_missing_servers(state: Any, registry: ToolRegistry) -> None:
"""Connect configured MCP servers that are not currently live.""" """Connect configured MCP servers that are not currently live."""
missing_servers = { async with _reload_lock(state):
name: cfg for name, cfg in state._mcp_servers.items() if name not in state._mcp_stacks if getattr(state, "_mcp_closing", False):
} return
if state._mcp_connecting or not missing_servers: missing_servers = {
return name: cfg for name, cfg in state._mcp_servers.items() if name not in state._mcp_stacks
state._mcp_connecting = True }
try: if state._mcp_connecting or not missing_servers:
connected = await connect_mcp_servers(missing_servers, registry) return
state._mcp_stacks.update(connected) state._mcp_connecting = True
_attach_reconnect_handlers(state, registry, connected) try:
if connected: connected = await connect_mcp_servers(missing_servers, registry)
logger.info("MCP connected servers: {}", sorted(connected)) if getattr(state, "_mcp_closing", False):
else: for connection in connected.values():
logger.warning("No MCP servers connected successfully (will retry next message)") await connection.aclose()
except asyncio.CancelledError: return
logger.warning("MCP connection cancelled (will retry next message)") state._mcp_stacks.update(connected)
except BaseException as e: _attach_reconnect_handlers(state, registry, connected)
logger.warning("Failed to connect MCP servers (will retry next message): {}", e) if connected:
finally: logger.info("MCP connected servers: {}", sorted(connected))
state._mcp_connecting = False else:
logger.warning("No MCP servers connected successfully (will retry next message)")
except asyncio.CancelledError:
logger.warning("MCP connection cancelled (will retry next message)")
except BaseException as e:
logger.warning("Failed to connect MCP servers (will retry next message): {}", e)
finally:
state._mcp_connecting = False
async def reload_servers(state: Any, registry: ToolRegistry) -> dict[str, Any]: async def reload_servers(state: Any, registry: ToolRegistry) -> dict[str, Any]:
"""Reconcile live MCP connections with the current config file.""" """Reconcile live MCP connections with the current config file."""
async with _reload_lock(state): async with _reload_lock(state):
if getattr(state, "_mcp_closing", False):
return {
"ok": False,
"message": "MCP connections are shutting down.",
"requires_restart": True,
}
try: try:
from nanobot.config.loader import load_config, resolve_config_env_vars from nanobot.config.loader import load_config, resolve_config_env_vars
@@ -1177,7 +1246,7 @@ async def reload_servers(state: Any, registry: ToolRegistry) -> dict[str, Any]:
tools_removed = 0 tools_removed = 0
for name in [*removed, *changed]: for name in [*removed, *changed]:
tools_removed += _unregister_server_tools(state, registry, name) tools_removed += _unregister_server_tools(state, registry, name)
_retire_server_stack(state, name) await _close_server(state, name)
state._mcp_servers = next_servers state._mcp_servers = next_servers
retry_missing = sorted( retry_missing = sorted(
@@ -1187,9 +1256,17 @@ async def reload_servers(state: Any, registry: ToolRegistry) -> dict[str, Any]:
) )
to_connect_names = sorted(set(added) | set(changed) | set(retry_missing)) to_connect_names = sorted(set(added) | set(changed) | set(retry_missing))
to_connect = {name: next_servers[name] for name in to_connect_names} to_connect = {name: next_servers[name] for name in to_connect_names}
connected: dict[str, AsyncExitStack] = {} connected: dict[str, MCPConnection] = {}
if to_connect: if to_connect:
connected = await connect_mcp_servers(to_connect, registry) connected = await connect_mcp_servers(to_connect, registry)
if getattr(state, "_mcp_closing", False):
for connection in connected.values():
await connection.aclose()
return {
"ok": False,
"message": "MCP connections are shutting down.",
"requires_restart": True,
}
state._mcp_stacks.update(connected) state._mcp_stacks.update(connected)
_attach_reconnect_handlers(state, registry, connected) _attach_reconnect_handlers(state, registry, connected)
@@ -1323,6 +1400,8 @@ async def _refresh_terminated_server(
stale_tool: Tool, stale_tool: Tool,
) -> Tool | None: ) -> Tool | None:
async with _reload_lock(state): async with _reload_lock(state):
if getattr(state, "_mcp_closing", False):
return None
cfg = state._mcp_servers.get(server_name) cfg = state._mcp_servers.get(server_name)
if cfg is None: if cfg is None:
logger.warning( logger.warning(
@@ -1341,9 +1420,13 @@ async def _refresh_terminated_server(
logger.warning("MCP server '{}' session terminated; refreshing connection", server_name) logger.warning("MCP server '{}' session terminated; refreshing connection", server_name)
_unregister_server_tools(state, registry, server_name) _unregister_server_tools(state, registry, server_name)
_retire_server_stack(state, server_name) await _close_server(state, server_name)
connected = await connect_mcp_servers({server_name: cfg}, registry) connected = await connect_mcp_servers({server_name: cfg}, registry)
if getattr(state, "_mcp_closing", False):
for connection in connected.values():
await connection.aclose()
return None
state._mcp_stacks.update(connected) state._mcp_stacks.update(connected)
_attach_reconnect_handlers(state, registry, connected) _attach_reconnect_handlers(state, registry, connected)
if server_name not in connected: if server_name not in connected:
@@ -1378,17 +1461,24 @@ def _unregister_server_tools(state: Any, registry: ToolRegistry, server_name: st
return removed return removed
def _retire_server_stack(state: Any, server_name: str) -> None: async def _close_server(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) stack = state._mcp_stacks.pop(server_name, None)
if stack is None: if stack is None:
return return
retired = getattr(state, "_mcp_retired_stacks", None) try:
if retired is not None: await stack.aclose()
retired.append((server_name, stack)) except (RuntimeError, BaseExceptionGroup):
logger.debug("MCP server '{}' cleanup error (can be ignored)", server_name)
async def close_mcp_servers(state: Any) -> None:
"""Close every MCP connection while excluding reconnect and hot reload."""
state._mcp_closing = True
async with _reload_lock(state):
connections = list(state._mcp_stacks.items())
state._mcp_stacks.clear()
for name, connection in connections:
try:
await connection.aclose()
except (RuntimeError, BaseExceptionGroup):
logger.debug("MCP server '{}' cleanup error (can be ignored)", name)
+20 -30
View File
@@ -74,15 +74,6 @@ class _FakeMcpTool(Tool):
return "ok" 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: def _make_loop(tmp_path, *, mcp_servers: dict | None = None) -> AgentLoop:
bus = MessageBus() bus = MessageBus()
provider = MagicMock() provider = MagicMock()
@@ -125,15 +116,26 @@ async def test_mcp_read_filter_drops_progress_notifications_without_progress_tok
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_close_mcp_swallows_sdk_cancel_scope_cleanup(tmp_path): async def test_owned_mcp_connection_closes_from_its_owner_task():
loop = _make_loop(tmp_path) close_requested = asyncio.Event()
stack = _CancelScopeStack() ready = asyncio.Event()
loop._mcp_stacks["remote"] = stack # type: ignore[assignment] tasks: dict[str, asyncio.Task] = {}
await loop.close_mcp() async def own_connection() -> None:
tasks["open"] = asyncio.current_task() # type: ignore[assignment]
ready.set()
await close_requested.wait()
tasks["close"] = asyncio.current_task() # type: ignore[assignment]
assert stack.closed is True owner = asyncio.create_task(own_connection())
assert loop._mcp_stacks == {} connection = mcp_runtime._OwnedMCPConnection(owner, close_requested)
await ready.wait()
await connection.aclose()
assert tasks["open"] is owner
assert tasks["close"] is owner
assert tasks["close"] is not asyncio.current_task()
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -241,9 +243,6 @@ async def test_reload_mcp_servers_adds_and_removes_tools_without_restart(
assert removed["removed"] == ["browserbase"] assert removed["removed"] == ["browserbase"]
assert not loop.tools.has("mcp_browserbase_navigate") assert not loop.tools.has("mcp_browserbase_navigate")
assert "browserbase" not in loop._mcp_stacks assert "browserbase" not in loop._mcp_stacks
assert closed == []
await loop.close_mcp()
assert closed == ["browserbase"] assert closed == ["browserbase"]
@@ -305,9 +304,6 @@ async def test_request_mcp_reload_reaches_runtime_control_without_restart(
assert result["removed"] == ["browserbase"] assert result["removed"] == ["browserbase"]
assert result["requires_restart"] is False assert result["requires_restart"] is False
assert not loop.tools.has("mcp_browserbase_navigate") assert not loop.tools.has("mcp_browserbase_navigate")
assert closed == []
await loop.close_mcp()
assert closed == ["browserbase"] assert closed == ["browserbase"]
@@ -403,15 +399,12 @@ async def test_mcp_tool_reconnects_after_session_terminated(
assert output == "recovered" assert output == "recovered"
assert connect_count == 2 assert connect_count == 2
assert closed == [] assert closed == ["remote"]
assert sessions[0].call_count == 1 assert sessions[0].call_count == 1
assert sessions[1].call_count == 1 assert sessions[1].call_count == 1
assert "remote" in loop._mcp_stacks assert "remote" in loop._mcp_stacks
assert loop.tools.get("mcp_remote_quote") is not old_tool assert loop.tools.get("mcp_remote_quote") is not old_tool
await loop.close_mcp()
assert closed == ["remote", "remote"]
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_mcp_reconnect_handler_uses_sanitized_server_prefix( async def test_mcp_reconnect_handler_uses_sanitized_server_prefix(
@@ -521,7 +514,4 @@ async def test_concurrent_mcp_reconnect_reuses_fresh_session(
assert outputs == ["fresh:alpha", "fresh:beta"] assert outputs == ["fresh:alpha", "fresh:beta"]
assert connect_count == 2 assert connect_count == 2
assert closed == [] assert closed == ["remote"]
await loop.close_mcp()
assert closed == ["remote", "remote"]
+28 -14
View File
@@ -15,7 +15,6 @@ import asyncio
import multiprocessing import multiprocessing
import socket import socket
import time import time
from contextlib import suppress
from unittest.mock import MagicMock from unittest.mock import MagicMock
import httpx import httpx
@@ -70,7 +69,7 @@ async def _wait_for_server(url: str, timeout: float = 10.0) -> bool:
deadline = time.monotonic() + timeout deadline = time.monotonic() + timeout
while time.monotonic() < deadline: while time.monotonic() < deadline:
try: try:
async with httpx.AsyncClient(timeout=2.0) as client: async with httpx.AsyncClient(timeout=2.0, trust_env=False) as client:
response = await client.get( response = await client.get(
url, url,
headers={"Accept": "text/event-stream"}, headers={"Accept": "text/event-stream"},
@@ -170,26 +169,30 @@ async def test_mcp_reconnect_after_session_timeout(tmp_path, mcp_server_url):
) )
loop = _make_loop(tmp_path, mcp_servers={"repro": cfg}) loop = _make_loop(tmp_path, mcp_servers={"repro": cfg})
await loop._connect_mcp() await asyncio.create_task(loop._connect_mcp())
assert "repro" in loop._mcp_stacks assert "repro" in loop._mcp_stacks
tool = loop.tools.get("mcp_repro_greet") tool = loop.tools.get("mcp_repro_greet")
assert isinstance(tool, MCPToolWrapper) assert isinstance(tool, MCPToolWrapper)
output = await tool.execute(name="first") output = await asyncio.create_task(tool.execute(name="first"))
assert "Hello, first" in output assert "Hello, first" in output
# Wait for the server-side idle timeout to terminate the session. # Wait for the server-side idle timeout to terminate the session.
await asyncio.sleep(_IDLE_TIMEOUT_SECONDS + 1) await asyncio.sleep(_IDLE_TIMEOUT_SECONDS + 1)
output = await tool.execute(name="second") output = await asyncio.create_task(tool.execute(name="second"))
assert "Hello, second" in output assert "Hello, second" in output
await loop.close_mcp() await asyncio.create_task(loop.close_mcp())
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_mcp_reconnect_during_shutdown_does_not_crash(tmp_path, mcp_server_url): async def test_mcp_reconnect_during_shutdown_does_not_crash(
tmp_path,
mcp_server_url,
monkeypatch: pytest.MonkeyPatch,
):
"""Simulate the production crash: shutdown while reconnect is in flight.""" """Simulate the production crash: shutdown while reconnect is in flight."""
cfg = MCPServerConfig( cfg = MCPServerConfig(
type="streamableHttp", type="streamableHttp",
@@ -199,15 +202,28 @@ async def test_mcp_reconnect_during_shutdown_does_not_crash(tmp_path, mcp_server
) )
loop = _make_loop(tmp_path, mcp_servers={"repro": cfg}) loop = _make_loop(tmp_path, mcp_servers={"repro": cfg})
await loop._connect_mcp() await asyncio.create_task(loop._connect_mcp())
tool = loop.tools.get("mcp_repro_greet") tool = loop.tools.get("mcp_repro_greet")
assert isinstance(tool, MCPToolWrapper) assert isinstance(tool, MCPToolWrapper)
await tool.execute(name="first") await asyncio.create_task(tool.execute(name="first"))
await asyncio.sleep(_IDLE_TIMEOUT_SECONDS + 1) await asyncio.sleep(_IDLE_TIMEOUT_SECONDS + 1)
reconnect_started = asyncio.Event()
finish_reconnect = asyncio.Event()
real_connect = mcp_module.connect_mcp_servers
async def gated_connect(*args, **kwargs):
reconnect_started.set()
await finish_reconnect.wait()
return await real_connect(*args, **kwargs)
monkeypatch.setattr(mcp_module, "connect_mcp_servers", gated_connect)
call_task = asyncio.create_task(tool.execute(name="second")) call_task = asyncio.create_task(tool.execute(name="second"))
loop.stop() await asyncio.wait_for(reconnect_started.wait(), timeout=5)
close_task = asyncio.create_task(loop.close_mcp())
await asyncio.sleep(0)
finish_reconnect.set()
unhandled: list[BaseException] = [] unhandled: list[BaseException] = []
@@ -219,13 +235,11 @@ async def test_mcp_reconnect_during_shutdown_does_not_crash(tmp_path, mcp_server
asyncio.get_running_loop().set_exception_handler(capture_unhandled) asyncio.get_running_loop().set_exception_handler(capture_unhandled)
try: try:
await asyncio.wait_for(asyncio.shield(call_task), timeout=15) await asyncio.wait_for(asyncio.gather(call_task, close_task), timeout=15)
except asyncio.CancelledError: except asyncio.CancelledError:
unhandled.append(asyncio.CancelledError("main task cancelled by leaked MCP cancel scope")) unhandled.append(asyncio.CancelledError("main task cancelled by leaked MCP cancel scope"))
except Exception as exc: except Exception as exc:
unhandled.append(exc) unhandled.append(exc)
with suppress(Exception):
await loop.close_mcp()
assert not unhandled, f"Unhandled exception leaked during reconnect/shutdown: {unhandled[0]}" assert not unhandled, f"Unhandled exception leaked during reconnect/shutdown: {unhandled[0]}"
assert loop._mcp_stacks == {}