refactor: move MCP lifecycle out of AgentLoop (#5343)

This commit is contained in:
chengyongru
2026-08-12 17:51:04 +08:00
committed by GitHub
parent 686dd0603e
commit 19997d20bb
39 changed files with 1192 additions and 846 deletions
+76 -26
View File
@@ -13,6 +13,7 @@ import pytest
import nanobot.agent.tools.mcp as mcp_mod
from nanobot.agent.tools.mcp import (
MCPPromptWrapper,
MCPProvider,
MCPResourceWrapper,
MCPToolWrapper,
_normalize_windows_stdio_command,
@@ -153,7 +154,7 @@ def _make_wrapper(session: object, *, timeout: float = 0.1) -> MCPToolWrapper:
@pytest.mark.asyncio
async def test_connect_missing_servers_propagates_external_cancellation(monkeypatch) -> None:
async def test_mcp_provider_connect_propagates_external_cancellation(monkeypatch) -> None:
started = asyncio.Event()
async def connect_mcp_servers(_servers: dict, _registry: ToolRegistry) -> dict:
@@ -161,24 +162,21 @@ async def test_connect_missing_servers_propagates_external_cancellation(monkeypa
await asyncio.sleep(60)
return {}
class State:
pass
state = State()
state._mcp_closing = False
state._mcp_servers = {"test": MCPServerConfig(command="fake")}
state._mcp_stacks = {}
state._mcp_connecting = False
provider = MCPProvider(
{"test": MCPServerConfig(command="fake")},
ToolRegistry(),
)
monkeypatch.setattr(mcp_mod, "connect_mcp_servers", connect_mcp_servers)
task = asyncio.create_task(mcp_mod.connect_missing_servers(state, ToolRegistry()))
task = asyncio.create_task(provider.connect())
await asyncio.wait_for(started.wait(), timeout=1.0)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert state._mcp_connecting is False
assert provider.connected_server_names == set()
assert provider.runtime_status() == {"test": "failed"}
@pytest.mark.asyncio
@@ -223,25 +221,20 @@ async def test_saved_oauth_http_403_projects_failed_runtime_without_details(
rejected_streamable_http,
)
class State:
pass
state = State()
state._mcp_closing = False
state._mcp_servers = {
"xmind": MCPServerConfig(
provider = MCPProvider(
{
"xmind": MCPServerConfig(
type="streamableHttp",
auth="oauth",
url="https://app.xmind.com/api/mcp",
)
}
state._mcp_stacks = {}
state._mcp_runtime_statuses = {}
state._mcp_connecting = False
)
},
ToolRegistry(),
)
await mcp_mod.connect_missing_servers(state, ToolRegistry())
await provider.connect()
snapshot = mcp_mod.runtime_status(state)
snapshot = provider.runtime_status()
assert snapshot == {"xmind": "failed"}
assert "saved-oauth-secret" not in str(snapshot)
assert "app.xmind.com" not in str(snapshot)
@@ -1263,6 +1256,59 @@ async def test_connect_mcp_servers_propagates_external_cancellation(
await asyncio.wait_for(closed.wait(), timeout=1.0)
@pytest.mark.asyncio
async def test_connect_mcp_servers_rolls_back_completed_batch_on_cancellation(
monkeypatch: pytest.MonkeyPatch,
) -> None:
slow_started = asyncio.Event()
closed: list[str] = []
sessions = {"fast": _make_fake_session(["demo"])}
class _SelectiveClientSession:
def __init__(self, read: object, _write: object) -> None:
self._session = sessions[str(read)]
async def __aenter__(self) -> object:
return self._session
async def __aexit__(self, exc_type, exc, tb) -> bool:
return False
@asynccontextmanager
async def _selective_stdio_client(params: object):
command = str(params.command)
try:
if command == "slow":
slow_started.set()
await asyncio.Event().wait()
yield command, object()
finally:
closed.append(command)
monkeypatch.setattr(sys.modules["mcp"], "ClientSession", _SelectiveClientSession)
monkeypatch.setattr(sys.modules["mcp.client.stdio"], "stdio_client", _selective_stdio_client)
registry = ToolRegistry()
task = asyncio.create_task(
connect_mcp_servers(
{
"fast": MCPServerConfig(command="fast"),
"slow": MCPServerConfig(command="slow"),
},
registry,
)
)
await asyncio.wait_for(slow_started.wait(), timeout=1.0)
assert registry.tool_names == ["mcp_fast_demo"]
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task
assert registry.tool_names == []
assert sorted(closed) == ["fast", "slow"]
@pytest.mark.asyncio
async def test_connect_mcp_servers_streamable_http_uses_finite_timeout(
fake_mcp_runtime: dict[str, object | None],
@@ -1900,7 +1946,11 @@ def test_long_server_name_tools_are_matched_by_server_name() -> None:
assert len(wrapper.name) == 64
assert not wrapper.name.startswith(mcp_mod._tool_prefix(server_name))
mcp_mod._attach_reconnect_handlers(SimpleNamespace(), registry, {server_name})
provider = MCPProvider(
{server_name: MCPServerConfig(command="fake")},
registry,
)
provider._attach_reconnect_handlers({server_name})
assert wrapper._reconnect is not None
assert other_wrapper._reconnect is None