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:
@@ -15,15 +15,13 @@ import asyncio
|
||||
import multiprocessing
|
||||
import socket
|
||||
import time
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from nanobot.agent.loop import AgentLoop
|
||||
from nanobot.agent.tools import mcp as mcp_module
|
||||
from nanobot.agent.tools.mcp import MCPToolWrapper
|
||||
from nanobot.bus.queue import MessageBus
|
||||
from nanobot.agent.tools.mcp import MCPProvider, MCPToolWrapper
|
||||
from nanobot.agent.tools.registry import ToolRegistry
|
||||
from nanobot.config.schema import MCPServerConfig
|
||||
from nanobot.security import network as security_network
|
||||
|
||||
@@ -113,18 +111,9 @@ def mcp_server_url():
|
||||
process.join(timeout=2.0)
|
||||
|
||||
|
||||
def _make_loop(tmp_path, *, mcp_servers: dict) -> 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,
|
||||
)
|
||||
def _make_provider(*, mcp_servers: dict) -> tuple[MCPProvider, ToolRegistry]:
|
||||
registry = ToolRegistry()
|
||||
return MCPProvider(mcp_servers, registry), registry
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
@@ -170,12 +159,12 @@ async def test_mcp_reconnect_after_session_timeout(tmp_path, mcp_server_url):
|
||||
tool_timeout=_TOOL_TIMEOUT_SECONDS,
|
||||
enabled_tools=["*"],
|
||||
)
|
||||
loop = _make_loop(tmp_path, mcp_servers={"repro": cfg})
|
||||
provider, registry = _make_provider(mcp_servers={"repro": cfg})
|
||||
|
||||
await asyncio.create_task(loop._connect_mcp())
|
||||
assert "repro" in loop._mcp_stacks
|
||||
await asyncio.create_task(provider.connect())
|
||||
assert provider.connected_server_names == {"repro"}
|
||||
|
||||
tool = loop.tools.get("mcp_repro_greet")
|
||||
tool = registry.get("mcp_repro_greet")
|
||||
assert isinstance(tool, MCPToolWrapper)
|
||||
|
||||
output = await asyncio.create_task(tool.execute(name="first"))
|
||||
@@ -187,7 +176,7 @@ async def test_mcp_reconnect_after_session_timeout(tmp_path, mcp_server_url):
|
||||
output = await asyncio.create_task(tool.execute(name="second"))
|
||||
assert "Hello, second" in output
|
||||
|
||||
await asyncio.create_task(loop.close_mcp())
|
||||
await asyncio.create_task(provider.aclose())
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -203,10 +192,10 @@ async def test_mcp_reconnect_during_shutdown_does_not_crash(
|
||||
tool_timeout=_TOOL_TIMEOUT_SECONDS,
|
||||
enabled_tools=["*"],
|
||||
)
|
||||
loop = _make_loop(tmp_path, mcp_servers={"repro": cfg})
|
||||
provider, registry = _make_provider(mcp_servers={"repro": cfg})
|
||||
|
||||
await asyncio.create_task(loop._connect_mcp())
|
||||
tool = loop.tools.get("mcp_repro_greet")
|
||||
await asyncio.create_task(provider.connect())
|
||||
tool = registry.get("mcp_repro_greet")
|
||||
assert isinstance(tool, MCPToolWrapper)
|
||||
|
||||
await asyncio.create_task(tool.execute(name="first"))
|
||||
@@ -224,7 +213,7 @@ async def test_mcp_reconnect_during_shutdown_does_not_crash(
|
||||
monkeypatch.setattr(mcp_module, "connect_mcp_servers", gated_connect)
|
||||
call_task = asyncio.create_task(tool.execute(name="second"))
|
||||
await asyncio.wait_for(reconnect_started.wait(), timeout=5)
|
||||
close_task = asyncio.create_task(loop.close_mcp())
|
||||
close_task = asyncio.create_task(provider.aclose())
|
||||
await asyncio.sleep(0)
|
||||
finish_reconnect.set()
|
||||
|
||||
@@ -245,4 +234,4 @@ async def test_mcp_reconnect_during_shutdown_does_not_crash(
|
||||
unhandled.append(exc)
|
||||
|
||||
assert not unhandled, f"Unhandled exception leaked during reconnect/shutdown: {unhandled[0]}"
|
||||
assert loop._mcp_stacks == {}
|
||||
assert provider.connected_server_names == set()
|
||||
|
||||
Reference in New Issue
Block a user