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:
@@ -1,4 +1,4 @@
|
||||
"""Tests for MCP connection lifecycle in AgentLoop."""
|
||||
"""Tests for the application-owned MCP provider lifecycle."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -6,7 +6,7 @@ import asyncio
|
||||
from contextlib import AsyncExitStack
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import anyio
|
||||
import pytest
|
||||
@@ -15,11 +15,10 @@ from mcp.shared.exceptions import McpError
|
||||
from mcp.shared.message import SessionMessage
|
||||
from mcp.types import ErrorData
|
||||
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.agent.tools import mcp as mcp_runtime
|
||||
from nanobot.agent.tools.base import Tool
|
||||
from nanobot.agent.tools.mcp import MCPResourceWrapper, MCPToolWrapper
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.agent.tools.mcp import MCPProvider, MCPResourceWrapper, MCPToolWrapper
|
||||
from nanobot.agent.tools.registry import ToolRegistry
|
||||
from nanobot.config.loader import load_config, save_config
|
||||
from nanobot.config.schema import MCPServerConfig
|
||||
|
||||
@@ -74,18 +73,20 @@ class _FakeMcpTool(Tool):
|
||||
return "ok"
|
||||
|
||||
|
||||
def _make_loop(tmp_path, *, mcp_servers: dict | None = None) -> AgentLoop:
|
||||
bus = MessageBus()
|
||||
provider = MagicMock()
|
||||
provider.get_default_model.return_value = "test-model"
|
||||
provider.generation.max_tokens = 4096
|
||||
return AgentLoop(
|
||||
bus=bus,
|
||||
provider=provider,
|
||||
workspace=tmp_path,
|
||||
model="test-model",
|
||||
mcp_servers=mcp_servers or {"test": object()},
|
||||
def _stdio_server(command: str = "test-mcp") -> MCPServerConfig:
|
||||
return MCPServerConfig(type="stdio", command=command)
|
||||
|
||||
|
||||
def _make_provider(
|
||||
*,
|
||||
mcp_servers: dict[str, MCPServerConfig] | None = None,
|
||||
) -> tuple[MCPProvider, ToolRegistry]:
|
||||
registry = ToolRegistry()
|
||||
provider = MCPProvider(
|
||||
mcp_servers if mcp_servers is not None else {"test": _stdio_server()},
|
||||
registry,
|
||||
)
|
||||
return provider, registry
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -140,7 +141,7 @@ async def test_owned_mcp_connection_closes_from_its_owner_task():
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connect_mcp_retries_when_no_servers_connect(tmp_path, monkeypatch: pytest.MonkeyPatch):
|
||||
loop = _make_loop(tmp_path)
|
||||
provider, _registry = _make_provider()
|
||||
attempts = 0
|
||||
|
||||
async def _fake_connect(_servers, _registry):
|
||||
@@ -150,12 +151,12 @@ async def test_connect_mcp_retries_when_no_servers_connect(tmp_path, monkeypatch
|
||||
|
||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||
|
||||
await loop._connect_mcp()
|
||||
await loop._connect_mcp()
|
||||
await provider.connect()
|
||||
await provider.connect()
|
||||
|
||||
assert attempts == 2
|
||||
assert loop._mcp_stacks == {}
|
||||
assert loop.mcp_runtime_status() == {"test": "failed"}
|
||||
assert provider.connected_server_names == set()
|
||||
assert provider.runtime_status() == {"test": "failed"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -168,7 +169,7 @@ async def test_connect_mcp_does_not_report_failure_before_oauth_authorization(
|
||||
auth="oauth",
|
||||
url="https://mcp.example.com/mcp",
|
||||
)
|
||||
loop = _make_loop(tmp_path, mcp_servers={"oauth-app": cfg})
|
||||
provider, _registry = _make_provider(mcp_servers={"oauth-app": cfg})
|
||||
connect = AsyncMock()
|
||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", connect)
|
||||
monkeypatch.setattr(
|
||||
@@ -176,19 +177,20 @@ async def test_connect_mcp_does_not_report_failure_before_oauth_authorization(
|
||||
lambda _name, _url: False,
|
||||
)
|
||||
|
||||
await loop._connect_mcp()
|
||||
await provider.connect()
|
||||
|
||||
connect.assert_not_awaited()
|
||||
assert loop.mcp_runtime_status() == {}
|
||||
assert provider.runtime_status() == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agent_loop_run_closes_mcp_from_connection_owner_task(
|
||||
async def test_mcp_provider_closes_connections_independently_from_agent_loop(
|
||||
tmp_path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
loop = _make_loop(tmp_path, mcp_servers={"playwright": object()})
|
||||
connected = asyncio.Event()
|
||||
provider, registry = _make_provider(
|
||||
mcp_servers={"playwright": _stdio_server("playwright")}
|
||||
)
|
||||
owner_tasks: list[asyncio.Task | None] = []
|
||||
closed_tasks: list[asyncio.Task | None] = []
|
||||
|
||||
@@ -203,40 +205,38 @@ async def test_agent_loop_run_closes_mcp_from_connection_owner_task(
|
||||
|
||||
async def _fake_connect(servers, _registry):
|
||||
stacks = {name: _OwnerCheckedStack() for name in servers}
|
||||
connected.set()
|
||||
return stacks
|
||||
|
||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||
|
||||
task = asyncio.create_task(loop.run())
|
||||
await asyncio.wait_for(connected.wait(), timeout=1)
|
||||
loop.stop()
|
||||
task.cancel()
|
||||
await asyncio.gather(task, return_exceptions=True)
|
||||
await provider.connect()
|
||||
registry.register(_FakeMcpTool("mcp_playwright_search"))
|
||||
await provider.aclose()
|
||||
|
||||
assert owner_tasks
|
||||
assert closed_tasks == owner_tasks
|
||||
assert loop._mcp_stacks == {}
|
||||
assert provider.connected_server_names == set()
|
||||
assert registry.get("mcp_playwright_search") is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_close_server_ignores_server_cancelled_error(tmp_path):
|
||||
loop = _make_loop(tmp_path)
|
||||
provider, _registry = _make_provider()
|
||||
|
||||
class _ServerCancelledStack:
|
||||
async def aclose(self) -> None:
|
||||
raise asyncio.CancelledError()
|
||||
|
||||
loop._mcp_stacks = {"test": _ServerCancelledStack()}
|
||||
provider._connections = {"test": _ServerCancelledStack()}
|
||||
|
||||
await mcp_runtime._close_server(loop, "test")
|
||||
await provider._close_server("test")
|
||||
|
||||
assert loop._mcp_stacks == {}
|
||||
assert provider.connected_server_names == set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_close_mcp_servers_continues_after_server_cancelled_error(tmp_path):
|
||||
loop = _make_loop(tmp_path)
|
||||
async def test_provider_close_continues_after_server_cancelled_error(tmp_path):
|
||||
provider, _registry = _make_provider()
|
||||
closed: list[str] = []
|
||||
|
||||
class _ServerCancelledStack:
|
||||
@@ -247,21 +247,53 @@ async def test_close_mcp_servers_continues_after_server_cancelled_error(tmp_path
|
||||
async def aclose(self) -> None:
|
||||
closed.append("second")
|
||||
|
||||
loop._mcp_stacks = {
|
||||
provider._connections = {
|
||||
"first": _ServerCancelledStack(),
|
||||
"second": _TrackedStack(),
|
||||
}
|
||||
|
||||
await mcp_runtime.close_mcp_servers(loop)
|
||||
await provider.aclose()
|
||||
|
||||
assert closed == ["second"]
|
||||
assert loop._mcp_stacks == {}
|
||||
assert provider.connected_server_names == set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_provider_close_finishes_other_connections_before_propagating_cancellation(
|
||||
tmp_path,
|
||||
):
|
||||
provider, _registry = _make_provider()
|
||||
started = asyncio.Event()
|
||||
closed: list[str] = []
|
||||
|
||||
class _BlockingStack:
|
||||
async def aclose(self) -> None:
|
||||
started.set()
|
||||
await asyncio.Event().wait()
|
||||
|
||||
class _TrackedStack:
|
||||
async def aclose(self) -> None:
|
||||
closed.append("second")
|
||||
|
||||
provider._connections = {
|
||||
"first": _BlockingStack(),
|
||||
"second": _TrackedStack(),
|
||||
}
|
||||
task = asyncio.create_task(provider.aclose())
|
||||
await asyncio.wait_for(started.wait(), timeout=1)
|
||||
|
||||
task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await task
|
||||
|
||||
assert closed == ["second"]
|
||||
assert provider.connected_server_names == set()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("close_all", [False, True], ids=["single", "all"])
|
||||
async def test_mcp_cleanup_re_raises_external_cancellation(tmp_path, close_all: bool):
|
||||
loop = _make_loop(tmp_path)
|
||||
provider, _registry = _make_provider()
|
||||
started = asyncio.Event()
|
||||
|
||||
class _BlockingStack:
|
||||
@@ -269,12 +301,12 @@ async def test_mcp_cleanup_re_raises_external_cancellation(tmp_path, close_all:
|
||||
started.set()
|
||||
await asyncio.Event().wait()
|
||||
|
||||
loop._mcp_stacks = {"test": _BlockingStack()}
|
||||
provider._connections = {"test": _BlockingStack()}
|
||||
|
||||
if close_all:
|
||||
task = asyncio.create_task(mcp_runtime.close_mcp_servers(loop))
|
||||
task = asyncio.create_task(provider.aclose())
|
||||
else:
|
||||
task = asyncio.create_task(mcp_runtime._close_server(loop, "test"))
|
||||
task = asyncio.create_task(provider._close_server("test"))
|
||||
await asyncio.wait_for(started.wait(), timeout=1)
|
||||
task.cancel()
|
||||
|
||||
@@ -312,41 +344,38 @@ async def test_reload_mcp_servers_adds_and_removes_tools_without_restart(
|
||||
return stacks
|
||||
|
||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||
loop = _make_loop(tmp_path, mcp_servers={})
|
||||
provider, registry = _make_provider(mcp_servers={})
|
||||
|
||||
added = await mcp_runtime.reload_servers(loop, loop.tools)
|
||||
added = await provider.reload()
|
||||
|
||||
assert added["ok"] is True
|
||||
assert added["added"] == ["browserbase"]
|
||||
assert loop.tools.has("mcp_browserbase_navigate")
|
||||
assert "browserbase" in loop._mcp_stacks
|
||||
assert registry.has("mcp_browserbase_navigate")
|
||||
assert provider.connected_server_names == {"browserbase"}
|
||||
|
||||
config = load_config()
|
||||
del config.tools.mcp_servers["browserbase"]
|
||||
save_config(config)
|
||||
|
||||
removed = await mcp_runtime.reload_servers(loop, loop.tools)
|
||||
removed = await provider.reload()
|
||||
|
||||
assert removed["ok"] is True
|
||||
assert removed["removed"] == ["browserbase"]
|
||||
assert not loop.tools.has("mcp_browserbase_navigate")
|
||||
assert "browserbase" not in loop._mcp_stacks
|
||||
assert not registry.has("mcp_browserbase_navigate")
|
||||
assert provider.connected_server_names == set()
|
||||
assert closed == ["browserbase"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_request_mcp_reload_reaches_runtime_control_without_restart(
|
||||
async def test_reload_is_a_direct_provider_operation_without_an_agent_loop(
|
||||
tmp_path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
config_path = tmp_path / "config.json"
|
||||
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
|
||||
config = load_config()
|
||||
config.tools.mcp_servers["browserbase"] = MCPServerConfig(
|
||||
browserbase = MCPServerConfig(
|
||||
type="stdio",
|
||||
command="browserbase-mcp",
|
||||
)
|
||||
save_config(config)
|
||||
configured: dict[str, MCPServerConfig] = {"browserbase": browserbase}
|
||||
|
||||
closed: list[str] = []
|
||||
|
||||
@@ -364,37 +393,68 @@ async def test_request_mcp_reload_reaches_runtime_control_without_restart(
|
||||
return stacks
|
||||
|
||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||
loop = _make_loop(tmp_path, mcp_servers={})
|
||||
registry = ToolRegistry()
|
||||
provider = MCPProvider({}, registry, server_loader=lambda: configured)
|
||||
|
||||
async def _handle_one_runtime_control() -> None:
|
||||
msg = await loop.bus.consume_inbound()
|
||||
handled = await mcp_runtime.handle_runtime_control(loop, msg, loop.tools)
|
||||
assert handled is True
|
||||
|
||||
consumer = asyncio.create_task(_handle_one_runtime_control())
|
||||
result = await mcp_runtime.request_mcp_reload(loop.bus, timeout=2.0)
|
||||
await consumer
|
||||
result = await provider.reload()
|
||||
|
||||
assert result["ok"] is True
|
||||
assert result["added"] == ["browserbase"]
|
||||
assert result["requires_restart"] is False
|
||||
assert loop.tools.has("mcp_browserbase_navigate")
|
||||
assert registry.has("mcp_browserbase_navigate")
|
||||
|
||||
config = load_config()
|
||||
del config.tools.mcp_servers["browserbase"]
|
||||
save_config(config)
|
||||
configured = {}
|
||||
|
||||
consumer = asyncio.create_task(_handle_one_runtime_control())
|
||||
result = await mcp_runtime.request_mcp_reload(loop.bus, timeout=2.0)
|
||||
await consumer
|
||||
result = await provider.reload()
|
||||
|
||||
assert result["ok"] is True
|
||||
assert result["removed"] == ["browserbase"]
|
||||
assert result["requires_restart"] is False
|
||||
assert not loop.tools.has("mcp_browserbase_navigate")
|
||||
assert not registry.has("mcp_browserbase_navigate")
|
||||
assert closed == ["browserbase"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reload_timeout_marks_attempted_server_failed_and_allows_retry(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
server = _stdio_server("slow-mcp")
|
||||
started = asyncio.Event()
|
||||
attempts = 0
|
||||
|
||||
async def _fake_connect(servers, _registry):
|
||||
nonlocal attempts
|
||||
attempts += 1
|
||||
if attempts == 1:
|
||||
started.set()
|
||||
await asyncio.Event().wait()
|
||||
stack = AsyncExitStack()
|
||||
await stack.__aenter__()
|
||||
return {name: stack for name in servers}
|
||||
|
||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||
provider = MCPProvider(
|
||||
{"test": server},
|
||||
ToolRegistry(),
|
||||
server_loader=lambda: {"test": server},
|
||||
)
|
||||
|
||||
reload_task = asyncio.create_task(provider.reload())
|
||||
await asyncio.wait_for(started.wait(), timeout=1.0)
|
||||
with pytest.raises(asyncio.TimeoutError):
|
||||
await asyncio.wait_for(reload_task, timeout=0.01)
|
||||
|
||||
assert provider.connected_server_names == set()
|
||||
assert provider.runtime_status() == {"test": "failed"}
|
||||
|
||||
result = await provider.reload()
|
||||
|
||||
assert result["ok"] is True
|
||||
assert provider.connected_server_names == {"test"}
|
||||
assert provider.runtime_status() == {"test": "connected"}
|
||||
await provider.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_reload_mcp_servers_retries_configured_server_without_live_stack(
|
||||
tmp_path,
|
||||
@@ -419,16 +479,18 @@ async def test_reload_mcp_servers_retries_configured_server_without_live_stack(
|
||||
return stacks
|
||||
|
||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||
loop = _make_loop(tmp_path, mcp_servers={"browserbase": config.tools.mcp_servers["browserbase"]})
|
||||
provider, registry = _make_provider(
|
||||
mcp_servers={"browserbase": config.tools.mcp_servers["browserbase"]}
|
||||
)
|
||||
|
||||
result = await mcp_runtime.reload_servers(loop, loop.tools)
|
||||
result = await provider.reload()
|
||||
|
||||
assert result["ok"] is True
|
||||
assert result["added"] == []
|
||||
assert result["changed"] == []
|
||||
assert result["retried"] == ["browserbase"]
|
||||
assert loop.tools.has("mcp_browserbase_navigate")
|
||||
await loop.close_mcp()
|
||||
assert registry.has("mcp_browserbase_navigate")
|
||||
await provider.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -465,16 +527,16 @@ async def test_reload_mcp_servers_skips_oauth_server_waiting_for_authorization(
|
||||
"nanobot.agent.tools.mcp_oauth.mcp_oauth_has_credentials",
|
||||
lambda name, _url: name == "linear",
|
||||
)
|
||||
loop = _make_loop(tmp_path, mcp_servers={"notion": notion})
|
||||
provider, _registry = _make_provider(mcp_servers={"notion": notion})
|
||||
|
||||
result = await mcp_runtime.reload_servers(loop, loop.tools)
|
||||
result = await provider.reload()
|
||||
|
||||
assert attempted == ["linear"]
|
||||
assert result["ok"] is True
|
||||
assert result["failed"] == []
|
||||
assert result["retried"] == []
|
||||
assert result["connected"] == ["linear"]
|
||||
await loop.close_mcp()
|
||||
await provider.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -482,7 +544,9 @@ async def test_mcp_tool_reconnects_after_session_terminated(
|
||||
tmp_path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
loop = _make_loop(tmp_path, mcp_servers={"remote": object()})
|
||||
provider, registry = _make_provider(
|
||||
mcp_servers={"remote": _stdio_server("remote")}
|
||||
)
|
||||
closed: list[str] = []
|
||||
sessions: list[Any] = []
|
||||
connect_count = 0
|
||||
@@ -525,8 +589,8 @@ async def test_mcp_tool_reconnects_after_session_terminated(
|
||||
|
||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||
|
||||
await loop._connect_mcp()
|
||||
old_tool = loop.tools.get("mcp_remote_quote")
|
||||
await provider.connect()
|
||||
old_tool = registry.get("mcp_remote_quote")
|
||||
assert isinstance(old_tool, MCPToolWrapper)
|
||||
|
||||
output = await old_tool.execute(symbol="AAPL")
|
||||
@@ -536,8 +600,8 @@ async def test_mcp_tool_reconnects_after_session_terminated(
|
||||
assert closed == ["remote"]
|
||||
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
|
||||
assert provider.connected_server_names == {"remote"}
|
||||
assert registry.get("mcp_remote_quote") is not old_tool
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -545,7 +609,9 @@ async def test_mcp_reconnect_handler_uses_sanitized_server_prefix(
|
||||
tmp_path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
loop = _make_loop(tmp_path, mcp_servers={"remote_": object()})
|
||||
provider, registry = _make_provider(
|
||||
mcp_servers={"remote_": _stdio_server("remote")}
|
||||
)
|
||||
connect_count = 0
|
||||
|
||||
class _FakeSession:
|
||||
@@ -578,15 +644,15 @@ async def test_mcp_reconnect_handler_uses_sanitized_server_prefix(
|
||||
|
||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||
|
||||
await loop._connect_mcp()
|
||||
old_tool = loop.tools.get("mcp_remote_quote")
|
||||
await provider.connect()
|
||||
old_tool = registry.get("mcp_remote_quote")
|
||||
assert isinstance(old_tool, MCPToolWrapper)
|
||||
|
||||
output = await old_tool.execute()
|
||||
|
||||
assert output == "recovered"
|
||||
assert connect_count == 2
|
||||
assert loop.tools.get("mcp_remote_quote") is not old_tool
|
||||
assert registry.get("mcp_remote_quote") is not old_tool
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -594,7 +660,9 @@ async def test_concurrent_mcp_reconnect_reuses_fresh_session(
|
||||
tmp_path,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
loop = _make_loop(tmp_path, mcp_servers={"remote": object()})
|
||||
provider, registry = _make_provider(
|
||||
mcp_servers={"remote": _stdio_server("remote")}
|
||||
)
|
||||
closed: list[str] = []
|
||||
connect_count = 0
|
||||
|
||||
@@ -638,9 +706,9 @@ async def test_concurrent_mcp_reconnect_reuses_fresh_session(
|
||||
|
||||
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
|
||||
|
||||
await loop._connect_mcp()
|
||||
old_alpha = loop.tools.get("mcp_remote_resource_alpha")
|
||||
old_beta = loop.tools.get("mcp_remote_resource_beta")
|
||||
await provider.connect()
|
||||
old_alpha = registry.get("mcp_remote_resource_alpha")
|
||||
old_beta = registry.get("mcp_remote_resource_beta")
|
||||
assert isinstance(old_alpha, MCPResourceWrapper)
|
||||
assert isinstance(old_beta, MCPResourceWrapper)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user