mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-15 16:49:24 +03:00
refactor: move MCP lifecycle out of AgentLoop (#5343)
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user