fix(mcp): keep transport cleanup in owner tasks
This commit is contained in:
@@ -44,6 +44,10 @@ async def connect_mcp(state: Any, tools: ToolRegistry) -> None:
|
|||||||
await mcp_tools.connect_missing_servers(state, tools)
|
await mcp_tools.connect_missing_servers(state, tools)
|
||||||
|
|
||||||
|
|
||||||
|
async def close_mcp(state: Any) -> None:
|
||||||
|
await mcp_tools.close_mcp_servers(state)
|
||||||
|
|
||||||
|
|
||||||
async def handle_runtime_control(state: Any, msg: InboundMessage, tools: ToolRegistry) -> bool:
|
async def handle_runtime_control(state: Any, msg: InboundMessage, tools: ToolRegistry) -> bool:
|
||||||
return await mcp_tools.handle_runtime_control(state, msg, tools)
|
return await mcp_tools.handle_runtime_control(state, msg, tools)
|
||||||
|
|
||||||
|
|||||||
+4
-15
@@ -7,7 +7,7 @@ import dataclasses
|
|||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
from collections.abc import Mapping
|
from collections.abc import Mapping
|
||||||
from contextlib import AsyncExitStack, nullcontext, suppress
|
from contextlib import nullcontext, suppress
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from enum import Enum, auto
|
from enum import Enum, auto
|
||||||
from functools import partial
|
from functools import partial
|
||||||
@@ -82,6 +82,7 @@ from nanobot.utils.runtime import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from nanobot.agent.tools.mcp import MCPConnection
|
||||||
from nanobot.config.schema import (
|
from nanobot.config.schema import (
|
||||||
ChannelsConfig,
|
ChannelsConfig,
|
||||||
ProviderConfig,
|
ProviderConfig,
|
||||||
@@ -360,8 +361,7 @@ class AgentLoop:
|
|||||||
self._unified_session = unified_session
|
self._unified_session = unified_session
|
||||||
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, MCPConnection] = {}
|
||||||
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] = []
|
||||||
@@ -1178,18 +1178,7 @@ 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()
|
||||||
stacks = [*self._mcp_retired_stacks, *self._mcp_stacks.items()]
|
await agent_context.close_mcp(self)
|
||||||
self._mcp_retired_stacks.clear()
|
|
||||||
for name, stack in stacks:
|
|
||||||
try:
|
|
||||||
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):
|
|
||||||
logger.debug("MCP server '{}' cleanup error (can be ignored)", name)
|
|
||||||
self._mcp_stacks.clear()
|
|
||||||
|
|
||||||
def _schedule_background(self, coro) -> None:
|
def _schedule_background(self, coro) -> None:
|
||||||
"""Schedule a coroutine as a tracked background task (drained on shutdown)."""
|
"""Schedule a coroutine as a tracked background task (drained on shutdown)."""
|
||||||
|
|||||||
+131
-41
@@ -9,7 +9,7 @@ import shutil
|
|||||||
import urllib.parse
|
import urllib.parse
|
||||||
from collections.abc import Awaitable, Callable
|
from collections.abc import Awaitable, Callable
|
||||||
from contextlib import AsyncExitStack, suppress
|
from contextlib import AsyncExitStack, suppress
|
||||||
from typing import Any, Mapping
|
from typing import Any, Mapping, Protocol
|
||||||
from weakref import WeakKeyDictionary
|
from weakref import WeakKeyDictionary
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
@@ -54,6 +54,26 @@ _RELOAD_LOCKS: WeakKeyDictionary[Any, asyncio.Lock] = WeakKeyDictionary()
|
|||||||
_ReconnectCallback = Callable[[str, str, Tool], Awaitable[Tool | None]]
|
_ReconnectCallback = Callable[[str, str, Tool], Awaitable[Tool | None]]
|
||||||
|
|
||||||
|
|
||||||
|
class MCPConnection(Protocol):
|
||||||
|
async def aclose(self) -> None: ...
|
||||||
|
|
||||||
|
|
||||||
|
class _OwnedMCPConnection:
|
||||||
|
"""Close an MCP transport from the task that originally opened it."""
|
||||||
|
|
||||||
|
def __init__(self, owner: asyncio.Task[None], close_requested: asyncio.Event) -> None:
|
||||||
|
self._owner = owner
|
||||||
|
self._close_requested = close_requested
|
||||||
|
|
||||||
|
async def aclose(self) -> None:
|
||||||
|
self._close_requested.set()
|
||||||
|
try:
|
||||||
|
await asyncio.shield(self._owner)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
if not self._owner.cancelled():
|
||||||
|
raise
|
||||||
|
|
||||||
|
|
||||||
def _is_malformed_mcp_progress_notification(message: Any) -> bool:
|
def _is_malformed_mcp_progress_notification(message: Any) -> bool:
|
||||||
payload = _mcp_jsonrpc_payload(message)
|
payload = _mcp_jsonrpc_payload(message)
|
||||||
if _payload_value(payload, "method") != "notifications/progress":
|
if _payload_value(payload, "method") != "notifications/progress":
|
||||||
@@ -814,19 +834,19 @@ class MCPPromptWrapper(_MCPWrapperBase):
|
|||||||
|
|
||||||
async def connect_mcp_servers(
|
async def connect_mcp_servers(
|
||||||
mcp_servers: dict, registry: ToolRegistry
|
mcp_servers: dict, registry: ToolRegistry
|
||||||
) -> dict[str, AsyncExitStack]:
|
) -> dict[str, MCPConnection]:
|
||||||
"""Connect to configured MCP servers and register their tools, resources, prompts.
|
"""Connect to configured MCP servers and register their tools, resources, prompts.
|
||||||
|
|
||||||
Returns a dict mapping server name -> its dedicated AsyncExitStack.
|
Returns one connection handle per server. Each handle keeps the task that
|
||||||
Each server gets its own stack to prevent cancel scope conflicts
|
entered the MCP SDK contexts alive so reconnect and shutdown can close
|
||||||
when multiple MCP servers are configured.
|
AnyIO cancel scopes from their owning task.
|
||||||
"""
|
"""
|
||||||
from mcp import ClientSession, StdioServerParameters
|
from mcp import ClientSession, StdioServerParameters
|
||||||
from mcp.client.sse import sse_client
|
from mcp.client.sse import sse_client
|
||||||
from mcp.client.stdio import stdio_client
|
from mcp.client.stdio import stdio_client
|
||||||
from mcp.client.streamable_http import streamable_http_client
|
from mcp.client.streamable_http import streamable_http_client
|
||||||
|
|
||||||
async def connect_single_server(name: str, cfg) -> tuple[str, AsyncExitStack | None]:
|
async def open_single_server(name: str, cfg) -> tuple[str, AsyncExitStack | None]:
|
||||||
server_stack = AsyncExitStack()
|
server_stack = AsyncExitStack()
|
||||||
await server_stack.__aenter__()
|
await server_stack.__aenter__()
|
||||||
|
|
||||||
@@ -1045,7 +1065,43 @@ async def connect_mcp_servers(
|
|||||||
await server_stack.aclose()
|
await server_stack.aclose()
|
||||||
return name, None
|
return name, None
|
||||||
|
|
||||||
server_stacks: dict[str, AsyncExitStack] = {}
|
async def connect_single_server(name: str, cfg) -> tuple[str, MCPConnection | None]:
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
ready: asyncio.Future[bool] = loop.create_future()
|
||||||
|
close_requested = asyncio.Event()
|
||||||
|
|
||||||
|
async def own_connection() -> None:
|
||||||
|
stack: AsyncExitStack | None = None
|
||||||
|
try:
|
||||||
|
_, stack = await open_single_server(name, cfg)
|
||||||
|
if not ready.done():
|
||||||
|
ready.set_result(stack is not None)
|
||||||
|
if stack is not None:
|
||||||
|
await close_requested.wait()
|
||||||
|
except BaseException as exc:
|
||||||
|
if not ready.done():
|
||||||
|
ready.set_exception(exc)
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
if stack is not None:
|
||||||
|
await stack.aclose()
|
||||||
|
|
||||||
|
owner = asyncio.create_task(own_connection(), name=f"mcp:{name}")
|
||||||
|
connection = _OwnedMCPConnection(owner, close_requested)
|
||||||
|
try:
|
||||||
|
connected = await ready
|
||||||
|
except BaseException:
|
||||||
|
close_requested.set()
|
||||||
|
owner.cancel()
|
||||||
|
with suppress(BaseException):
|
||||||
|
await asyncio.shield(owner)
|
||||||
|
raise
|
||||||
|
if not connected:
|
||||||
|
await connection.aclose()
|
||||||
|
return name, None
|
||||||
|
return name, connection
|
||||||
|
|
||||||
|
server_stacks: dict[str, MCPConnection] = {}
|
||||||
|
|
||||||
for name, cfg in mcp_servers.items():
|
for name, cfg in mcp_servers.items():
|
||||||
try:
|
try:
|
||||||
@@ -1124,31 +1180,44 @@ def runtime_lines(
|
|||||||
|
|
||||||
async def connect_missing_servers(state: Any, registry: ToolRegistry) -> None:
|
async def connect_missing_servers(state: Any, registry: ToolRegistry) -> None:
|
||||||
"""Connect configured MCP servers that are not currently live."""
|
"""Connect configured MCP servers that are not currently live."""
|
||||||
missing_servers = {
|
async with _reload_lock(state):
|
||||||
name: cfg for name, cfg in state._mcp_servers.items() if name not in state._mcp_stacks
|
if getattr(state, "_mcp_closing", False):
|
||||||
}
|
return
|
||||||
if state._mcp_connecting or not missing_servers:
|
missing_servers = {
|
||||||
return
|
name: cfg for name, cfg in state._mcp_servers.items() if name not in state._mcp_stacks
|
||||||
state._mcp_connecting = True
|
}
|
||||||
try:
|
if state._mcp_connecting or not missing_servers:
|
||||||
connected = await connect_mcp_servers(missing_servers, registry)
|
return
|
||||||
state._mcp_stacks.update(connected)
|
state._mcp_connecting = True
|
||||||
_attach_reconnect_handlers(state, registry, connected)
|
try:
|
||||||
if connected:
|
connected = await connect_mcp_servers(missing_servers, registry)
|
||||||
logger.info("MCP connected servers: {}", sorted(connected))
|
if getattr(state, "_mcp_closing", False):
|
||||||
else:
|
for connection in connected.values():
|
||||||
logger.warning("No MCP servers connected successfully (will retry next message)")
|
await connection.aclose()
|
||||||
except asyncio.CancelledError:
|
return
|
||||||
logger.warning("MCP connection cancelled (will retry next message)")
|
state._mcp_stacks.update(connected)
|
||||||
except BaseException as e:
|
_attach_reconnect_handlers(state, registry, connected)
|
||||||
logger.warning("Failed to connect MCP servers (will retry next message): {}", e)
|
if connected:
|
||||||
finally:
|
logger.info("MCP connected servers: {}", sorted(connected))
|
||||||
state._mcp_connecting = False
|
else:
|
||||||
|
logger.warning("No MCP servers connected successfully (will retry next message)")
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
logger.warning("MCP connection cancelled (will retry next message)")
|
||||||
|
except BaseException as e:
|
||||||
|
logger.warning("Failed to connect MCP servers (will retry next message): {}", e)
|
||||||
|
finally:
|
||||||
|
state._mcp_connecting = False
|
||||||
|
|
||||||
|
|
||||||
async def reload_servers(state: Any, registry: ToolRegistry) -> dict[str, Any]:
|
async def reload_servers(state: Any, registry: ToolRegistry) -> dict[str, Any]:
|
||||||
"""Reconcile live MCP connections with the current config file."""
|
"""Reconcile live MCP connections with the current config file."""
|
||||||
async with _reload_lock(state):
|
async with _reload_lock(state):
|
||||||
|
if getattr(state, "_mcp_closing", False):
|
||||||
|
return {
|
||||||
|
"ok": False,
|
||||||
|
"message": "MCP connections are shutting down.",
|
||||||
|
"requires_restart": True,
|
||||||
|
}
|
||||||
try:
|
try:
|
||||||
from nanobot.config.loader import load_config, resolve_config_env_vars
|
from nanobot.config.loader import load_config, resolve_config_env_vars
|
||||||
|
|
||||||
@@ -1177,7 +1246,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)
|
||||||
_retire_server_stack(state, name)
|
await _close_server(state, name)
|
||||||
|
|
||||||
state._mcp_servers = next_servers
|
state._mcp_servers = next_servers
|
||||||
retry_missing = sorted(
|
retry_missing = sorted(
|
||||||
@@ -1187,9 +1256,17 @@ async def reload_servers(state: Any, registry: ToolRegistry) -> dict[str, Any]:
|
|||||||
)
|
)
|
||||||
to_connect_names = sorted(set(added) | set(changed) | set(retry_missing))
|
to_connect_names = sorted(set(added) | set(changed) | set(retry_missing))
|
||||||
to_connect = {name: next_servers[name] for name in to_connect_names}
|
to_connect = {name: next_servers[name] for name in to_connect_names}
|
||||||
connected: dict[str, AsyncExitStack] = {}
|
connected: dict[str, MCPConnection] = {}
|
||||||
if to_connect:
|
if to_connect:
|
||||||
connected = await connect_mcp_servers(to_connect, registry)
|
connected = await connect_mcp_servers(to_connect, registry)
|
||||||
|
if getattr(state, "_mcp_closing", False):
|
||||||
|
for connection in connected.values():
|
||||||
|
await connection.aclose()
|
||||||
|
return {
|
||||||
|
"ok": False,
|
||||||
|
"message": "MCP connections are shutting down.",
|
||||||
|
"requires_restart": True,
|
||||||
|
}
|
||||||
state._mcp_stacks.update(connected)
|
state._mcp_stacks.update(connected)
|
||||||
_attach_reconnect_handlers(state, registry, connected)
|
_attach_reconnect_handlers(state, registry, connected)
|
||||||
|
|
||||||
@@ -1323,6 +1400,8 @@ async def _refresh_terminated_server(
|
|||||||
stale_tool: Tool,
|
stale_tool: Tool,
|
||||||
) -> Tool | None:
|
) -> Tool | None:
|
||||||
async with _reload_lock(state):
|
async with _reload_lock(state):
|
||||||
|
if getattr(state, "_mcp_closing", False):
|
||||||
|
return None
|
||||||
cfg = state._mcp_servers.get(server_name)
|
cfg = state._mcp_servers.get(server_name)
|
||||||
if cfg is None:
|
if cfg is None:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
@@ -1341,9 +1420,13 @@ 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)
|
||||||
_retire_server_stack(state, server_name)
|
await _close_server(state, server_name)
|
||||||
|
|
||||||
connected = await connect_mcp_servers({server_name: cfg}, registry)
|
connected = await connect_mcp_servers({server_name: cfg}, registry)
|
||||||
|
if getattr(state, "_mcp_closing", False):
|
||||||
|
for connection in connected.values():
|
||||||
|
await connection.aclose()
|
||||||
|
return None
|
||||||
state._mcp_stacks.update(connected)
|
state._mcp_stacks.update(connected)
|
||||||
_attach_reconnect_handlers(state, registry, connected)
|
_attach_reconnect_handlers(state, registry, connected)
|
||||||
if server_name not in connected:
|
if server_name not in connected:
|
||||||
@@ -1378,17 +1461,24 @@ def _unregister_server_tools(state: Any, registry: ToolRegistry, server_name: st
|
|||||||
return removed
|
return removed
|
||||||
|
|
||||||
|
|
||||||
def _retire_server_stack(state: Any, server_name: str) -> None:
|
async def _close_server(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
|
||||||
retired = getattr(state, "_mcp_retired_stacks", None)
|
try:
|
||||||
if retired is not None:
|
await stack.aclose()
|
||||||
retired.append((server_name, stack))
|
except (RuntimeError, BaseExceptionGroup):
|
||||||
|
logger.debug("MCP server '{}' cleanup error (can be ignored)", server_name)
|
||||||
|
|
||||||
|
|
||||||
|
async def close_mcp_servers(state: Any) -> None:
|
||||||
|
"""Close every MCP connection while excluding reconnect and hot reload."""
|
||||||
|
state._mcp_closing = True
|
||||||
|
async with _reload_lock(state):
|
||||||
|
connections = list(state._mcp_stacks.items())
|
||||||
|
state._mcp_stacks.clear()
|
||||||
|
for name, connection in connections:
|
||||||
|
try:
|
||||||
|
await connection.aclose()
|
||||||
|
except (RuntimeError, BaseExceptionGroup):
|
||||||
|
logger.debug("MCP server '{}' cleanup error (can be ignored)", name)
|
||||||
|
|||||||
@@ -74,15 +74,6 @@ 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()
|
||||||
@@ -125,15 +116,26 @@ async def test_mcp_read_filter_drops_progress_notifications_without_progress_tok
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_close_mcp_swallows_sdk_cancel_scope_cleanup(tmp_path):
|
async def test_owned_mcp_connection_closes_from_its_owner_task():
|
||||||
loop = _make_loop(tmp_path)
|
close_requested = asyncio.Event()
|
||||||
stack = _CancelScopeStack()
|
ready = asyncio.Event()
|
||||||
loop._mcp_stacks["remote"] = stack # type: ignore[assignment]
|
tasks: dict[str, asyncio.Task] = {}
|
||||||
|
|
||||||
await loop.close_mcp()
|
async def own_connection() -> None:
|
||||||
|
tasks["open"] = asyncio.current_task() # type: ignore[assignment]
|
||||||
|
ready.set()
|
||||||
|
await close_requested.wait()
|
||||||
|
tasks["close"] = asyncio.current_task() # type: ignore[assignment]
|
||||||
|
|
||||||
assert stack.closed is True
|
owner = asyncio.create_task(own_connection())
|
||||||
assert loop._mcp_stacks == {}
|
connection = mcp_runtime._OwnedMCPConnection(owner, close_requested)
|
||||||
|
await ready.wait()
|
||||||
|
|
||||||
|
await connection.aclose()
|
||||||
|
|
||||||
|
assert tasks["open"] is owner
|
||||||
|
assert tasks["close"] is owner
|
||||||
|
assert tasks["close"] is not asyncio.current_task()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -241,9 +243,6 @@ 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"]
|
||||||
|
|
||||||
|
|
||||||
@@ -305,9 +304,6 @@ 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"]
|
||||||
|
|
||||||
|
|
||||||
@@ -403,15 +399,12 @@ 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 == []
|
assert closed == ["remote"]
|
||||||
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(
|
||||||
@@ -521,7 +514,4 @@ 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 == []
|
assert closed == ["remote"]
|
||||||
|
|
||||||
await loop.close_mcp()
|
|
||||||
assert closed == ["remote", "remote"]
|
|
||||||
|
|||||||
@@ -15,7 +15,6 @@ import asyncio
|
|||||||
import multiprocessing
|
import multiprocessing
|
||||||
import socket
|
import socket
|
||||||
import time
|
import time
|
||||||
from contextlib import suppress
|
|
||||||
from unittest.mock import MagicMock
|
from unittest.mock import MagicMock
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
@@ -70,7 +69,7 @@ async def _wait_for_server(url: str, timeout: float = 10.0) -> bool:
|
|||||||
deadline = time.monotonic() + timeout
|
deadline = time.monotonic() + timeout
|
||||||
while time.monotonic() < deadline:
|
while time.monotonic() < deadline:
|
||||||
try:
|
try:
|
||||||
async with httpx.AsyncClient(timeout=2.0) as client:
|
async with httpx.AsyncClient(timeout=2.0, trust_env=False) as client:
|
||||||
response = await client.get(
|
response = await client.get(
|
||||||
url,
|
url,
|
||||||
headers={"Accept": "text/event-stream"},
|
headers={"Accept": "text/event-stream"},
|
||||||
@@ -170,26 +169,30 @@ async def test_mcp_reconnect_after_session_timeout(tmp_path, mcp_server_url):
|
|||||||
)
|
)
|
||||||
loop = _make_loop(tmp_path, mcp_servers={"repro": cfg})
|
loop = _make_loop(tmp_path, mcp_servers={"repro": cfg})
|
||||||
|
|
||||||
await loop._connect_mcp()
|
await asyncio.create_task(loop._connect_mcp())
|
||||||
assert "repro" in loop._mcp_stacks
|
assert "repro" in loop._mcp_stacks
|
||||||
|
|
||||||
tool = loop.tools.get("mcp_repro_greet")
|
tool = loop.tools.get("mcp_repro_greet")
|
||||||
assert isinstance(tool, MCPToolWrapper)
|
assert isinstance(tool, MCPToolWrapper)
|
||||||
|
|
||||||
output = await tool.execute(name="first")
|
output = await asyncio.create_task(tool.execute(name="first"))
|
||||||
assert "Hello, first" in output
|
assert "Hello, first" in output
|
||||||
|
|
||||||
# Wait for the server-side idle timeout to terminate the session.
|
# Wait for the server-side idle timeout to terminate the session.
|
||||||
await asyncio.sleep(_IDLE_TIMEOUT_SECONDS + 1)
|
await asyncio.sleep(_IDLE_TIMEOUT_SECONDS + 1)
|
||||||
|
|
||||||
output = await tool.execute(name="second")
|
output = await asyncio.create_task(tool.execute(name="second"))
|
||||||
assert "Hello, second" in output
|
assert "Hello, second" in output
|
||||||
|
|
||||||
await loop.close_mcp()
|
await asyncio.create_task(loop.close_mcp())
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_mcp_reconnect_during_shutdown_does_not_crash(tmp_path, mcp_server_url):
|
async def test_mcp_reconnect_during_shutdown_does_not_crash(
|
||||||
|
tmp_path,
|
||||||
|
mcp_server_url,
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
):
|
||||||
"""Simulate the production crash: shutdown while reconnect is in flight."""
|
"""Simulate the production crash: shutdown while reconnect is in flight."""
|
||||||
cfg = MCPServerConfig(
|
cfg = MCPServerConfig(
|
||||||
type="streamableHttp",
|
type="streamableHttp",
|
||||||
@@ -199,15 +202,28 @@ async def test_mcp_reconnect_during_shutdown_does_not_crash(tmp_path, mcp_server
|
|||||||
)
|
)
|
||||||
loop = _make_loop(tmp_path, mcp_servers={"repro": cfg})
|
loop = _make_loop(tmp_path, mcp_servers={"repro": cfg})
|
||||||
|
|
||||||
await loop._connect_mcp()
|
await asyncio.create_task(loop._connect_mcp())
|
||||||
tool = loop.tools.get("mcp_repro_greet")
|
tool = loop.tools.get("mcp_repro_greet")
|
||||||
assert isinstance(tool, MCPToolWrapper)
|
assert isinstance(tool, MCPToolWrapper)
|
||||||
|
|
||||||
await tool.execute(name="first")
|
await asyncio.create_task(tool.execute(name="first"))
|
||||||
await asyncio.sleep(_IDLE_TIMEOUT_SECONDS + 1)
|
await asyncio.sleep(_IDLE_TIMEOUT_SECONDS + 1)
|
||||||
|
|
||||||
|
reconnect_started = asyncio.Event()
|
||||||
|
finish_reconnect = asyncio.Event()
|
||||||
|
real_connect = mcp_module.connect_mcp_servers
|
||||||
|
|
||||||
|
async def gated_connect(*args, **kwargs):
|
||||||
|
reconnect_started.set()
|
||||||
|
await finish_reconnect.wait()
|
||||||
|
return await real_connect(*args, **kwargs)
|
||||||
|
|
||||||
|
monkeypatch.setattr(mcp_module, "connect_mcp_servers", gated_connect)
|
||||||
call_task = asyncio.create_task(tool.execute(name="second"))
|
call_task = asyncio.create_task(tool.execute(name="second"))
|
||||||
loop.stop()
|
await asyncio.wait_for(reconnect_started.wait(), timeout=5)
|
||||||
|
close_task = asyncio.create_task(loop.close_mcp())
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
finish_reconnect.set()
|
||||||
|
|
||||||
unhandled: list[BaseException] = []
|
unhandled: list[BaseException] = []
|
||||||
|
|
||||||
@@ -219,13 +235,11 @@ async def test_mcp_reconnect_during_shutdown_does_not_crash(tmp_path, mcp_server
|
|||||||
asyncio.get_running_loop().set_exception_handler(capture_unhandled)
|
asyncio.get_running_loop().set_exception_handler(capture_unhandled)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
await asyncio.wait_for(asyncio.shield(call_task), timeout=15)
|
await asyncio.wait_for(asyncio.gather(call_task, close_task), timeout=15)
|
||||||
except asyncio.CancelledError:
|
except asyncio.CancelledError:
|
||||||
unhandled.append(asyncio.CancelledError("main task cancelled by leaked MCP cancel scope"))
|
unhandled.append(asyncio.CancelledError("main task cancelled by leaked MCP cancel scope"))
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
unhandled.append(exc)
|
unhandled.append(exc)
|
||||||
|
|
||||||
with suppress(Exception):
|
|
||||||
await loop.close_mcp()
|
|
||||||
|
|
||||||
assert not unhandled, f"Unhandled exception leaked during reconnect/shutdown: {unhandled[0]}"
|
assert not unhandled, f"Unhandled exception leaked during reconnect/shutdown: {unhandled[0]}"
|
||||||
|
assert loop._mcp_stacks == {}
|
||||||
|
|||||||
Reference in New Issue
Block a user