fix(mcp): defer stale stack cleanup during reconnect
This commit is contained in:
@@ -361,6 +361,7 @@ class AgentLoop:
|
|||||||
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, AsyncExitStack] = {}
|
||||||
|
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] = []
|
||||||
@@ -1177,9 +1178,15 @@ 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()
|
||||||
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:
|
try:
|
||||||
await stack.aclose()
|
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):
|
except (RuntimeError, BaseExceptionGroup):
|
||||||
logger.debug("MCP server '{}' cleanup error (can be ignored)", name)
|
logger.debug("MCP server '{}' cleanup error (can be ignored)", name)
|
||||||
self._mcp_stacks.clear()
|
self._mcp_stacks.clear()
|
||||||
|
|||||||
@@ -1177,7 +1177,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)
|
||||||
await _close_server(state, name)
|
_retire_server_stack(state, name)
|
||||||
|
|
||||||
state._mcp_servers = next_servers
|
state._mcp_servers = next_servers
|
||||||
retry_missing = sorted(
|
retry_missing = sorted(
|
||||||
@@ -1341,7 +1341,7 @@ 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)
|
||||||
await _close_server(state, server_name)
|
_retire_server_stack(state, server_name)
|
||||||
|
|
||||||
connected = await connect_mcp_servers({server_name: cfg}, registry)
|
connected = await connect_mcp_servers({server_name: cfg}, registry)
|
||||||
state._mcp_stacks.update(connected)
|
state._mcp_stacks.update(connected)
|
||||||
@@ -1378,11 +1378,17 @@ def _unregister_server_tools(state: Any, registry: ToolRegistry, server_name: st
|
|||||||
return removed
|
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)
|
stack = state._mcp_stacks.pop(server_name, None)
|
||||||
if stack is None:
|
if stack is None:
|
||||||
return
|
return
|
||||||
try:
|
retired = getattr(state, "_mcp_retired_stacks", None)
|
||||||
await stack.aclose()
|
if retired is not None:
|
||||||
except (RuntimeError, BaseExceptionGroup):
|
retired.append((server_name, stack))
|
||||||
logger.debug("MCP server '{}' cleanup error (can be ignored)", server_name)
|
|
||||||
|
|||||||
@@ -74,6 +74,15 @@ 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()
|
||||||
@@ -115,6 +124,18 @@ async def test_mcp_read_filter_drops_progress_notifications_without_progress_tok
|
|||||||
assert forwarded == [tool_change, valid_progress]
|
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
|
@pytest.mark.asyncio
|
||||||
async def test_connect_mcp_retries_when_no_servers_connect(tmp_path, monkeypatch: pytest.MonkeyPatch):
|
async def test_connect_mcp_retries_when_no_servers_connect(tmp_path, monkeypatch: pytest.MonkeyPatch):
|
||||||
loop = _make_loop(tmp_path)
|
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 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"]
|
||||||
|
|
||||||
|
|
||||||
@@ -281,6 +305,9 @@ 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"]
|
||||||
|
|
||||||
|
|
||||||
@@ -376,12 +403,15 @@ 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 == ["remote"]
|
assert closed == []
|
||||||
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(
|
||||||
@@ -491,4 +521,7 @@ 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 == ["remote"]
|
assert closed == []
|
||||||
|
|
||||||
|
await loop.close_mcp()
|
||||||
|
assert closed == ["remote", "remote"]
|
||||||
|
|||||||
Reference in New Issue
Block a user