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
+7
View File
@@ -51,6 +51,13 @@ Main files:
- feeds tool results back into the model; - feeds tool results back into the model;
- stops when a final answer is produced or runtime limits are hit. - stops when a final answer is produced or runtime limits are hit.
MCP connections are application-owned infrastructure. Composition roots create
an `MCPProvider`, share its `ToolRegistry` with `AgentLoop`, await `connect()`
before use, and guarantee `aclose()` during shutdown; the loop does not manage
that lifecycle. `AgentLoop.from_config()` therefore requires a caller-owned
`ToolRegistry`; callers using MCP share it with their application-owned
`MCPProvider`.
Keep this split in mind when debugging. If a problem is about channel routing, session keys, workspace selection, or outbound delivery, start in `agent/loop.py`. If it is about provider calls, tool calls, streaming, or iteration limits, start in `agent/runner.py`. Keep this split in mind when debugging. If a problem is about channel routing, session keys, workspace selection, or outbound delivery, start in `agent/loop.py`. If it is about provider calls, tool calls, streaming, or iteration limits, start in `agent/runner.py`.
## Providers ## Providers
+1 -19
View File
@@ -42,29 +42,11 @@ def session_extra(metadata: Mapping[str, Any] | None) -> dict[str, Any]:
) )
async def connect_mcp(state: Any, tools: ToolRegistry) -> None:
await mcp_tools.connect_missing_servers(state, tools)
def mcp_runtime_status(state: Any) -> dict[str, mcp_tools.MCPRuntimeStatus]:
return mcp_tools.runtime_status(state)
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:
if msg.metadata.get(INBOUND_META_RUNTIME_CONTROL) == RUNTIME_CONTROL_SESSION_DISCARD: if msg.metadata.get(INBOUND_META_RUNTIME_CONTROL) == RUNTIME_CONTROL_SESSION_DISCARD:
await state.discard_session(msg.session_key) await state.discard_session(msg.session_key)
return True return True
for handler in ( return await image_generation_tools.handle_runtime_control(state, msg, tools)
image_generation_tools.handle_runtime_control,
mcp_tools.handle_runtime_control,
):
if await handler(state, msg, tools):
return True
return False
class ContextBuilder: class ContextBuilder:
+18 -33
View File
@@ -95,11 +95,9 @@ from nanobot.utils.runtime import (
) )
if TYPE_CHECKING: if TYPE_CHECKING:
from nanobot.agent.tools.mcp import MCPConnection, MCPRuntimeStatus
from nanobot.config.schema import ( from nanobot.config.schema import (
ChannelsConfig, ChannelsConfig,
Config, Config,
MCPServerConfig,
ProviderConfig, ProviderConfig,
ToolsConfig, ToolsConfig,
) )
@@ -271,7 +269,7 @@ class AgentLoop:
cron_service: CronService | None = None, cron_service: CronService | None = None,
restrict_to_workspace: bool = False, restrict_to_workspace: bool = False,
session_manager: SessionManager | None = None, session_manager: SessionManager | None = None,
mcp_servers: dict[str, MCPServerConfig] | None = None, tool_registry: ToolRegistry | None = None,
channels_config: ChannelsConfig | None = None, channels_config: ChannelsConfig | None = None,
timezone: str | None = None, timezone: str | None = None,
session_ttl_minutes: int = 0, session_ttl_minutes: int = 0,
@@ -379,7 +377,7 @@ class AgentLoop:
self.context = ContextBuilder(workspace, timezone=timezone, disabled_skills=disabled_skills) self.context = ContextBuilder(workspace, timezone=timezone, disabled_skills=disabled_skills)
self.sessions = session_manager or SessionManager(workspace) self.sessions = session_manager or SessionManager(workspace)
self.sessions.set_file_cap_archiver(self.context.memory.raw_archive) self.sessions.set_file_cap_archiver(self.context.memory.raw_archive)
self.tools = ToolRegistry() self.tools = tool_registry if tool_registry is not None else ToolRegistry()
# One file-read/write tracker per logical session. The tool registry is # One file-read/write tracker per logical session. The tool registry is
# shared by this loop, so tools resolve the active state via contextvars. # shared by this loop, so tools resolve the active state via contextvars.
self._file_state_store = FileStateStore() self._file_state_store = FileStateStore()
@@ -399,15 +397,11 @@ 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_stacks: dict[str, MCPConnection] = {}
self._mcp_runtime_statuses: dict[str, MCPRuntimeStatus] = {}
self._mcp_connecting = False
self._runtime_context_providers: list[RuntimeContextProvider] = [] self._runtime_context_providers: list[RuntimeContextProvider] = []
self._active_tasks: dict[str, set[asyncio.Task[Any]]] = {} self._active_tasks: dict[str, set[asyncio.Task[Any]]] = {}
self._discarding_sessions: set[str] = set() self._discarding_sessions: set[str] = set()
self._background_tasks: set[asyncio.Task[Any]] = set() self._background_tasks: set[asyncio.Task[Any]] = set()
self._close_mcp_lock = asyncio.Lock() self._close_lock = asyncio.Lock()
self._session_locks: weakref.WeakValueDictionary[str, asyncio.Lock] = ( self._session_locks: weakref.WeakValueDictionary[str, asyncio.Lock] = (
weakref.WeakValueDictionary() weakref.WeakValueDictionary()
) )
@@ -464,10 +458,15 @@ class AgentLoop:
cls, cls,
config: Config, config: Config,
bus: MessageBus | None = None, bus: MessageBus | None = None,
*,
tool_registry: ToolRegistry,
**extra: Any, **extra: Any,
) -> AgentLoop: ) -> AgentLoop:
"""Create an AgentLoop from config with the common parameter set. """Create an AgentLoop from config with the common parameter set.
The tool registry is caller-owned so application composition can share
it with infrastructure such as an ``MCPProvider``.
Extra keyword arguments are forwarded to ``AgentLoop.__init__``, Extra keyword arguments are forwarded to ``AgentLoop.__init__``,
allowing callers to override or extend the standard config-derived allowing callers to override or extend the standard config-derived
parameters (e.g. ``cron_service``, ``session_manager``). parameters (e.g. ``cron_service``, ``session_manager``).
@@ -486,8 +485,6 @@ class AgentLoop:
config, config,
provider_snapshot_loader, provider_snapshot_loader,
) )
from nanobot.agent.plugins import agent_plugin_mcp_servers
return cls( return cls(
bus=bus, bus=bus,
provider=provider, provider=provider,
@@ -502,7 +499,6 @@ class AgentLoop:
provider_retry_mode=defaults.provider_retry_mode, provider_retry_mode=defaults.provider_retry_mode,
tool_hint_max_length=defaults.tool_hint_max_length, tool_hint_max_length=defaults.tool_hint_max_length,
restrict_to_workspace=config.tools.restrict_to_workspace, restrict_to_workspace=config.tools.restrict_to_workspace,
mcp_servers=agent_plugin_mcp_servers(config.workspace_path, config.tools.mcp_servers),
channels_config=config.channels, channels_config=config.channels,
timezone=defaults.timezone, timezone=defaults.timezone,
unified_session=defaults.unified_session, unified_session=defaults.unified_session,
@@ -517,6 +513,7 @@ class AgentLoop:
restart_mode=config.gateway.restart_mode, restart_mode=config.gateway.restart_mode,
provider_snapshot_loader=provider_snapshot_loader, provider_snapshot_loader=provider_snapshot_loader,
preset_snapshot_loader=preset_snapshot_loader, preset_snapshot_loader=preset_snapshot_loader,
tool_registry=tool_registry,
**extra, **extra,
) )
@@ -643,14 +640,6 @@ class AgentLoop:
logger.info("Registered {} tools: {}", len(registered), registered) logger.info("Registered {} tools: {}", len(registered), registered)
async def _connect_mcp(self) -> None:
"""Connect configured MCP servers."""
await agent_context.connect_mcp(self, self.tools)
def mcp_runtime_status(self) -> dict[str, MCPRuntimeStatus]:
"""Return connection state learned from real MCP runtime attempts."""
return agent_context.mcp_runtime_status(self)
def register_runtime_context_provider( def register_runtime_context_provider(
self, self,
provider: RuntimeContextProvider, provider: RuntimeContextProvider,
@@ -1162,7 +1151,6 @@ class AgentLoop:
"""Run the agent loop, dispatching messages as tasks to stay responsive to /stop.""" """Run the agent loop, dispatching messages as tasks to stay responsive to /stop."""
self._running = True self._running = True
try: try:
await self._connect_mcp()
logger.info("Agent loop started") logger.info("Agent loop started")
while self._running: while self._running:
@@ -1253,8 +1241,7 @@ class AgentLoop:
active_tasks.add(task) active_tasks.add(task)
task.add_done_callback(active_tasks.discard) task.add_done_callback(active_tasks.discard)
finally: finally:
# MCP stdio transports use AnyIO cancel scopes; close them from the task that opened them. await self.aclose()
await self.close_mcp()
async def _dispatch(self, msg: InboundMessage) -> None: async def _dispatch(self, msg: InboundMessage) -> None:
"""Process a message: per-session serial, cross-session concurrent.""" """Process a message: per-session serial, cross-session concurrent."""
@@ -1372,24 +1359,24 @@ class AgentLoop:
await delivery.idle() await delivery.idle()
await self._publish_next_deferred_automation_turn(session_key) await self._publish_next_deferred_automation_turn(session_key)
async def close_mcp(self) -> None: async def aclose(self) -> None:
"""Stop active work, then close exec, subagent, and MCP resources. """Stop active work, then close resources owned by the agent loop.
Resource teardown must still run if cancellation interrupts task draining. Resource teardown must still run if cancellation interrupts task draining.
Gateway shutdown deliberately bounds this coroutine, so keeping the cleanup Gateway shutdown deliberately bounds this coroutine, so keeping the cleanup
phase in ``finally`` prevents a timed-out background task from leaving phase in ``finally`` prevents a timed-out background task from leaving
subprocess transports alive after the event loop closes. subprocess transports alive after the event loop closes.
""" """
# The agent loop closes itself from ``run()`` while gateway shutdown also # The loop closes itself from ``run()`` while application shutdown also
# performs a guaranteed final close. Serialize those owners so they cannot # performs a guaranteed final close. Serialize those owners so they cannot
# tear down the same subprocess transports concurrently. # tear down the same resources concurrently.
close_lock = getattr(self, "_close_mcp_lock", None) close_lock = getattr(self, "_close_lock", None)
if close_lock is None: if close_lock is None:
close_lock = self._close_mcp_lock = asyncio.Lock() close_lock = self._close_lock = asyncio.Lock()
async with close_lock: async with close_lock:
await self._close_mcp_unlocked() await self._aclose_unlocked()
async def _close_mcp_unlocked(self) -> None: async def _aclose_unlocked(self) -> None:
errors: list[BaseException] = [] errors: list[BaseException] = []
active_task_groups = getattr(self, "_active_tasks", {}) active_task_groups = getattr(self, "_active_tasks", {})
active_tasks = tuple({task for tasks in active_task_groups.values() for task in tasks}) active_tasks = tuple({task for tasks in active_task_groups.values() for task in tasks})
@@ -1412,7 +1399,6 @@ class AgentLoop:
cleanup_steps = ( cleanup_steps = (
self.subagents.close, self.subagents.close,
self._exec_session_manager.close_all, self._exec_session_manager.close_all,
lambda: agent_context.close_mcp(self),
) )
for cleanup in cleanup_steps: for cleanup in cleanup_steps:
try: try:
@@ -2301,7 +2287,6 @@ class AgentLoop:
"""Process an external message directly and return the outbound payload.""" """Process an external message directly and return the outbound payload."""
if channel == "system": if channel == "system":
raise ValueError("channel 'system' is reserved for internal messages") raise ValueError("channel 'system' is reserved for internal messages")
await self._connect_mcp()
metadata: dict[str, Any] = {} metadata: dict[str, Any] = {}
if not persist_user_message: if not persist_user_message:
metadata[turn_continuation.SKIP_USER_PERSIST_META] = True metadata[turn_continuation.SKIP_USER_PERSIST_META] = True
+377 -396
View File
@@ -1,4 +1,6 @@
"""MCP client: connects to MCP servers and wraps their tools as native nanobot tools.""" """MCP client and dynamic tool-provider lifecycle."""
from __future__ import annotations
import asyncio import asyncio
import hashlib import hashlib
@@ -7,23 +9,15 @@ import os
import re import re
import shutil import shutil
import urllib.parse import urllib.parse
from collections.abc import AsyncIterator, Awaitable, Callable from collections.abc import AsyncIterator, Awaitable, Callable, Iterable, Mapping
from contextlib import AsyncExitStack, suppress from contextlib import AsyncExitStack, suppress
from typing import TYPE_CHECKING, Any, Literal, Mapping, Protocol, cast from typing import TYPE_CHECKING, Any, Literal, Protocol, cast
from weakref import WeakKeyDictionary
import httpx import httpx
from loguru import logger from loguru import logger
from nanobot.agent.tools.base import Tool, ToolResult from nanobot.agent.tools.base import Tool, ToolResult
from nanobot.agent.tools.registry import ToolRegistry from nanobot.agent.tools.registry import ToolRegistry
from nanobot.bus.events import (
INBOUND_META_RUNTIME_CONTROL,
RUNTIME_CONTROL_ACK,
RUNTIME_CONTROL_MCP_RELOAD,
InboundMessage,
)
from nanobot.bus.queue import MessageBus
from nanobot.security.network import ( from nanobot.security.network import (
PinnedDNSAsyncTransport, PinnedDNSAsyncTransport,
env_proxy_applies_to_url, env_proxy_applies_to_url,
@@ -39,7 +33,7 @@ if TYPE_CHECKING:
from mcp.types import Tool as MCPToolDefinition from mcp.types import Tool as MCPToolDefinition
from nanobot.agent.tools.mcp_oauth import MCPOAuthHandlers from nanobot.agent.tools.mcp_oauth import MCPOAuthHandlers
from nanobot.config.schema import MCPServerConfig from nanobot.config.schema import Config, MCPServerConfig
# Transient connection errors that warrant a single retry. # Transient connection errors that warrant a single retry.
# These typically happen when an MCP server restarts or a network # These typically happen when an MCP server restarts or a network
@@ -60,18 +54,37 @@ _WINDOWS_SHELL_LAUNCHERS: frozenset[str] = frozenset(("npx", "npm", "pnpm", "yar
# Characters allowed in tool names by model providers (Anthropic, OpenAI, etc.). # Characters allowed in tool names by model providers (Anthropic, OpenAI, etc.).
# Replace anything outside [a-zA-Z0-9_-] with underscore and collapse runs. # Replace anything outside [a-zA-Z0-9_-] with underscore and collapse runs.
_SANITIZE_RE = re.compile(r"_+") _SANITIZE_RE = re.compile(r"_+")
_RELOAD_LOCKS: WeakKeyDictionary[Any, asyncio.Lock] = WeakKeyDictionary()
_ReconnectCallback = Callable[[str, str, Tool], Awaitable[Tool | None]] _ReconnectCallback = Callable[[str, str, Tool], Awaitable[Tool | None]]
MCPServerLoader = Callable[[], Mapping[str, "MCPServerConfig"]]
MCPRuntimeStatus = Literal["connecting", "connected", "failed"] MCPRuntimeStatus = Literal["connecting", "connected", "failed"]
_MCP_RUNTIME_STATUSES: frozenset[MCPRuntimeStatus] = frozenset(
("connecting", "connected", "failed")
)
class MCPConnection(Protocol): class MCPConnection(Protocol):
async def aclose(self) -> None: ... async def aclose(self) -> None: ...
async def _close_mcp_connection(name: str, connection: MCPConnection) -> None:
try:
await connection.aclose()
except asyncio.CancelledError:
if task_is_cancelling():
raise
logger.debug("MCP server '{}' cleanup error (can be ignored)", name)
except (RuntimeError, BaseExceptionGroup):
logger.debug("MCP server '{}' cleanup error (can be ignored)", name)
async def _close_mcp_connections(connections: Mapping[str, MCPConnection]) -> None:
cancellation: asyncio.CancelledError | None = None
for name, connection in connections.items():
try:
await _close_mcp_connection(name, connection)
except asyncio.CancelledError as exc:
cancellation = cancellation or exc
if cancellation is not None:
raise cancellation
class _OwnedMCPConnection: class _OwnedMCPConnection:
"""Close an MCP transport from the task that originally opened it.""" """Close an MCP transport from the task that originally opened it."""
@@ -492,11 +505,11 @@ class _MCPWrapperBase(Tool):
"""Common reconnect handling for wrappers bound to one MCP server session.""" """Common reconnect handling for wrappers bound to one MCP server session."""
_plugin_discoverable = False _plugin_discoverable = False
_session: "ClientSession" _session: ClientSession
_server_name: str _server_name: str
_name: str _name: str
def _set_mcp_connection(self, session: "ClientSession", server_name: str) -> None: def _set_mcp_connection(self, session: ClientSession, server_name: str) -> None:
self._session = session self._session = session
self._server_name = server_name self._server_name = server_name
self._reconnect: _ReconnectCallback | None = None self._reconnect: _ReconnectCallback | None = None
@@ -586,9 +599,9 @@ class MCPToolWrapper(_MCPWrapperBase):
def __init__( def __init__(
self, self,
session: "ClientSession", session: ClientSession,
server_name: str, server_name: str,
tool_def: "MCPToolDefinition", tool_def: MCPToolDefinition,
tool_timeout: int = 30, tool_timeout: int = 30,
): ):
self._set_mcp_connection(session, server_name) self._set_mcp_connection(session, server_name)
@@ -748,9 +761,9 @@ class MCPResourceWrapper(_MCPWrapperBase):
def __init__( def __init__(
self, self,
session: "ClientSession", session: ClientSession,
server_name: str, server_name: str,
resource_def: "Resource", resource_def: Resource,
resource_timeout: int = 30, resource_timeout: int = 30,
): ):
self._set_mcp_connection(session, server_name) self._set_mcp_connection(session, server_name)
@@ -852,9 +865,9 @@ class MCPPromptWrapper(_MCPWrapperBase):
def __init__( def __init__(
self, self,
session: "ClientSession", session: ClientSession,
server_name: str, server_name: str,
prompt_def: "Prompt", prompt_def: Prompt,
prompt_timeout: int = 30, prompt_timeout: int = 30,
): ):
self._set_mcp_connection(session, server_name) self._set_mcp_connection(session, server_name)
@@ -985,10 +998,10 @@ class MCPPromptWrapper(_MCPWrapperBase):
async def connect_mcp_servers( async def connect_mcp_servers(
mcp_servers: "dict[str, MCPServerConfig]", mcp_servers: dict[str, MCPServerConfig],
registry: ToolRegistry, registry: ToolRegistry,
*, *,
oauth_handlers: Mapping[str, "MCPOAuthHandlers"] | None = None, oauth_handlers: Mapping[str, MCPOAuthHandlers] | None = None,
) -> dict[str, MCPConnection]: ) -> 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.
@@ -1002,7 +1015,7 @@ async def connect_mcp_servers(
from mcp.client.streamable_http import streamable_http_client from mcp.client.streamable_http import streamable_http_client
async def open_single_server( async def open_single_server(
name: str, cfg: "MCPServerConfig", server_stack: AsyncExitStack name: str, cfg: MCPServerConfig, server_stack: AsyncExitStack
) -> bool: ) -> bool:
try: try:
transport_type = cfg.type transport_type = cfg.type
@@ -1244,7 +1257,7 @@ async def connect_mcp_servers(
return False return False
async def connect_single_server( async def connect_single_server(
name: str, cfg: "MCPServerConfig" name: str, cfg: MCPServerConfig
) -> tuple[str, MCPConnection | None]: ) -> tuple[str, MCPConnection | None]:
loop = asyncio.get_running_loop() loop = asyncio.get_running_loop()
ready: asyncio.Future[bool] = loop.create_future() ready: asyncio.Future[bool] = loop.create_future()
@@ -1282,15 +1295,29 @@ async def connect_mcp_servers(
return name, connection return name, connection
server_stacks: dict[str, MCPConnection] = {} server_stacks: dict[str, MCPConnection] = {}
attempted_names: list[str] = []
for name, cfg in mcp_servers.items(): try:
for name, cfg in mcp_servers.items():
attempted_names.append(name)
try:
result = await connect_single_server(name, cfg)
except Exception as e:
_log_mcp_connection_failure(name, e)
continue
if result[1] is not None:
server_stacks[result[0]] = result[1]
except BaseException:
# Callers can bound readiness/reload with a timeout. If cancellation
# interrupts a later server, ownership of earlier connections has not
# transferred yet, so roll the whole batch back before propagating it.
for name in attempted_names:
_unregister_server_tools(registry, name)
try: try:
result = await connect_single_server(name, cfg) await _close_mcp_connections(server_stacks)
except Exception as e: except BaseException as cleanup_exc:
_log_mcp_connection_failure(name, e) logger.debug("MCP batch rollback cleanup error (can be ignored): {}", cleanup_exc)
continue raise
if result[1] is not None:
server_stacks[result[0]] = result[1]
return server_stacks return server_stacks
@@ -1301,369 +1328,357 @@ def session_extra(metadata: Mapping[str, Any] | None) -> dict[str, Any]:
return {"mcp_presets": mcp_presets} if isinstance(mcp_presets, list) and mcp_presets else {} return {"mcp_presets": mcp_presets} if isinstance(mcp_presets, list) and mcp_presets else {}
def _runtime_status_store( def _configured_servers(config: Config) -> dict[str, MCPServerConfig]:
state: Any, from nanobot.agent.plugins import agent_plugin_mcp_servers
*,
create: bool = False, return agent_plugin_mcp_servers(
) -> dict[str, MCPRuntimeStatus] | None: config.workspace_path,
raw_statuses: object = getattr(state, "_mcp_runtime_statuses", None) config.tools.mcp_servers,
if isinstance(raw_statuses, dict): )
return cast(dict[str, MCPRuntimeStatus], raw_statuses)
if not create:
return None
statuses: dict[str, MCPRuntimeStatus] = {}
state._mcp_runtime_statuses = statuses
return statuses
def runtime_status(state: Any) -> dict[str, MCPRuntimeStatus]: def _load_current_servers() -> dict[str, MCPServerConfig]:
"""Return the latest connection-attempt result for configured MCP servers.""" from nanobot.config.loader import load_config, resolve_config_env_vars
statuses = _runtime_status_store(state)
raw_configured: object = getattr(state, "_mcp_servers", None) return _configured_servers(resolve_config_env_vars(load_config()))
if statuses is None or not isinstance(raw_configured, dict):
return {}
configured = cast(dict[str, Any], raw_configured)
return {
name: status
for name, status in statuses.items()
if name in configured and status in _MCP_RUNTIME_STATUSES
}
def _set_runtime_status( class MCPProvider:
state: Any, """Own configured MCP connections and their dynamic tool registrations."""
server_names: Mapping[str, Any] | set[str] | list[str] | tuple[str, ...],
status: MCPRuntimeStatus,
) -> None:
statuses = _runtime_status_store(state, create=True)
assert statuses is not None
for name in server_names:
statuses[name] = status
def __init__(
self,
servers: Mapping[str, MCPServerConfig],
registry: ToolRegistry,
*,
server_loader: MCPServerLoader | None = None,
) -> None:
self._servers = dict(servers)
self._registry = registry
self._server_loader = server_loader or _load_current_servers
self._connections: dict[str, MCPConnection] = {}
self._runtime_statuses: dict[str, MCPRuntimeStatus] = {}
self._lock = asyncio.Lock()
self._closing = False
def _record_connection_result( @classmethod
state: Any, def from_config(
attempted: Mapping[str, Any] | set[str] | list[str] | tuple[str, ...], cls,
connected: Mapping[str, Any] | set[str] | list[str] | tuple[str, ...], config: Config,
) -> None: registry: ToolRegistry,
attempted_names = set(attempted) *,
connected_names = set(connected) server_loader: MCPServerLoader | None = None,
_set_runtime_status(state, connected_names, "connected") ) -> MCPProvider:
_set_runtime_status(state, attempted_names - connected_names, "failed") return cls(
_configured_servers(config),
registry,
server_loader=server_loader,
)
@property
def configured_server_names(self) -> set[str]:
return set(self._servers)
async def connect_missing_servers(state: Any, registry: ToolRegistry) -> None: @property
"""Connect configured MCP servers that are not currently live.""" def connected_server_names(self) -> set[str]:
async with _reload_lock(state): return set(self._connections)
if getattr(state, "_mcp_closing", False):
return def runtime_status(self) -> dict[str, MCPRuntimeStatus]:
configured_missing = { """Return the latest connection-attempt result for configured servers."""
name: cfg for name, cfg in state._mcp_servers.items() if name not in state._mcp_stacks return {
name: status
for name, status in self._runtime_statuses.items()
if name in self._servers
} }
oauth_servers = {
name: cfg def _set_runtime_status(
for name, cfg in configured_missing.items() self,
if getattr(cfg, "auth", None) == "oauth" server_names: Iterable[str],
} status: MCPRuntimeStatus,
authorization_pending: set[str] = set() ) -> None:
if oauth_servers: for name in server_names:
self._runtime_statuses[name] = status
def _record_connection_result(
self,
attempted: Iterable[str],
connected: Iterable[str],
) -> None:
attempted_names = set(attempted)
connected_names = set(connected)
self._set_runtime_status(connected_names, "connected")
self._set_runtime_status(attempted_names - connected_names, "failed")
async def connect(self) -> None:
"""Connect configured servers that are not currently live."""
async with self._lock:
if self._closing:
return
configured_missing = {
name: cfg
for name, cfg in self._servers.items()
if name not in self._connections
}
oauth_servers = {
name: cfg
for name, cfg in configured_missing.items()
if cfg.auth == "oauth"
}
authorization_pending: set[str] = set()
if oauth_servers:
from nanobot.agent.tools.mcp_oauth import mcp_oauth_has_credentials
authorization_pending = {
name
for name, cfg in oauth_servers.items()
if not mcp_oauth_has_credentials(name, cfg.url)
}
for name in authorization_pending:
self._runtime_statuses.pop(name, None)
missing_servers = {
name: cfg
for name, cfg in configured_missing.items()
if name not in authorization_pending
}
if not missing_servers:
return
self._set_runtime_status(missing_servers, "connecting")
try:
connected = await connect_mcp_servers(missing_servers, self._registry)
if self._closing:
await _close_mcp_connections(connected)
return
self._connections.update(connected)
self._record_connection_result(missing_servers, connected)
self._attach_reconnect_handlers(connected)
if connected:
logger.info("MCP connected servers: {}", sorted(connected))
else:
logger.warning(
"No MCP servers connected successfully "
"(will retry on the next readiness check)"
)
except asyncio.CancelledError:
self._set_runtime_status(missing_servers, "failed")
if task_is_cancelling():
raise
logger.warning(
"MCP connection cancelled (will retry on the next readiness check)"
)
except BaseException as exc:
self._set_runtime_status(missing_servers, "failed")
logger.warning(
"Failed to connect MCP servers "
"(will retry on the next readiness check): {}",
exc,
)
async def reload(self) -> dict[str, Any]:
"""Reconcile live MCP connections with the current configuration."""
async with self._lock:
if self._closing:
return self._closing_result()
try:
next_servers = dict(self._server_loader())
except Exception as exc:
logger.warning("MCP hot reload could not read config: {}", exc)
return {
"ok": False,
"message": "Could not reload MCP config. Restart nanobot to pick up changes.",
"requires_restart": True,
"error": str(exc),
}
current_servers = dict(self._servers)
current_names = set(current_servers)
next_names = set(next_servers)
from nanobot.agent.tools.mcp_oauth import mcp_oauth_has_credentials from nanobot.agent.tools.mcp_oauth import mcp_oauth_has_credentials
authorization_pending = { authorization_pending = {
name name
for name, cfg in oauth_servers.items() for name, cfg in next_servers.items()
if not mcp_oauth_has_credentials(name, cfg.url) if cfg.auth == "oauth" and not mcp_oauth_has_credentials(name, cfg.url)
} }
statuses = _runtime_status_store(state) removed = sorted(current_names - next_names)
if statuses is not None: added = sorted(next_names - current_names)
for name in authorization_pending: changed = sorted(
statuses.pop(name, None) name
missing_servers = { for name in current_names & next_names
name: cfg if _server_signature(current_servers[name])
for name, cfg in configured_missing.items() != _server_signature(next_servers[name])
if name not in authorization_pending )
}
if state._mcp_connecting or not missing_servers: tools_removed = 0
return for name in [*removed, *changed]:
state._mcp_connecting = True tools_removed += _unregister_server_tools(self._registry, name)
_set_runtime_status(state, missing_servers, "connecting") await self._close_server(name)
try:
connected = await connect_mcp_servers(missing_servers, registry) for name in [*removed, *authorization_pending]:
if getattr(state, "_mcp_closing", False): self._runtime_statuses.pop(name, None)
for connection in connected.values():
await connection.aclose()
return
state._mcp_stacks.update(connected)
_record_connection_result(state, missing_servers, connected)
_attach_reconnect_handlers(state, registry, connected)
if connected:
logger.info("MCP connected servers: {}", sorted(connected))
else:
logger.warning("No MCP servers connected successfully (will retry next message)")
except asyncio.CancelledError:
if task_is_cancelling():
raise
_set_runtime_status(state, missing_servers, "failed")
logger.warning("MCP connection cancelled (will retry next message)")
except BaseException as e:
_set_runtime_status(state, missing_servers, "failed")
logger.warning("Failed to connect MCP servers (will retry next message): {}", e)
finally:
state._mcp_connecting = False
self._servers = next_servers
retry_missing = sorted(
name
for name in next_names
if name not in self._connections
and name not in set(added) | set(changed)
and name not in authorization_pending
)
to_connect_names = sorted(
(set(added) | set(changed) | set(retry_missing))
- authorization_pending
)
to_connect = {name: next_servers[name] for name in to_connect_names}
connected: dict[str, MCPConnection] = {}
if to_connect:
self._set_runtime_status(to_connect, "connecting")
try:
connected = await connect_mcp_servers(to_connect, self._registry)
except BaseException:
self._set_runtime_status(to_connect, "failed")
raise
if self._closing:
await _close_mcp_connections(connected)
return self._closing_result()
self._connections.update(connected)
self._record_connection_result(to_connect, connected)
self._attach_reconnect_handlers(connected)
async def reload_servers(state: Any, registry: ToolRegistry) -> dict[str, Any]: failed = sorted(set(to_connect) - set(connected))
"""Reconcile live MCP connections with the current config file.""" unchanged = not removed and not added and not changed and not retry_missing
async with _reload_lock(state): ok = not failed
if getattr(state, "_mcp_closing", False): if failed:
return { message = (
"ok": False, "MCP config reloaded, but some servers did not connect: "
"message": "MCP connections are shutting down.", + ", ".join(failed)
"requires_restart": True, )
} elif unchanged:
try: message = "MCP config is already live."
from nanobot.agent.plugins import agent_plugin_mcp_servers elif retry_missing and not added and not changed and not removed:
from nanobot.config.loader import load_config, resolve_config_env_vars message = "MCP connections refreshed without restarting nanobot."
else:
message = "MCP config reloaded without restarting nanobot."
config = resolve_config_env_vars(load_config()) logger.info(
next_servers = agent_plugin_mcp_servers( "MCP hot reload: added={} changed={} removed={} retried={} "
config.workspace_path, "connected={} failed={} tools_removed={}",
config.tools.mcp_servers, added,
changed,
removed,
retry_missing,
sorted(connected),
failed,
tools_removed,
) )
except Exception as exc:
logger.warning("MCP hot reload could not read config: {}", exc)
return { return {
"ok": False, "ok": ok,
"message": "Could not reload MCP config. Restart nanobot to pick up changes.", "message": message,
"requires_restart": True, "added": added,
"error": str(exc), "changed": changed,
"removed": removed,
"retried": retry_missing,
"connected": sorted(self._connections),
"configured": sorted(self._servers),
"failed": failed,
"tools_removed": tools_removed,
"requires_restart": False,
} }
current_servers = dict(state._mcp_servers) @staticmethod
current_names = set(current_servers) def _closing_result() -> dict[str, Any]:
next_names = set(next_servers)
from nanobot.agent.tools.mcp_oauth import mcp_oauth_has_credentials
authorization_pending = {
name
for name, cfg in next_servers.items()
if cfg.auth == "oauth" and not mcp_oauth_has_credentials(name, cfg.url)
}
removed = sorted(current_names - next_names)
added = sorted(next_names - current_names)
changed = sorted(
name
for name in current_names & next_names
if _server_signature(current_servers[name]) != _server_signature(next_servers[name])
)
tools_removed = 0
for name in [*removed, *changed]:
tools_removed += _unregister_server_tools(registry, name)
await _close_server(state, name)
runtime_statuses = _runtime_status_store(state)
if runtime_statuses is not None:
for name in [*removed, *authorization_pending]:
runtime_statuses.pop(name, None)
state._mcp_servers = next_servers
retry_missing = sorted(
name
for name in next_names
if name not in state._mcp_stacks
and name not in set(added) | set(changed)
and name not in authorization_pending
)
to_connect_names = sorted(
(set(added) | set(changed) | set(retry_missing)) - authorization_pending
)
to_connect = {name: next_servers[name] for name in to_connect_names}
connected: dict[str, MCPConnection] = {}
if to_connect:
_set_runtime_status(state, to_connect, "connecting")
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)
_record_connection_result(state, to_connect, connected)
_attach_reconnect_handlers(state, registry, connected)
failed = sorted(set(to_connect) - set(connected))
unchanged = not removed and not added and not changed and not retry_missing
ok = not failed
if failed:
message = "MCP config reloaded, but some servers did not connect: " + ", ".join(failed)
elif unchanged:
message = "MCP config is already live."
elif retry_missing and not added and not changed and not removed:
message = "MCP connections refreshed without restarting nanobot."
else:
message = "MCP config reloaded without restarting nanobot."
logger.info(
"MCP hot reload: added={} changed={} removed={} retried={} connected={} failed={} tools_removed={}",
added,
changed,
removed,
retry_missing,
sorted(connected),
failed,
tools_removed,
)
return {
"ok": ok,
"message": message,
"added": added,
"changed": changed,
"removed": removed,
"retried": retry_missing,
"connected": sorted(state._mcp_stacks),
"configured": sorted(state._mcp_servers),
"failed": failed,
"tools_removed": tools_removed,
"requires_restart": False,
}
async def request_mcp_reload(
bus: MessageBus,
*,
timeout: float = 15.0,
) -> dict[str, Any]:
"""Ask the running agent loop to reconcile live MCP connections."""
loop = asyncio.get_running_loop()
ack: asyncio.Future[dict[str, Any]] = loop.create_future()
await bus.publish_inbound(
InboundMessage(
channel="system",
sender_id="webui-settings",
chat_id="runtime",
content=RUNTIME_CONTROL_MCP_RELOAD,
metadata={
INBOUND_META_RUNTIME_CONTROL: RUNTIME_CONTROL_MCP_RELOAD,
RUNTIME_CONTROL_ACK: ack,
},
)
)
try:
result = await asyncio.wait_for(ack, timeout=timeout)
except asyncio.TimeoutError:
return { return {
"ok": False, "ok": False,
"message": "MCP hot reload timed out. Restart nanobot to pick up changes.", "message": "MCP connections are shutting down.",
"requires_restart": True, "requires_restart": True,
} }
return result if isinstance(cast(object, result), dict) else {
"ok": False,
"message": "MCP hot reload returned an unexpected response.",
"requires_restart": True,
}
def _attach_reconnect_handlers(self, server_names: Iterable[str]) -> None:
async def reconnect(
server_name: str,
tool_name: str,
stale_tool: Tool,
) -> Tool | None:
return await self._refresh_terminated_server(
server_name,
tool_name,
stale_tool,
)
async def handle_runtime_control(state: Any, msg: InboundMessage, registry: ToolRegistry) -> bool: for server_name in server_names:
metadata = msg.metadata if isinstance(cast(object, msg.metadata), dict) else {} for tool_name in list(self._registry.tool_names):
control = metadata.get(INBOUND_META_RUNTIME_CONTROL) tool = self._registry.get(tool_name)
if control != RUNTIME_CONTROL_MCP_RELOAD: if not _tool_belongs_to_server(tool, tool_name, server_name):
return False continue
if isinstance(tool, _MCPWrapperBase):
tool.set_reconnect_handler(reconnect)
ack = metadata.get(RUNTIME_CONTROL_ACK) async def _refresh_terminated_server(
try: self,
result = await reload_servers(state, registry) server_name: str,
except Exception as exc: tool_name: str,
logger.exception("MCP hot reload failed") stale_tool: Tool,
result = { ) -> Tool | None:
"ok": False, async with self._lock:
"message": "MCP hot reload failed. Restart nanobot to pick up changes.", if self._closing:
"requires_restart": True, return None
"error": str(exc), cfg = self._servers.get(server_name)
} if cfg is None:
if isinstance(ack, asyncio.Future) and not ack.done(): logger.warning(
cast(asyncio.Future[dict[str, Any]], ack).set_result(result) "MCP server '{}' session terminated but is no longer configured",
return True server_name,
)
return None
current_tool = self._registry.get(tool_name)
if (
current_tool is not None
and current_tool is not stale_tool
and server_name in self._connections
):
return current_tool
def _reload_lock(state: Any) -> asyncio.Lock:
try:
return _RELOAD_LOCKS[state]
except KeyError:
lock = asyncio.Lock()
_RELOAD_LOCKS[state] = lock
return lock
def _attach_reconnect_handlers(
state: Any,
registry: ToolRegistry,
server_names: Mapping[str, Any] | set[str] | list[str] | tuple[str, ...],
) -> None:
async def reconnect(server_name: str, tool_name: str, stale_tool: Tool) -> Tool | None:
return await _refresh_terminated_server(
state,
registry,
server_name,
tool_name,
stale_tool,
)
for server_name in server_names:
for tool_name in list(registry.tool_names):
tool = registry.get(tool_name)
if not _tool_belongs_to_server(tool, tool_name, server_name):
continue
if isinstance(tool, _MCPWrapperBase):
tool.set_reconnect_handler(reconnect)
async def _refresh_terminated_server(
state: Any,
registry: ToolRegistry,
server_name: str,
tool_name: str,
stale_tool: Tool,
) -> Tool | None:
async with _reload_lock(state):
if getattr(state, "_mcp_closing", False):
return None
cfg = state._mcp_servers.get(server_name)
if cfg is None:
logger.warning( logger.warning(
"MCP server '{}' session terminated but is no longer configured", "MCP server '{}' session terminated; refreshing connection",
server_name, server_name,
) )
return None _unregister_server_tools(self._registry, server_name)
await self._close_server(server_name)
current_tool = registry.get(tool_name) self._set_runtime_status({server_name}, "connecting")
if ( connected = await connect_mcp_servers(
current_tool is not None {server_name: cfg},
and current_tool is not stale_tool self._registry,
and server_name in state._mcp_stacks )
): if self._closing:
return current_tool await _close_mcp_connections(connected)
return None
self._connections.update(connected)
self._record_connection_result({server_name}, connected)
self._attach_reconnect_handlers(connected)
if server_name not in connected:
logger.warning(
"MCP server '{}' reconnect failed after session termination",
server_name,
)
return None
return self._registry.get(tool_name)
logger.warning("MCP server '{}' session terminated; refreshing connection", server_name) async def _close_server(self, server_name: str) -> None:
_unregister_server_tools(registry, server_name) connection = self._connections.pop(server_name, None)
await _close_server(state, server_name) if connection is None:
return
await _close_mcp_connection(server_name, connection)
_set_runtime_status(state, {server_name}, "connecting") async def aclose(self) -> None:
connected = await connect_mcp_servers({server_name: cfg}, registry) """Close every connection while excluding reconnect and hot reload."""
if getattr(state, "_mcp_closing", False): self._closing = True
for connection in connected.values(): async with self._lock:
await connection.aclose() connections = dict(self._connections)
return None self._connections.clear()
state._mcp_stacks.update(connected) self._runtime_statuses.clear()
_record_connection_result(state, {server_name}, connected) for name in self._servers:
_attach_reconnect_handlers(state, registry, connected) _unregister_server_tools(self._registry, name)
if server_name not in connected: await _close_mcp_connections(connections)
logger.warning("MCP server '{}' reconnect failed after session termination", server_name)
return None
return registry.get(tool_name)
def _server_signature(cfg: Any) -> Any: def _server_signature(cfg: Any) -> Any:
@@ -1690,37 +1705,3 @@ def _unregister_server_tools(registry: ToolRegistry, server_name: str) -> int:
registry.unregister(tool_name) registry.unregister(tool_name)
removed += 1 removed += 1
return removed return removed
async def _close_server(state: Any, server_name: str) -> None:
stack = state._mcp_stacks.pop(server_name, None)
if stack is None:
return
try:
await stack.aclose()
except asyncio.CancelledError:
if task_is_cancelling():
raise
logger.debug("MCP server '{}' cleanup error (can be ignored)", server_name)
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()
statuses = _runtime_status_store(state)
if statuses is not None:
statuses.clear()
for name, connection in connections:
try:
await connection.aclose()
except asyncio.CancelledError:
if task_is_cancelling():
raise
logger.debug("MCP server '{}' cleanup error (can be ignored)", name)
except (RuntimeError, BaseExceptionGroup):
logger.debug("MCP server '{}' cleanup error (can be ignored)", name)
+1 -1
View File
@@ -78,7 +78,7 @@ class MyTool(Tool):
"runner", "sessions", "consolidator", "runner", "sessions", "consolidator",
"dream", "auto_compact", "context", "commands", "dream", "auto_compact", "context", "commands",
# Sensitive runtime state (credentials, message routing, task tracking) # Sensitive runtime state (credentials, message routing, task tracking)
"_mcp_servers", "_mcp_stacks", "_pending_queues", "_pending_queues",
"_session_locks", "_active_tasks", "_background_tasks", "_session_locks", "_active_tasks", "_background_tasks",
# Security boundaries (inspect + modify both blocked) # Security boundaries (inspect + modify both blocked)
"restrict_to_workspace", "channels_config", "restrict_to_workspace", "channels_config",
+23 -10
View File
@@ -48,6 +48,7 @@ _AGENT_LOOP_KEY = web.AppKey[Any]("agent_loop")
_MODEL_NAME_KEY = web.AppKey[str]("model_name") _MODEL_NAME_KEY = web.AppKey[str]("model_name")
_REQUEST_TIMEOUT_KEY = web.AppKey[float]("request_timeout") _REQUEST_TIMEOUT_KEY = web.AppKey[float]("request_timeout")
_SESSION_LOCKS_KEY = web.AppKey[dict[str, asyncio.Lock]]("session_locks") _SESSION_LOCKS_KEY = web.AppKey[dict[str, asyncio.Lock]]("session_locks")
_PREPARE_AGENT_KEY = web.AppKey[Callable[[], Awaitable[None]] | None]("prepare_agent")
_MISSING = object() _MISSING = object()
@@ -66,6 +67,17 @@ def _app_value(
return app.get(legacy_key, default) return app.get(legacy_key, default)
async def _prepare_agent(app: Any) -> None:
prepare: Callable[[], Awaitable[None]] | None = _app_value(
app,
_PREPARE_AGENT_KEY,
"prepare_agent",
None,
)
if prepare is not None:
await prepare()
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
# Response helpers # Response helpers
# --------------------------------------------------------------------------- # ---------------------------------------------------------------------------
@@ -346,8 +358,9 @@ async def handle_chat_completions(request: web.Request) -> web.Response | web.St
nonlocal stream_failed nonlocal stream_failed
try: try:
async with session_lock: async with session_lock:
response = await asyncio.wait_for( async with asyncio.timeout(timeout_s):
agent_loop.process_direct( await _prepare_agent(request.app)
response = await agent_loop.process_direct(
content=text, content=text,
media=media_paths if media_paths else None, media=media_paths if media_paths else None,
session_key=session_key, session_key=session_key,
@@ -355,9 +368,7 @@ async def handle_chat_completions(request: web.Request) -> web.Response | web.St
chat_id=API_CHAT_ID, chat_id=API_CHAT_ID,
on_stream=_on_stream, on_stream=_on_stream,
on_stream_end=_on_stream_end, on_stream_end=_on_stream_end,
), )
timeout=timeout_s,
)
if not emitted_content: if not emitted_content:
response_text = _response_text(response) response_text = _response_text(response)
if response_text.strip(): if response_text.strip():
@@ -390,16 +401,15 @@ async def handle_chat_completions(request: web.Request) -> web.Response | web.St
try: try:
async with session_lock: async with session_lock:
try: try:
response = await asyncio.wait_for( async with asyncio.timeout(timeout_s):
agent_loop.process_direct( await _prepare_agent(request.app)
response = await agent_loop.process_direct(
content=text, content=text,
media=media_paths if media_paths else None, media=media_paths if media_paths else None,
session_key=session_key, session_key=session_key,
channel="api", channel="api",
chat_id=API_CHAT_ID, chat_id=API_CHAT_ID,
), )
timeout=timeout_s,
)
response_text = _response_text(response) response_text = _response_text(response)
if not response_text or not response_text.strip(): if not response_text or not response_text.strip():
logger.warning("Empty response for session {}, using fallback", session_key) logger.warning("Empty response for session {}, using fallback", session_key)
@@ -452,6 +462,7 @@ def create_app(
model_name: str = "nanobot", model_name: str = "nanobot",
request_timeout: float = 120.0, request_timeout: float = 120.0,
api_key: str = "", api_key: str = "",
prepare_agent: Callable[[], Awaitable[None]] | None = None,
) -> web.Application: ) -> web.Application:
"""Create the aiohttp application. """Create the aiohttp application.
@@ -460,12 +471,14 @@ def create_app(
model_name: Model name reported in responses. model_name: Model name reported in responses.
request_timeout: Per-request timeout in seconds. request_timeout: Per-request timeout in seconds.
api_key: Optional API key for Bearer-token authentication on API routes. api_key: Optional API key for Bearer-token authentication on API routes.
prepare_agent: Optional application-owned readiness callback run before each turn.
""" """
app = web.Application(client_max_size=20 * 1024 * 1024) # 20MB for base64 images app = web.Application(client_max_size=20 * 1024 * 1024) # 20MB for base64 images
app[_AGENT_LOOP_KEY] = agent_loop app[_AGENT_LOOP_KEY] = agent_loop
app[_MODEL_NAME_KEY] = model_name app[_MODEL_NAME_KEY] = model_name
app[_REQUEST_TIMEOUT_KEY] = request_timeout app[_REQUEST_TIMEOUT_KEY] = request_timeout
app[_SESSION_LOCKS_KEY] = {} # per-user locks, keyed by session_key app[_SESSION_LOCKS_KEY] = {} # per-user locks, keyed by session_key
app[_PREPARE_AGENT_KEY] = prepare_agent
@web.middleware @web.middleware
async def auth_middleware( async def auth_middleware(
-1
View File
@@ -16,7 +16,6 @@ OUTBOUND_META_AGENT_UI = "_agent_ui"
# loop to update runtime state without going through a user session. # loop to update runtime state without going through a user session.
INBOUND_META_RUNTIME_CONTROL = "_runtime_control" INBOUND_META_RUNTIME_CONTROL = "_runtime_control"
RUNTIME_CONTROL_ACK = "_ack" RUNTIME_CONTROL_ACK = "_ack"
RUNTIME_CONTROL_MCP_RELOAD = "mcp_reload"
RUNTIME_CONTROL_IMAGE_GENERATION_RELOAD = "image_generation_reload" RUNTIME_CONTROL_IMAGE_GENERATION_RELOAD = "image_generation_reload"
RUNTIME_CONTROL_SESSION_DISCARD = "session_discard" RUNTIME_CONTROL_SESSION_DISCARD = "session_discard"
+4 -1
View File
@@ -5,7 +5,7 @@ from __future__ import annotations
import asyncio import asyncio
import hashlib import hashlib
import inspect import inspect
from collections.abc import Callable, Iterable, Mapping from collections.abc import Awaitable, Callable, Iterable, Mapping
from contextlib import suppress from contextlib import suppress
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Any, cast from typing import TYPE_CHECKING, Any, cast
@@ -101,6 +101,7 @@ class ChannelManager:
webui_runtime_surface: str = "browser", webui_runtime_surface: str = "browser",
webui_runtime_capabilities: dict[str, Any] | None = None, webui_runtime_capabilities: dict[str, Any] | None = None,
webui_mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None, webui_mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None,
webui_mcp_reload: Callable[[], Awaitable[dict[str, Any]]] | None = None,
webui_skill_state_action: Callable[[set[str]], None] | None = None, webui_skill_state_action: Callable[[set[str]], None] | None = None,
config_path: Path | None = None, config_path: Path | None = None,
): ):
@@ -121,6 +122,7 @@ class ChannelManager:
self._webui_runtime_surface = webui_runtime_surface self._webui_runtime_surface = webui_runtime_surface
self._webui_runtime_capabilities = dict(webui_runtime_capabilities or {}) self._webui_runtime_capabilities = dict(webui_runtime_capabilities or {})
self._webui_mcp_runtime_status = webui_mcp_runtime_status self._webui_mcp_runtime_status = webui_mcp_runtime_status
self._webui_mcp_reload = webui_mcp_reload
self._webui_skill_state_action = webui_skill_state_action self._webui_skill_state_action = webui_skill_state_action
self.channels: dict[str, BaseChannel] = {} self.channels: dict[str, BaseChannel] = {}
self._channel_owners: dict[str, str] = {} self._channel_owners: dict[str, str] = {}
@@ -190,6 +192,7 @@ class ChannelManager:
channel_feature_action=self.apply_channel_feature_action, channel_feature_action=self.apply_channel_feature_action,
channel_runtime_status=self.get_status, channel_runtime_status=self.get_status,
mcp_runtime_status=self._webui_mcp_runtime_status, mcp_runtime_status=self._webui_mcp_runtime_status,
mcp_reload=self._webui_mcp_reload,
skill_state_action=self._webui_skill_state_action, skill_state_action=self._webui_skill_state_action,
logger=logger, logger=logger,
) )
@@ -75,6 +75,7 @@ def _make_handler(
local_trigger_pending_ids: Any | None = None, local_trigger_pending_ids: Any | None = None,
channel_feature_action: Any | None = None, channel_feature_action: Any | None = None,
channel_runtime_status: Any | None = None, channel_runtime_status: Any | None = None,
mcp_reload: Any | None = None,
) -> GatewayServices: ) -> GatewayServices:
config = WebSocketConfig.model_validate(cfg) if isinstance(cfg, dict) else cfg config = WebSocketConfig.model_validate(cfg) if isinstance(cfg, dict) else cfg
workspace = workspace_path or Path.cwd() workspace = workspace_path or Path.cwd()
@@ -94,6 +95,7 @@ def _make_handler(
local_trigger_pending_ids=local_trigger_pending_ids, local_trigger_pending_ids=local_trigger_pending_ids,
channel_feature_action=channel_feature_action, channel_feature_action=channel_feature_action,
channel_runtime_status=channel_runtime_status, channel_runtime_status=channel_runtime_status,
mcp_reload=mcp_reload,
) )
@@ -111,6 +113,7 @@ def _ch(
local_trigger_pending_ids: Any | None = None, local_trigger_pending_ids: Any | None = None,
channel_feature_action: Any | None = None, channel_feature_action: Any | None = None,
channel_runtime_status: Any | None = None, channel_runtime_status: Any | None = None,
mcp_reload: Any | None = None,
**extra: Any, **extra: Any,
) -> WebSocketChannel: ) -> WebSocketChannel:
cfg: dict[str, Any] = { cfg: dict[str, Any] = {
@@ -134,6 +137,7 @@ def _ch(
local_trigger_pending_ids=local_trigger_pending_ids, local_trigger_pending_ids=local_trigger_pending_ids,
channel_feature_action=channel_feature_action, channel_feature_action=channel_feature_action,
channel_runtime_status=channel_runtime_status, channel_runtime_status=channel_runtime_status,
mcp_reload=mcp_reload,
) )
return InProcessHttpChannel(cfg, bus, gateway=gateway) return InProcessHttpChannel(cfg, bus, gateway=gateway)
@@ -2054,14 +2058,15 @@ async def test_mcp_presets_routes_require_token_and_return_payload(
_custom_action, _custom_action,
) )
async def _hot_reload(_bus): async def _hot_reload():
return {"ok": True, "message": "MCP config reloaded.", "requires_restart": False} return {"ok": True, "message": "MCP config reloaded.", "requires_restart": False}
monkeypatch.setattr( channel = _ch(
"nanobot.webui.settings_routes.request_mcp_reload", bus,
_hot_reload, session_manager=_seed_session(tmp_path),
port=29913,
mcp_reload=_hot_reload,
) )
channel = _ch(bus, session_manager=_seed_session(tmp_path), port=29913)
server_task = asyncio.create_task(channel.start()) server_task = asyncio.create_task(channel.start())
try: try:
deny = await _http_get("http://127.0.0.1:29913/api/settings/mcp-presets") deny = await _http_get("http://127.0.0.1:29913/api/settings/mcp-presets")
+38 -23
View File
@@ -13,6 +13,8 @@ from rich.console import Console
from nanobot import __logo__ from nanobot import __logo__
from nanobot.agent.hooks import create_file_edit_activity_hook from nanobot.agent.hooks import create_file_edit_activity_hook
from nanobot.agent.loop import AgentLoop from nanobot.agent.loop import AgentLoop
from nanobot.agent.tools.mcp import MCPProvider
from nanobot.agent.tools.registry import ToolRegistry
from nanobot.bus.outbound_events import ( from nanobot.bus.outbound_events import (
StreamDeltaEvent, StreamDeltaEvent,
StreamedResponseEvent, StreamedResponseEvent,
@@ -84,6 +86,8 @@ def agent(
# Create cron service with workspace-scoped store # Create cron service with workspace-scoped store
cron_store_path = runtime_config.workspace_path / "cron" / "jobs.json" cron_store_path = runtime_config.workspace_path / "cron" / "jobs.json"
cron = CronService(cron_store_path) cron = CronService(cron_store_path)
tools = ToolRegistry()
mcp_provider = MCPProvider.from_config(runtime_config, tools)
_set_nanobot_logs(logs) _set_nanobot_logs(logs)
@@ -95,6 +99,7 @@ def agent(
cron_service=cron, cron_service=cron,
image_generation_provider_configs=image_gen_provider_configs(runtime_config), image_generation_provider_configs=image_gen_provider_configs(runtime_config),
hook_factories=[create_file_edit_activity_hook], hook_factories=[create_file_edit_activity_hook],
tool_registry=tools,
) )
except ValueError as exc: except ValueError as exc:
_print_agent_start_error(exc) _print_agent_start_error(exc)
@@ -106,6 +111,12 @@ def agent(
render_markdown=False, render_markdown=False,
) )
async def _close_runtime() -> None:
try:
await agent_loop.aclose()
finally:
await mcp_provider.aclose()
# Shared reference for progress callbacks # Shared reference for progress callbacks
_thinking: ThinkingSpinner | None = None _thinking: ThinkingSpinner | None = None
@@ -149,30 +160,33 @@ def agent(
if message: if message:
# Single message mode — direct call, no bus needed # Single message mode — direct call, no bus needed
async def run_once() -> None: async def run_once() -> None:
renderer = StreamRenderer( try:
render_markdown=markdown, await mcp_provider.connect()
bot_name=runtime_config.agents.defaults.bot_name, renderer = StreamRenderer(
bot_icon=runtime_config.agents.defaults.bot_icon,
)
response = await agent_loop.process_direct(
message,
session_id,
on_progress=_make_progress(renderer),
on_stream=renderer.on_delta,
on_stream_end=renderer.on_end,
)
if not renderer.streamed:
await renderer.close()
print_kwargs: dict[str, Any] = {}
if renderer.header_printed:
print_kwargs["show_header"] = False
cli_terminal._print_agent_response(
response.content if response else "",
render_markdown=markdown, render_markdown=markdown,
metadata=response.metadata if response else None, bot_name=runtime_config.agents.defaults.bot_name,
**print_kwargs, bot_icon=runtime_config.agents.defaults.bot_icon,
) )
await agent_loop.close_mcp() response = await agent_loop.process_direct(
message,
session_id,
on_progress=_make_progress(renderer),
on_stream=renderer.on_delta,
on_stream_end=renderer.on_end,
)
if not renderer.streamed:
await renderer.close()
print_kwargs: dict[str, Any] = {}
if renderer.header_printed:
print_kwargs["show_header"] = False
cli_terminal._print_agent_response(
response.content if response else "",
render_markdown=markdown,
metadata=response.metadata if response else None,
**print_kwargs,
)
finally:
await _close_runtime()
asyncio.run(run_once()) asyncio.run(run_once())
else: else:
@@ -209,6 +223,7 @@ def agent(
signal.signal(signal.SIGPIPE, signal.SIG_IGN) signal.signal(signal.SIGPIPE, signal.SIG_IGN)
async def run_interactive() -> None: async def run_interactive() -> None:
await mcp_provider.connect()
bus_task = asyncio.create_task(agent_loop.run()) bus_task = asyncio.create_task(agent_loop.run())
turn_done = asyncio.Event() turn_done = asyncio.Event()
turn_done.set() turn_done.set()
@@ -347,6 +362,6 @@ def agent(
agent_loop.stop() agent_loop.stop()
outbound_task.cancel() outbound_task.cancel()
await asyncio.gather(bus_task, outbound_task, return_exceptions=True) await asyncio.gather(bus_task, outbound_task, return_exceptions=True)
await agent_loop.close_mcp() await _close_runtime()
asyncio.run(run_interactive()) asyncio.run(run_interactive())
+11 -2
View File
@@ -49,6 +49,8 @@ from nanobot import __logo__, __version__ # noqa: E402
from nanobot import optional_features as feature_support # noqa: E402 from nanobot import optional_features as feature_support # noqa: E402
from nanobot.agent.hooks import create_file_edit_activity_hook # noqa: E402 from nanobot.agent.hooks import create_file_edit_activity_hook # noqa: E402
from nanobot.agent.loop import AgentLoop # noqa: E402 from nanobot.agent.loop import AgentLoop # noqa: E402
from nanobot.agent.tools.mcp import MCPProvider # noqa: E402
from nanobot.agent.tools.registry import ToolRegistry # noqa: E402
from nanobot.cli import terminal as cli_terminal # noqa: E402 from nanobot.cli import terminal as cli_terminal # noqa: E402
from nanobot.cli.agent import agent # noqa: E402 from nanobot.cli.agent import agent # noqa: E402
from nanobot.cli.gateway import create_gateway_app # noqa: E402 from nanobot.cli.gateway import create_gateway_app # noqa: E402
@@ -351,12 +353,15 @@ def serve(
sync_workspace_templates(runtime_config.workspace_path) sync_workspace_templates(runtime_config.workspace_path)
bus = MessageBus() bus = MessageBus()
session_manager = SessionManager(runtime_config.workspace_path) session_manager = SessionManager(runtime_config.workspace_path)
tools = ToolRegistry()
mcp_provider = MCPProvider.from_config(runtime_config, tools)
try: try:
agent_loop = AgentLoop.from_config( agent_loop = AgentLoop.from_config(
runtime_config, bus, runtime_config, bus,
session_manager=session_manager, session_manager=session_manager,
image_generation_provider_configs=image_gen_provider_configs(runtime_config), image_generation_provider_configs=image_gen_provider_configs(runtime_config),
hook_factories=[create_file_edit_activity_hook], hook_factories=[create_file_edit_activity_hook],
tool_registry=tools,
) )
except ValueError as exc: except ValueError as exc:
console.print(f"[red]Error: {exc}[/red]") console.print(f"[red]Error: {exc}[/red]")
@@ -378,13 +383,17 @@ def serve(
api_app = create_app( api_app = create_app(
agent_loop, model_name=model_name, request_timeout=timeout, agent_loop, model_name=model_name, request_timeout=timeout,
api_key=api_key, api_key=api_key,
prepare_agent=mcp_provider.connect,
) )
async def on_startup(_app: Any) -> None: async def on_startup(_app: Any) -> None:
await agent_loop._connect_mcp() await mcp_provider.connect()
async def on_cleanup(_app: Any) -> None: async def on_cleanup(_app: Any) -> None:
await agent_loop.close_mcp() try:
await agent_loop.aclose()
finally:
await mcp_provider.aclose()
api_app.on_startup.append(on_startup) api_app.on_startup.append(on_startup)
api_app.on_cleanup.append(on_cleanup) api_app.on_cleanup.append(on_cleanup)
+39 -17
View File
@@ -14,6 +14,8 @@ from rich.console import Console
from nanobot import __logo__, __version__ from nanobot import __logo__, __version__
from nanobot.agent.hooks import create_file_edit_activity_hook from nanobot.agent.hooks import create_file_edit_activity_hook
from nanobot.agent.loop import AgentLoop from nanobot.agent.loop import AgentLoop
from nanobot.agent.tools.mcp import MCPProvider
from nanobot.agent.tools.registry import ToolRegistry
from nanobot.cli import terminal as cli_terminal from nanobot.cli import terminal as cli_terminal
from nanobot.cli.runtime_config import _migrate_cron_store from nanobot.cli.runtime_config import _migrate_cron_store
from nanobot.cli.webui_support import ( from nanobot.cli.webui_support import (
@@ -233,6 +235,7 @@ def _print_gateway_health_endpoint(host: str, port: int) -> None:
async def _close_gateway_runtime( async def _close_gateway_runtime(
agent: AgentLoop, agent: AgentLoop,
mcp_provider: MCPProvider,
channels: Any, channels: Any,
tasks: list[asyncio.Task[Any]], tasks: list[asyncio.Task[Any]],
runtime_tasks: asyncio.Future[list[Any]] | None, runtime_tasks: asyncio.Future[list[Any]] | None,
@@ -240,18 +243,13 @@ async def _close_gateway_runtime(
task_wait_timeout: float = 15.0, task_wait_timeout: float = 15.0,
close_timeout: float = 15.0, close_timeout: float = 15.0,
) -> None: ) -> None:
"""Cancel runtime tasks, then deterministically close agent resources. """Cancel runtime tasks, then deterministically close application resources.
Order matters: runtime tasks (including the agent loop and any in-flight Order matters: runtime tasks (including the agent loop and any in-flight
turn) are cancelled and awaited -- bounded -- before exec sessions, turn) are cancelled and awaited -- bounded -- before the loop-owned resources
subagents, and MCP servers are torn down, so no active turn is using a and the application-owned MCP provider are torn down. The final close is
shared resource when it closes. The final close is bounded and idempotent: bounded and idempotent, so it also covers a cancelled or incomplete loop
the agent loop's own finally also calls ``close_mcp()``, so this runs again cleanup without leaving subprocess transports alive past ``loop.close()``.
as a no-op when that path already completed, and as the guaranteed final
close when it was skipped or cut short (which previously left asyncio
subprocess transports alive past ``loop.close()``, producing
"RuntimeError: Event loop is closed" noise and potentially orphaned
processes at interpreter exit).
""" """
# Some SDKs swallow task cancellation while attempting to reconnect. # Some SDKs swallow task cancellation while attempting to reconnect.
# Close channel transports before waiting for their runners to exit. # Close channel transports before waiting for their runners to exit.
@@ -272,10 +270,14 @@ async def _close_gateway_runtime(
task.cancel() task.cancel()
if runtime_tasks is not None and not runtime_tasks.done(): if runtime_tasks is not None and not runtime_tasks.done():
runtime_tasks.cancel() runtime_tasks.cancel()
try: for label, close in (
await asyncio.wait_for(agent.close_mcp(), timeout=close_timeout) ("agent", agent.aclose),
except BaseException as exc: # noqa: BLE001 - shutdown must proceed ("MCP provider", mcp_provider.aclose),
logger.warning("Gateway shutdown: agent resource cleanup incomplete: {}", exc) ):
try:
await asyncio.wait_for(close(), timeout=close_timeout)
except BaseException as exc: # noqa: BLE001 - shutdown must proceed
logger.warning("Gateway shutdown: {} cleanup incomplete: {}", label, exc)
# Retrieving an already-finished gather prevents noisy unhandled exceptions, # Retrieving an already-finished gather prevents noisy unhandled exceptions,
# but never wait for it here: its children were bounded individually above. # but never wait for it here: its children were bounded individually above.
if runtime_tasks is not None and runtime_tasks.done(): if runtime_tasks is not None and runtime_tasks.done():
@@ -414,6 +416,9 @@ def _run_gateway(
route_policy=WebuiTurnRoutePolicy(session_manager), route_policy=WebuiTurnRoutePolicy(session_manager),
) )
tools = ToolRegistry()
mcp_provider = MCPProvider.from_config(config, tools)
# Create agent with cron service # Create agent with cron service
agent = AgentLoop.from_config( agent = AgentLoop.from_config(
config, bus, config, bus,
@@ -431,6 +436,7 @@ def _run_gateway(
hooks=[TokenUsageHook(timezone_name=config.agents.defaults.timezone)], hooks=[TokenUsageHook(timezone_name=config.agents.defaults.timezone)],
local_trigger_store=trigger_store, local_trigger_store=trigger_store,
hook_factories=[create_file_edit_activity_hook], hook_factories=[create_file_edit_activity_hook],
tool_registry=tools,
) )
def _schedule_webui_background(awaitable: Awaitable[None]) -> None: def _schedule_webui_background(awaitable: Awaitable[None]) -> None:
agent.schedule_background(cast(Coroutine[Any, Any, None], awaitable)) agent.schedule_background(cast(Coroutine[Any, Any, None], awaitable))
@@ -512,6 +518,7 @@ def _run_gateway(
prompt, last_cursor = result prompt, last_cursor = result
key = dream_session_key() key = dream_session_key()
dream_runtime = agent.dream_runtime() dream_runtime = agent.dream_runtime()
await mcp_provider.connect()
resp = await agent.process_direct( resp = await agent.process_direct(
prompt, prompt,
session_key=key, session_key=key,
@@ -589,6 +596,7 @@ def _run_gateway(
if isinstance(message_tool, MessageTool): if isinstance(message_tool, MessageTool):
suppress_token = message_tool.set_suppress_delivery(True) suppress_token = message_tool.set_suppress_delivery(True)
try: try:
await mcp_provider.connect()
resp = await agent.process_direct( resp = await agent.process_direct(
prompt, prompt,
session_key="heartbeat", session_key="heartbeat",
@@ -668,7 +676,8 @@ def _run_gateway(
webui_static_dist=webui_static_dist, webui_static_dist=webui_static_dist,
webui_runtime_surface=webui_runtime_surface, webui_runtime_surface=webui_runtime_surface,
webui_runtime_capabilities=webui_runtime_capabilities, webui_runtime_capabilities=webui_runtime_capabilities,
webui_mcp_runtime_status=agent.mcp_runtime_status, webui_mcp_runtime_status=mcp_provider.runtime_status,
webui_mcp_reload=mcp_provider.reload,
webui_skill_state_action=_webui_skill_state_action, webui_skill_state_action=_webui_skill_state_action,
config_path=Path(config_path), config_path=Path(config_path),
) )
@@ -844,6 +853,13 @@ def _run_gateway(
await cron.start() await cron.start()
# Re-read once on first admission to close the watcher subscription window. # Re-read once on first admission to close the watcher subscription window.
agent.runtime_resolver.invalidate() agent.runtime_resolver.invalidate()
async def _run_agent() -> None:
try:
await mcp_provider.connect()
await agent.run()
finally:
await mcp_provider.aclose()
tasks = [ tasks = [
asyncio.create_task( asyncio.create_task(
watch_config_file( watch_config_file(
@@ -852,7 +868,7 @@ def _run_gateway(
), ),
name="nanobot-config-watcher", name="nanobot-config-watcher",
), ),
asyncio.create_task(agent.run(), name="nanobot-agent-loop"), asyncio.create_task(_run_agent(), name="nanobot-agent-loop"),
asyncio.create_task(channels.start_all(), name="nanobot-channels"), asyncio.create_task(channels.start_all(), name="nanobot-channels"),
asyncio.create_task( asyncio.create_task(
run_local_trigger_queue( run_local_trigger_queue(
@@ -910,7 +926,13 @@ def _run_gateway(
agent.stop() agent.stop()
# Cancel runtime tasks first, then deterministically close # Cancel runtime tasks first, then deterministically close
# exec/MCP resources while the event loop is still alive. # exec/MCP resources while the event loop is still alive.
await _close_gateway_runtime(agent, channels, tasks, runtime_tasks) await _close_gateway_runtime(
agent,
mcp_provider,
channels,
tasks,
runtime_tasks,
)
# Flush all cached sessions to durable storage before exit. # Flush all cached sessions to durable storage before exit.
# This prevents data loss on filesystems with write-back # This prevents data loss on filesystems with write-back
# caching (rclone VFS, NFS, FUSE mounts, etc.). # caching (rclone VFS, NFS, FUSE mounts, etc.).
+24 -4
View File
@@ -10,6 +10,8 @@ from typing import Any
from nanobot.agent.hook import AgentHook, SDKCaptureHook from nanobot.agent.hook import AgentHook, SDKCaptureHook
from nanobot.agent.hooks import create_file_edit_activity_hook from nanobot.agent.hooks import create_file_edit_activity_hook
from nanobot.agent.loop import AgentLoop from nanobot.agent.loop import AgentLoop
from nanobot.agent.tools.mcp import MCPProvider
from nanobot.agent.tools.registry import ToolRegistry
from nanobot.config.schema import Config from nanobot.config.schema import Config
from nanobot.providers.image_generation import image_gen_provider_configs from nanobot.providers.image_generation import image_gen_provider_configs
from nanobot.sdk.clients import MemoryClient, RuntimeClient, SessionClient from nanobot.sdk.clients import MemoryClient, RuntimeClient, SessionClient
@@ -71,9 +73,16 @@ class Nanobot:
print(result.content) print(result.content)
""" """
def __init__(self, loop: AgentLoop, *, config: Config | None = None) -> None: def __init__(
self,
loop: AgentLoop,
*,
config: Config | None = None,
mcp_provider: MCPProvider | None = None,
) -> None:
self._loop = loop self._loop = loop
self._config = config self._config = config
self._mcp_provider = mcp_provider
self.sessions = SessionClient(loop) self.sessions = SessionClient(loop)
self.memory = MemoryClient(loop) self.memory = MemoryClient(loop)
self.runtime = RuntimeClient(loop) self.runtime = RuntimeClient(loop)
@@ -120,12 +129,15 @@ class Nanobot:
elif model_preset is not None: elif model_preset is not None:
config.agents.defaults.model_preset = model_preset config.agents.defaults.model_preset = model_preset
tools = ToolRegistry()
mcp_provider = MCPProvider.from_config(config, tools)
loop = AgentLoop.from_config( loop = AgentLoop.from_config(
config, config,
image_generation_provider_configs=image_gen_provider_configs(config), image_generation_provider_configs=image_gen_provider_configs(config),
hook_factories=[create_file_edit_activity_hook], hook_factories=[create_file_edit_activity_hook],
tool_registry=tools,
) )
return cls(loop, config=config) return cls(loop, config=config, mcp_provider=mcp_provider)
async def run( async def run(
self, self,
@@ -178,6 +190,8 @@ class Nanobot:
) )
if runtime is not None: if runtime is not None:
kwargs["runtime"] = runtime kwargs["runtime"] = runtime
if self._mcp_provider is not None:
await self._mcp_provider.connect()
response = await self._loop.process_direct( response = await self._loop.process_direct(
message, message,
**kwargs, **kwargs,
@@ -259,6 +273,8 @@ class Nanobot:
if override_runtime is not None: if override_runtime is not None:
kwargs["runtime"] = override_runtime kwargs["runtime"] = override_runtime
try: try:
if self._mcp_provider is not None:
await self._mcp_provider.connect()
response = await self._loop.process_direct( response = await self._loop.process_direct(
message, message,
**kwargs, **kwargs,
@@ -327,8 +343,12 @@ class Nanobot:
await run.aclose() await run.aclose()
async def aclose(self) -> None: async def aclose(self) -> None:
"""Release resources held by this instance (MCP connections, etc.).""" """Release resources held by this instance."""
await self._loop.close_mcp() try:
await self._loop.aclose()
finally:
if self._mcp_provider is not None:
await self._mcp_provider.aclose()
async def __aenter__(self) -> Nanobot: async def __aenter__(self) -> Nanobot:
return self return self
+3 -1
View File
@@ -2,7 +2,7 @@
from __future__ import annotations from __future__ import annotations
from collections.abc import Mapping from collections.abc import Awaitable, Mapping
from dataclasses import dataclass from dataclasses import dataclass
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Any, Callable from typing import TYPE_CHECKING, Any, Callable
@@ -66,6 +66,7 @@ def build_gateway_services(
channel_feature_action: Callable[..., Any] | None = None, channel_feature_action: Callable[..., Any] | None = None,
channel_runtime_status: Callable[[], dict[str, Any]] | None = None, channel_runtime_status: Callable[[], dict[str, Any]] | None = None,
mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None, mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None,
mcp_reload: Callable[[], Awaitable[dict[str, Any]]] | None = None,
skill_state_action: Callable[[set[str]], None] | None = None, skill_state_action: Callable[[set[str]], None] | None = None,
logger: Any = default_logger, logger: Any = default_logger,
) -> GatewayServices: ) -> GatewayServices:
@@ -119,6 +120,7 @@ def build_gateway_services(
channel_feature_action=channel_feature_action, channel_feature_action=channel_feature_action,
channel_runtime_status=channel_runtime_status, channel_runtime_status=channel_runtime_status,
mcp_runtime_status=mcp_runtime_status, mcp_runtime_status=mcp_runtime_status,
mcp_reload=mcp_reload,
skill_state_action=skill_state_action, skill_state_action=skill_state_action,
log=logger, log=logger,
) )
+33 -4
View File
@@ -5,14 +5,13 @@ from __future__ import annotations
import asyncio import asyncio
import html import html
import json import json
from collections.abc import Callable, Mapping from collections.abc import Awaitable, Callable, Mapping
from typing import Any, cast from typing import Any, cast
from websockets.http11 import Request as WsRequest from websockets.http11 import Request as WsRequest
from websockets.http11 import Response from websockets.http11 import Response
from nanobot.agent.tools.image_generation import request_image_generation_reload from nanobot.agent.tools.image_generation import request_image_generation_reload
from nanobot.agent.tools.mcp import request_mcp_reload
from nanobot.agent.tools.mcp_oauth import MCP_OAUTH_CALLBACK_PATH from nanobot.agent.tools.mcp_oauth import MCP_OAUTH_CALLBACK_PATH
from nanobot.api.runtime import ApiRuntime, api_runtime_paths from nanobot.api.runtime import ApiRuntime, api_runtime_paths
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
@@ -71,6 +70,7 @@ _WEBUI_MUTATION_PAYLOAD_ATTR = "_nanobot_webui_mutation_payload"
_WEBUI_MUTATION_REQUEST_ATTR = "_nanobot_webui_mutation_request" _WEBUI_MUTATION_REQUEST_ATTR = "_nanobot_webui_mutation_request"
_CHANNEL_CONNECT_ACTIONS = frozenset({"start", "poll", "cancel"}) _CHANNEL_CONNECT_ACTIONS = frozenset({"start", "poll", "cancel"})
_MCP_OAUTH_CALLBACK_URL_MAX_BYTES = 8 * 1024 _MCP_OAUTH_CALLBACK_URL_MAX_BYTES = 8 * 1024
_MCP_RELOAD_TIMEOUT_SECONDS = 15.0
_query_first = contracts.query_first _query_first = contracts.query_first
@@ -227,6 +227,7 @@ class WebUISettingsRouter:
channel_feature_action: Callable[..., Any] | None = None, channel_feature_action: Callable[..., Any] | None = None,
channel_runtime_status: Callable[[], dict[str, Any]] | None = None, channel_runtime_status: Callable[[], dict[str, Any]] | None = None,
mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None, mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None,
mcp_reload: Callable[[], Awaitable[dict[str, Any]]] | None = None,
mcp_oauth_redirect_uri: Callable[[WsRequest], str] | None = None, mcp_oauth_redirect_uri: Callable[[WsRequest], str] | None = None,
) -> None: ) -> None:
self.settings = settings self.settings = settings
@@ -241,6 +242,7 @@ class WebUISettingsRouter:
self._channel_feature_action = channel_feature_action self._channel_feature_action = channel_feature_action
self._channel_runtime_status = channel_runtime_status self._channel_runtime_status = channel_runtime_status
self._mcp_runtime_status = mcp_runtime_status self._mcp_runtime_status = mcp_runtime_status
self._mcp_reload = mcp_reload
self._mcp_oauth_redirect_uri = mcp_oauth_redirect_uri self._mcp_oauth_redirect_uri = mcp_oauth_redirect_uri
self._mcp_oauth = McpOAuthManager() self._mcp_oauth = McpOAuthManager()
self._restart_sections: set[str] = set() self._restart_sections: set[str] = set()
@@ -472,7 +474,7 @@ class WebUISettingsRouter:
approve_code=approve_code, approve_code=approve_code,
deny_code=deny_code, deny_code=deny_code,
mcp_presets_action=mcp_presets_settings_action, mcp_presets_action=mcp_presets_settings_action,
reload_mcp=lambda: request_mcp_reload(self.bus), reload_mcp=self._reload_mcp_runtime,
mcp_runtime_status=self._mcp_runtime_status, mcp_runtime_status=self._mcp_runtime_status,
check_for_update=check_for_update, check_for_update=check_for_update,
channel_feature_action=self._channel_feature_action, channel_feature_action=self._channel_feature_action,
@@ -499,6 +501,33 @@ class WebUISettingsRouter:
self._restart_sections.discard("image") self._restart_sections.discard("image")
return updated return updated
async def _reload_mcp_runtime(self) -> dict[str, Any]:
if self._mcp_reload is None:
return {
"ok": False,
"message": "MCP runtime reload is unavailable. Restart nanobot to apply changes.",
"requires_restart": True,
}
try:
return await asyncio.wait_for(
self._mcp_reload(),
timeout=_MCP_RELOAD_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
return {
"ok": False,
"message": "MCP hot reload timed out. Restart nanobot to pick up changes.",
"requires_restart": True,
}
except Exception as exc:
self.logger.exception("MCP hot reload failed")
return {
"ok": False,
"message": "MCP hot reload failed. Restart nanobot to pick up changes.",
"requires_restart": True,
"error": str(exc),
}
def _parse_mcp_settings_query(self, request: WsRequest) -> QueryParams: def _parse_mcp_settings_query(self, request: WsRequest) -> QueryParams:
return self._query(request) return self._query(request)
@@ -622,7 +651,7 @@ class WebUISettingsRouter:
name, name,
cfg, cfg,
redirect_uri, redirect_uri,
reload_mcp=lambda: request_mcp_reload(self.bus), reload_mcp=self._reload_mcp_runtime,
reset_credentials=reset, reset_credentials=reset,
) )
except Exception as exc: except Exception as exc:
+3 -1
View File
@@ -14,7 +14,7 @@ import json
import mimetypes import mimetypes
import re import re
import time import time
from collections.abc import Callable, Mapping from collections.abc import Awaitable, Callable, Mapping
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING, Any, cast from typing import TYPE_CHECKING, Any, cast
from urllib.parse import quote, unquote, urlsplit, urlunsplit from urllib.parse import quote, unquote, urlsplit, urlunsplit
@@ -308,6 +308,7 @@ class GatewayHTTPHandler:
channel_feature_action: Callable[..., Any] | None = None, channel_feature_action: Callable[..., Any] | None = None,
channel_runtime_status: Callable[[], dict[str, Any]] | None = None, channel_runtime_status: Callable[[], dict[str, Any]] | None = None,
mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None, mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None,
mcp_reload: Callable[[], Awaitable[dict[str, Any]]] | None = None,
skill_state_action: Callable[[set[str]], None] | None = None, skill_state_action: Callable[[set[str]], None] | None = None,
log: Any = logger, log: Any = logger,
) -> None: ) -> None:
@@ -351,6 +352,7 @@ class GatewayHTTPHandler:
channel_feature_action=channel_feature_action, channel_feature_action=channel_feature_action,
channel_runtime_status=channel_runtime_status, channel_runtime_status=channel_runtime_status,
mcp_runtime_status=mcp_runtime_status, mcp_runtime_status=mcp_runtime_status,
mcp_reload=mcp_reload,
mcp_oauth_redirect_uri=self._mcp_oauth_redirect_uri, mcp_oauth_redirect_uri=self._mcp_oauth_redirect_uri,
) )
-3
View File
@@ -46,7 +46,6 @@ def make_loop(
context_window_tokens: int = 128_000, context_window_tokens: int = 128_000,
session_ttl_minutes: int = 0, session_ttl_minutes: int = 0,
unified_session: bool = False, unified_session: bool = False,
mcp_servers: dict | None = None,
tools_config=None, tools_config=None,
model_presets: dict | None = None, model_presets: dict | None = None,
hooks: list | None = None, hooks: list | None = None,
@@ -72,8 +71,6 @@ def make_loop(
session_ttl_minutes=session_ttl_minutes, session_ttl_minutes=session_ttl_minutes,
unified_session=unified_session, unified_session=unified_session,
) )
if mcp_servers is not None:
kwargs["mcp_servers"] = mcp_servers
if tools_config is not None: if tools_config is not None:
kwargs["tools_config"] = tools_config kwargs["tools_config"] = tools_config
if model_presets is not None: if model_presets is not None:
+46 -41
View File
@@ -8,6 +8,7 @@ from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
from nanobot.agent.loop import AgentLoop from nanobot.agent.loop import AgentLoop
from nanobot.agent.tools.registry import ToolRegistry
from nanobot.bus.events import InboundMessage from nanobot.bus.events import InboundMessage
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.command import CommandContext from nanobot.command import CommandContext
@@ -193,7 +194,11 @@ class TestIdleScanThrottling:
}) })
provider = MagicMock() provider = MagicMock()
provider.get_default_model.return_value = "test-model" provider.get_default_model.return_value = "test-model"
loop = AgentLoop.from_config(config, provider=provider) loop = AgentLoop.from_config(
config,
tool_registry=ToolRegistry(),
provider=provider,
)
loop.auto_compact.check_expired = MagicMock() loop.auto_compact.check_expired = MagicMock()
loop._check_expired_sessions_if_due() loop._check_expired_sessions_if_due()
@@ -310,7 +315,7 @@ class TestAutoCompact:
assert loop.auto_compact._is_expired(ts) is True assert loop.auto_compact._is_expired(ts) is True
ts2 = datetime.now() - timedelta(minutes=14, seconds=59) ts2 = datetime.now() - timedelta(minutes=14, seconds=59)
assert loop.auto_compact._is_expired(ts2) is False assert loop.auto_compact._is_expired(ts2) is False
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_is_expired_string_timestamp(self, tmp_path): async def test_is_expired_string_timestamp(self, tmp_path):
@@ -320,7 +325,7 @@ class TestAutoCompact:
assert loop.auto_compact._is_expired(ts) is True assert loop.auto_compact._is_expired(ts) is True
assert loop.auto_compact._is_expired(None) is False assert loop.auto_compact._is_expired(None) is False
assert loop.auto_compact._is_expired("") is False assert loop.auto_compact._is_expired("") is False
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_check_expired_only_archives_expired_sessions(self, tmp_path): async def test_check_expired_only_archives_expired_sessions(self, tmp_path):
@@ -343,7 +348,7 @@ class TestAutoCompact:
active_after = loop.sessions.get_or_create("cli:active") active_after = loop.sessions.get_or_create("cli:active")
assert len(active_after.messages) == 1 assert len(active_after.messages) == 1
assert active_after.messages[0]["content"] == "recent" assert active_after.messages[0]["content"] == "recent"
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_auto_compact_archives_full_tail_without_deleting_history(self, tmp_path): async def test_auto_compact_archives_full_tail_without_deleting_history(self, tmp_path):
@@ -367,7 +372,7 @@ class TestAutoCompact:
assert len(visible) == loop.auto_compact._RECENT_SUFFIX_MESSAGES assert len(visible) == loop.auto_compact._RECENT_SUFFIX_MESSAGES
assert visible[0]["content"] == "msg user 2" assert visible[0]["content"] == "msg user 2"
assert visible[-1]["content"] == "msg assistant 5" assert visible[-1]["content"] == "msg assistant 5"
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_auto_compact_extends_recent_suffix_to_user_turn(self, tmp_path): async def test_auto_compact_extends_recent_suffix_to_user_turn(self, tmp_path):
@@ -398,7 +403,7 @@ class TestAutoCompact:
for m in visible for m in visible
for tc in (m.get("tool_calls") or []) for tc in (m.get("tool_calls") or [])
) )
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_auto_compact_stores_summary(self, tmp_path): async def test_auto_compact_stores_summary(self, tmp_path):
@@ -422,7 +427,7 @@ class TestAutoCompact:
assert len(session_after.get_history(max_messages=12)) == ( assert len(session_after.get_history(max_messages=12)) == (
loop.auto_compact._RECENT_SUFFIX_MESSAGES loop.auto_compact._RECENT_SUFFIX_MESSAGES
) )
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_auto_compact_empty_session(self, tmp_path): async def test_auto_compact_empty_session(self, tmp_path):
@@ -436,7 +441,7 @@ class TestAutoCompact:
session_after = loop.sessions.get_or_create("cli:test") session_after = loop.sessions.get_or_create("cli:test")
assert len(session_after.messages) == 0 assert len(session_after.messages) == 0
assert "cli:test" not in loop.auto_compact._summaries assert "cli:test" not in loop.auto_compact._summaries
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_auto_compact_respects_last_consolidated(self, tmp_path): async def test_auto_compact_respects_last_consolidated(self, tmp_path):
@@ -455,7 +460,7 @@ class TestAutoCompact:
await loop.auto_compact._archive("cli:test", runtime=loop.llm_runtime()) await loop.auto_compact._archive("cli:test", runtime=loop.llm_runtime())
assert len(archived_messages) == 10 assert len(archived_messages) == 10
await loop.close_mcp() await loop.aclose()
class TestAutoCompactIdleDetection: class TestAutoCompactIdleDetection:
@@ -474,7 +479,7 @@ class TestAutoCompactIdleDetection:
session_after = loop.sessions.get_or_create("cli:test") session_after = loop.sessions.get_or_create("cli:test")
assert any(m["content"] == "old message" for m in session_after.messages) assert any(m["content"] == "old message" for m in session_after.messages)
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_auto_compact_triggers_on_idle(self, tmp_path): async def test_auto_compact_triggers_on_idle(self, tmp_path):
@@ -503,7 +508,7 @@ class TestAutoCompactIdleDetection:
for m in session_after.get_history(max_messages=len(session_after.messages)) for m in session_after.get_history(max_messages=len(session_after.messages))
) )
assert any(m["content"] == "new msg" for m in session_after.messages) assert any(m["content"] == "new msg" for m in session_after.messages)
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_no_auto_compact_when_active(self, tmp_path): async def test_no_auto_compact_when_active(self, tmp_path):
@@ -517,7 +522,7 @@ class TestAutoCompactIdleDetection:
session_after = loop.sessions.get_or_create("cli:test") session_after = loop.sessions.get_or_create("cli:test")
assert any(m["content"] == "recent message" for m in session_after.messages) assert any(m["content"] == "recent message" for m in session_after.messages)
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_auto_compact_does_not_affect_priority_commands(self, tmp_path): async def test_auto_compact_does_not_affect_priority_commands(self, tmp_path):
@@ -540,7 +545,7 @@ class TestAutoCompactIdleDetection:
# Session should be untouched since priority commands skip _process_message # Session should be untouched since priority commands skip _process_message
session_after = loop.sessions.get_or_create("cli:test") session_after = loop.sessions.get_or_create("cli:test")
assert any(m["content"] == "old message" for m in session_after.messages) assert any(m["content"] == "old message" for m in session_after.messages)
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_auto_compact_with_slash_new(self, tmp_path): async def test_auto_compact_with_slash_new(self, tmp_path):
@@ -562,7 +567,7 @@ class TestAutoCompactIdleDetection:
session_after = loop.sessions.get_or_create("cli:test") session_after = loop.sessions.get_or_create("cli:test")
assert len(session_after.messages) == 0 assert len(session_after.messages) == 0
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_shortcut_command_persisted_with_command_flag(self, tmp_path): async def test_shortcut_command_persisted_with_command_flag(self, tmp_path):
@@ -581,7 +586,7 @@ class TestAutoCompactIdleDetection:
assert session_after.messages[1]["role"] == "assistant" assert session_after.messages[1]["role"] == "assistant"
assert session_after.messages[1].get("_command") is True assert session_after.messages[1].get("_command") is True
assert AgentLoop._PENDING_USER_TURN_KEY not in session_after.metadata assert AgentLoop._PENDING_USER_TURN_KEY not in session_after.metadata
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_shortcut_command_excluded_from_get_history(self, tmp_path): async def test_shortcut_command_excluded_from_get_history(self, tmp_path):
@@ -597,7 +602,7 @@ class TestAutoCompactIdleDetection:
assert len(history) == 2 assert len(history) == 2
assert all(m["content"] != "/help" for m in history) assert all(m["content"] != "/help" for m in history)
assert all(m["content"] != "help text" for m in history) assert all(m["content"] != "help text" for m in history)
await loop.close_mcp() await loop.aclose()
class TestAutoCompactSystemMessages: class TestAutoCompactSystemMessages:
@@ -628,7 +633,7 @@ class TestAutoCompactSystemMessages:
m["content"] == "old user 0" m["content"] == "old user 0"
for m in session_after.get_history(max_messages=len(session_after.messages)) for m in session_after.get_history(max_messages=len(session_after.messages))
) )
await loop.close_mcp() await loop.aclose()
class TestAutoCompactEdgeCases: class TestAutoCompactEdgeCases:
@@ -656,7 +661,7 @@ class TestAutoCompactEdgeCases:
# "(nothing)" summary should not be stored # "(nothing)" summary should not be stored
assert "cli:test" not in loop.auto_compact._summaries assert "cli:test" not in loop.auto_compact._summaries
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_auto_compact_archive_failure_preserves_raw_history(self, tmp_path): async def test_auto_compact_archive_failure_preserves_raw_history(self, tmp_path):
@@ -677,7 +682,7 @@ class TestAutoCompactEdgeCases:
loop.auto_compact._RECENT_SUFFIX_MESSAGES loop.auto_compact._RECENT_SUFFIX_MESSAGES
) )
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_auto_compact_preserves_runtime_checkpoint_before_check(self, tmp_path): async def test_auto_compact_preserves_runtime_checkpoint_before_check(self, tmp_path):
@@ -709,7 +714,7 @@ class TestAutoCompactEdgeCases:
assert any(m["content"] == "previous message" for m in session_after.messages) assert any(m["content"] == "previous message" for m in session_after.messages)
assert any(m["content"] == "interrupted response" for m in session_after.messages) assert any(m["content"] == "interrupted response" for m in session_after.messages)
await loop.close_mcp() await loop.aclose()
class TestAutoCompactIntegration: class TestAutoCompactIntegration:
@@ -779,7 +784,7 @@ class TestAutoCompactIntegration:
# The new message should be processed (response exists) # The new message should be processed (response exists)
assert response is not None assert response is not None
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_runtime_context_markers_not_persisted_for_multi_paragraph_turn(self, tmp_path): async def test_runtime_context_markers_not_persisted_for_multi_paragraph_turn(self, tmp_path):
@@ -807,7 +812,7 @@ class TestAutoCompactIntegration:
content = str(persisted.get("content", "")) content = str(persisted.get("content", ""))
assert "[Runtime Context" not in content assert "[Runtime Context" not in content
assert "[/Runtime Context]" not in content assert "[/Runtime Context]" not in content
await loop.close_mcp() await loop.aclose()
class TestProactiveAutoCompact: class TestProactiveAutoCompact:
@@ -870,7 +875,7 @@ class TestProactiveAutoCompact:
session_after = loop.sessions.get_or_create("cli:test") session_after = loop.sessions.get_or_create("cli:test")
assert len(session_after.messages) == 1 assert len(session_after.messages) == 1
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_proactive_archive_on_idle_tick(self, tmp_path): async def test_proactive_archive_on_idle_tick(self, tmp_path):
@@ -897,7 +902,7 @@ class TestProactiveAutoCompact:
entry = loop.auto_compact._summaries.get("cli:test") entry = loop.auto_compact._summaries.get("cli:test")
assert entry is not None assert entry is not None
assert entry[0] == "User chatted about old things." assert entry[0] == "User chatted about old things."
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_proactive_archive_skips_dream_sessions(self, tmp_path): async def test_proactive_archive_skips_dream_sessions(self, tmp_path):
@@ -918,7 +923,7 @@ class TestProactiveAutoCompact:
assert _fake_compact.state["count"] == 0 assert _fake_compact.state["count"] == 0
assert "dream:20260602-155256" not in loop.auto_compact._archiving assert "dream:20260602-155256" not in loop.auto_compact._archiving
assert "dream:20260602-155256" not in loop.auto_compact._summaries assert "dream:20260602-155256" not in loop.auto_compact._summaries
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_no_proactive_archive_when_active(self, tmp_path): async def test_no_proactive_archive_when_active(self, tmp_path):
@@ -932,7 +937,7 @@ class TestProactiveAutoCompact:
session_after = loop.sessions.get_or_create("cli:test") session_after = loop.sessions.get_or_create("cli:test")
assert len(session_after.messages) == 1 assert len(session_after.messages) == 1
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_no_duplicate_archive(self, tmp_path): async def test_no_duplicate_archive(self, tmp_path):
@@ -968,7 +973,7 @@ class TestProactiveAutoCompact:
# Clean up # Clean up
block_forever.set() block_forever.set()
await _drain_background_tasks(loop) await _drain_background_tasks(loop)
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_proactive_archive_error_does_not_block(self, tmp_path): async def test_proactive_archive_error_does_not_block(self, tmp_path):
@@ -989,7 +994,7 @@ class TestProactiveAutoCompact:
# Key should be removed from _archiving (finally block) # Key should be removed from _archiving (finally block)
assert "cli:test" not in loop.auto_compact._archiving assert "cli:test" not in loop.auto_compact._archiving
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_proactive_archive_skips_empty_sessions(self, tmp_path): async def test_proactive_archive_skips_empty_sessions(self, tmp_path):
@@ -1005,7 +1010,7 @@ class TestProactiveAutoCompact:
# Empty session should not produce a summary # Empty session should not produce a summary
assert "cli:test" not in loop.auto_compact._summaries assert "cli:test" not in loop.auto_compact._summaries
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_skip_expired_session_with_active_agent_task(self, tmp_path): async def test_skip_expired_session_with_active_agent_task(self, tmp_path):
@@ -1026,7 +1031,7 @@ class TestProactiveAutoCompact:
session_after = loop.sessions.get_or_create("cli:test") session_after = loop.sessions.get_or_create("cli:test")
assert len(session_after.messages) == 12 # All messages preserved assert len(session_after.messages) == 12 # All messages preserved
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_archive_after_active_task_completes(self, tmp_path): async def test_archive_after_active_task_completes(self, tmp_path):
@@ -1047,7 +1052,7 @@ class TestProactiveAutoCompact:
# Second tick: task completed, should archive # Second tick: task completed, should archive
await self._run_check_expired(loop) await self._run_check_expired(loop)
assert _fake_compact.state["count"] == 1 assert _fake_compact.state["count"] == 1
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_partial_active_set_only_archives_inactive_expired(self, tmp_path): async def test_partial_active_set_only_archives_inactive_expired(self, tmp_path):
@@ -1083,7 +1088,7 @@ class TestProactiveAutoCompact:
assert len(s2_after.messages) == 12 # Preserved assert len(s2_after.messages) == 12 # Preserved
s3_after = loop.sessions.get_or_create("cli:recent") s3_after = loop.sessions.get_or_create("cli:recent")
assert len(s3_after.messages) == 1 # Preserved assert len(s3_after.messages) == 1 # Preserved
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_no_reschedule_after_successful_archive(self, tmp_path): async def test_no_reschedule_after_successful_archive(self, tmp_path):
@@ -1104,7 +1109,7 @@ class TestProactiveAutoCompact:
# Second tick: should NOT re-schedule because the session has no removable tail. # Second tick: should NOT re-schedule because the session has no removable tail.
await self._run_check_expired(loop) await self._run_check_expired(loop)
assert _fake_compact.state["count"] == 1 # Still 1, not re-scheduled assert _fake_compact.state["count"] == 1 # Still 1, not re-scheduled
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_empty_session_does_not_schedule_idle_compact(self, tmp_path): async def test_empty_session_does_not_schedule_idle_compact(self, tmp_path):
@@ -1124,7 +1129,7 @@ class TestProactiveAutoCompact:
await self._run_check_expired(loop) await self._run_check_expired(loop)
assert _fake_compact.state["count"] == 0 assert _fake_compact.state["count"] == 0
assert "cli:test" not in loop.auto_compact._summaries assert "cli:test" not in loop.auto_compact._summaries
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_session_can_be_compacted_again_after_new_messages(self, tmp_path): async def test_session_can_be_compacted_again_after_new_messages(self, tmp_path):
@@ -1155,7 +1160,7 @@ class TestProactiveAutoCompact:
# Second compact cycle should succeed # Second compact cycle should succeed
await loop.auto_compact._archive("cli:test", runtime=loop.llm_runtime()) await loop.auto_compact._archive("cli:test", runtime=loop.llm_runtime())
assert _fake_compact.state["count"] == 2 assert _fake_compact.state["count"] == 2
await loop.close_mcp() await loop.aclose()
class TestSummaryPersistence: class TestSummaryPersistence:
@@ -1182,7 +1187,7 @@ class TestSummaryPersistence:
assert meta is not None assert meta is not None
assert meta["text"] == "User said hello." assert meta["text"] == "User said hello."
assert "last_active" in meta assert "last_active" in meta
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_summary_recovered_after_restart(self, tmp_path): async def test_summary_recovered_after_restart(self, tmp_path):
@@ -1218,7 +1223,7 @@ class TestSummaryPersistence:
assert "Previous conversation summary" in summary assert "Previous conversation summary" in summary
# _last_summary persists in metadata for restart survival. # _last_summary persists in metadata for restart survival.
assert "_last_summary" in reloaded.metadata assert "_last_summary" in reloaded.metadata
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_metadata_persists_for_restart(self, tmp_path): async def test_metadata_persists_for_restart(self, tmp_path):
@@ -1246,7 +1251,7 @@ class TestSummaryPersistence:
assert "Summary." in summary2 assert "Summary." in summary2
# _last_summary persists in metadata for restart survival. # _last_summary persists in metadata for restart survival.
assert "_last_summary" in reloaded.metadata assert "_last_summary" in reloaded.metadata
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_metadata_cleanup_on_inmemory_path(self, tmp_path): async def test_metadata_cleanup_on_inmemory_path(self, tmp_path):
@@ -1272,7 +1277,7 @@ class TestSummaryPersistence:
assert summary is not None assert summary is not None
# _last_summary persists in metadata for restart survival. # _last_summary persists in metadata for restart survival.
assert "_last_summary" in reloaded.metadata assert "_last_summary" in reloaded.metadata
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_new_summary_overrides_old(self, tmp_path): async def test_new_summary_overrides_old(self, tmp_path):
@@ -1314,7 +1319,7 @@ class TestSummaryPersistence:
_, summary2 = loop.auto_compact.prepare_session(reloaded, "cli:test") _, summary2 = loop.auto_compact.prepare_session(reloaded, "cli:test")
assert summary2 is not None assert summary2 is not None
assert "Second summary." in summary2 assert "Second summary." in summary2
await loop.close_mcp() await loop.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_new_command_clears_last_summary(self, tmp_path): async def test_new_command_clears_last_summary(self, tmp_path):
@@ -1342,4 +1347,4 @@ class TestSummaryPersistence:
# After /new, metadata should no longer contain _last_summary # After /new, metadata should no longer contain _last_summary
fresh = loop.sessions.get_or_create("cli:test") fresh = loop.sessions.get_or_create("cli:test")
assert "_last_summary" not in fresh.metadata assert "_last_summary" not in fresh.metadata
await loop.close_mcp() await loop.aclose()
+5 -5
View File
@@ -538,7 +538,7 @@ class TestNewCommandArchival:
session_after = loop.sessions.get_or_create("cli:test") session_after = loop.sessions.get_or_create("cli:test")
assert len(session_after.messages) == 0 assert len(session_after.messages) == 0
await loop.close_mcp() await loop.aclose()
assert call_count == 1 assert call_count == 1
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -572,7 +572,7 @@ class TestNewCommandArchival:
assert response is not None assert response is not None
assert "new session started" in response.content.lower() assert "new session started" in response.content.lower()
await loop.close_mcp() await loop.aclose()
assert archived_count == 3 assert archived_count == 3
assert archived_session_key == "cli:test" assert archived_session_key == "cli:test"
@@ -603,8 +603,8 @@ class TestNewCommandArchival:
assert loop.sessions.get_or_create("cli:test").messages == [] assert loop.sessions.get_or_create("cli:test").messages == []
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_close_mcp_drains_background_tasks(self, tmp_path: Path) -> None: async def test_aclose_drains_background_tasks(self, tmp_path: Path) -> None:
"""close_mcp waits for background tasks to complete.""" """aclose waits for background tasks to complete."""
from nanobot.bus.events import InboundMessage from nanobot.bus.events import InboundMessage
loop = self._make_loop(tmp_path) loop = self._make_loop(tmp_path)
@@ -632,5 +632,5 @@ class TestNewCommandArchival:
assert not archived.is_set() assert not archived.is_set()
release_archive.set() release_archive.set()
await loop.close_mcp() await loop.aclose()
assert archived.is_set() assert archived.is_set()
@@ -93,7 +93,6 @@ async def test_process_direct_websocket_clears_run_status(tmp_path) -> None:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_process_direct_reuses_existing_session_lock(tmp_path) -> None: async def test_process_direct_reuses_existing_session_lock(tmp_path) -> None:
loop = _make_loop(tmp_path) loop = _make_loop(tmp_path)
loop._connect_mcp = AsyncMock()
session_key = "api:fixed" session_key = "api:fixed"
lock = loop._session_locks.setdefault(session_key, asyncio.Lock()) lock = loop._session_locks.setdefault(session_key, asyncio.Lock())
await lock.acquire() await lock.acquire()
-3
View File
@@ -1519,13 +1519,11 @@ async def test_run_agent_loop_goal_continue_message_reads_latest_metadata(
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_process_direct_rejects_reserved_system_channel(tmp_path: Path) -> None: async def test_process_direct_rejects_reserved_system_channel(tmp_path: Path) -> None:
loop = _make_full_loop(tmp_path) loop = _make_full_loop(tmp_path)
loop._connect_mcp = AsyncMock() # type: ignore[method-assign]
loop._process_message = AsyncMock(return_value=None) # type: ignore[method-assign] loop._process_message = AsyncMock(return_value=None) # type: ignore[method-assign]
with pytest.raises(ValueError, match="reserved for internal messages"): with pytest.raises(ValueError, match="reserved for internal messages"):
await loop.process_direct("external input", channel="system") await loop.process_direct("external input", channel="system")
loop._connect_mcp.assert_not_awaited()
loop._process_message.assert_not_awaited() loop._process_message.assert_not_awaited()
@@ -1534,7 +1532,6 @@ async def test_process_direct_skip_user_persist_does_not_save_retry_user(
tmp_path: Path, tmp_path: Path,
) -> None: ) -> None:
loop = _make_full_loop(tmp_path) loop = _make_full_loop(tmp_path)
loop._connect_mcp = AsyncMock()
session = loop.sessions.get_or_create("api:default") session = loop.sessions.get_or_create("api:default")
session.add_message("user", "hello") session.add_message("user", "hello")
session.add_message("assistant", "previous empty-response attempt") session.add_message("assistant", "previous empty-response attempt")
+1 -2
View File
@@ -123,8 +123,7 @@ async def test_session_discard_control_cancels_active_turn(tmp_path, monkeypatch
await asyncio.sleep(0) await asyncio.sleep(0)
loop.provider.chat_with_retry = AsyncMock(side_effect=block_provider) loop.provider.chat_with_retry = AsyncMock(side_effect=block_provider)
monkeypatch.setattr(loop, "_connect_mcp", AsyncMock()) monkeypatch.setattr(loop, "aclose", AsyncMock())
monkeypatch.setattr(loop, "close_mcp", AsyncMock())
terminate_exec_sessions = AsyncMock(return_value=1) terminate_exec_sessions = AsyncMock(return_value=1)
monkeypatch.setattr( monkeypatch.setattr(
loop._exec_session_manager, loop._exec_session_manager,
+48
View File
@@ -1,4 +1,5 @@
import asyncio import asyncio
import inspect
from pathlib import Path from pathlib import Path
from unittest.mock import AsyncMock, MagicMock from unittest.mock import AsyncMock, MagicMock
@@ -11,8 +12,10 @@ from nanobot.agent.tools.context import (
current_request_context, current_request_context,
reset_request_context, reset_request_context,
) )
from nanobot.agent.tools.registry import ToolRegistry
from nanobot.bus.events import InboundMessage from nanobot.bus.events import InboundMessage
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.config.schema import Config
from nanobot.providers.base import LLMResponse, ToolCallRequest from nanobot.providers.base import LLMResponse, ToolCallRequest
from nanobot.session.turn_continuation import INTERNAL_CONTINUATION_META from nanobot.session.turn_continuation import INTERNAL_CONTINUATION_META
@@ -56,6 +59,51 @@ class _Tools:
return (self.tool, arguments, None) if name == "cron" else (None, arguments, None) return (self.tool, arguments, None) if name == "cron" else (None, arguments, None)
def test_loop_registers_default_tools_in_injected_registry(tmp_path: Path) -> None:
provider = MagicMock()
provider.get_default_model.return_value = "test-model"
registry = ToolRegistry()
loop = AgentLoop(
bus=MessageBus(),
provider=provider,
workspace=tmp_path,
tool_registry=registry,
)
assert loop.tools is registry
assert registry.has("read_file")
def _config_for_loop(tmp_path: Path) -> Config:
return Config.model_validate({"agents": {"defaults": {"workspace": str(tmp_path)}}})
def _provider_for_loop() -> MagicMock:
provider = MagicMock()
provider.get_default_model.return_value = "test-model"
return provider
def test_loop_from_config_requires_caller_owned_registry(tmp_path: Path) -> None:
signature = inspect.signature(AgentLoop.from_config)
with pytest.raises(TypeError, match="tool_registry"):
signature.bind(_config_for_loop(tmp_path))
def test_loop_from_config_uses_caller_owned_registry(tmp_path: Path) -> None:
registry = ToolRegistry()
loop = AgentLoop.from_config(
_config_for_loop(tmp_path),
tool_registry=registry,
provider=_provider_for_loop(),
)
assert loop.tools is registry
assert loop.tools.has("read_file")
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_loop_binds_request_context_for_tool_execution(tmp_path: Path) -> None: async def test_loop_binds_request_context_for_tool_execution(tmp_path: Path) -> None:
provider = MagicMock() provider = MagicMock()
+165 -97
View File
@@ -1,4 +1,4 @@
"""Tests for MCP connection lifecycle in AgentLoop.""" """Tests for the application-owned MCP provider lifecycle."""
from __future__ import annotations from __future__ import annotations
@@ -6,7 +6,7 @@ import asyncio
from contextlib import AsyncExitStack from contextlib import AsyncExitStack
from types import SimpleNamespace from types import SimpleNamespace
from typing import Any from typing import Any
from unittest.mock import AsyncMock, MagicMock from unittest.mock import AsyncMock
import anyio import anyio
import pytest import pytest
@@ -15,11 +15,10 @@ from mcp.shared.exceptions import McpError
from mcp.shared.message import SessionMessage from mcp.shared.message import SessionMessage
from mcp.types import ErrorData from mcp.types import ErrorData
from nanobot.agent.loop import AgentLoop
from nanobot.agent.tools import mcp as mcp_runtime from nanobot.agent.tools import mcp as mcp_runtime
from nanobot.agent.tools.base import Tool from nanobot.agent.tools.base import Tool
from nanobot.agent.tools.mcp import MCPResourceWrapper, MCPToolWrapper from nanobot.agent.tools.mcp import MCPProvider, MCPResourceWrapper, MCPToolWrapper
from nanobot.bus.queue import MessageBus from nanobot.agent.tools.registry import ToolRegistry
from nanobot.config.loader import load_config, save_config from nanobot.config.loader import load_config, save_config
from nanobot.config.schema import MCPServerConfig from nanobot.config.schema import MCPServerConfig
@@ -74,18 +73,20 @@ class _FakeMcpTool(Tool):
return "ok" return "ok"
def _make_loop(tmp_path, *, mcp_servers: dict | None = None) -> AgentLoop: def _stdio_server(command: str = "test-mcp") -> MCPServerConfig:
bus = MessageBus() return MCPServerConfig(type="stdio", command=command)
provider = MagicMock()
provider.get_default_model.return_value = "test-model"
provider.generation.max_tokens = 4096 def _make_provider(
return AgentLoop( *,
bus=bus, mcp_servers: dict[str, MCPServerConfig] | None = None,
provider=provider, ) -> tuple[MCPProvider, ToolRegistry]:
workspace=tmp_path, registry = ToolRegistry()
model="test-model", provider = MCPProvider(
mcp_servers=mcp_servers or {"test": object()}, mcp_servers if mcp_servers is not None else {"test": _stdio_server()},
registry,
) )
return provider, registry
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -140,7 +141,7 @@ async def test_owned_mcp_connection_closes_from_its_owner_task():
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_connect_mcp_retries_when_no_servers_connect(tmp_path, monkeypatch: pytest.MonkeyPatch): 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 attempts = 0
async def _fake_connect(_servers, _registry): 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) monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
await loop._connect_mcp() await provider.connect()
await loop._connect_mcp() await provider.connect()
assert attempts == 2 assert attempts == 2
assert loop._mcp_stacks == {} assert provider.connected_server_names == set()
assert loop.mcp_runtime_status() == {"test": "failed"} assert provider.runtime_status() == {"test": "failed"}
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -168,7 +169,7 @@ async def test_connect_mcp_does_not_report_failure_before_oauth_authorization(
auth="oauth", auth="oauth",
url="https://mcp.example.com/mcp", 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() connect = AsyncMock()
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", connect) monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", connect)
monkeypatch.setattr( monkeypatch.setattr(
@@ -176,19 +177,20 @@ async def test_connect_mcp_does_not_report_failure_before_oauth_authorization(
lambda _name, _url: False, lambda _name, _url: False,
) )
await loop._connect_mcp() await provider.connect()
connect.assert_not_awaited() connect.assert_not_awaited()
assert loop.mcp_runtime_status() == {} assert provider.runtime_status() == {}
@pytest.mark.asyncio @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, tmp_path,
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
): ):
loop = _make_loop(tmp_path, mcp_servers={"playwright": object()}) provider, registry = _make_provider(
connected = asyncio.Event() mcp_servers={"playwright": _stdio_server("playwright")}
)
owner_tasks: list[asyncio.Task | None] = [] owner_tasks: list[asyncio.Task | None] = []
closed_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): async def _fake_connect(servers, _registry):
stacks = {name: _OwnerCheckedStack() for name in servers} stacks = {name: _OwnerCheckedStack() for name in servers}
connected.set()
return stacks return stacks
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect) monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
task = asyncio.create_task(loop.run()) await provider.connect()
await asyncio.wait_for(connected.wait(), timeout=1) registry.register(_FakeMcpTool("mcp_playwright_search"))
loop.stop() await provider.aclose()
task.cancel()
await asyncio.gather(task, return_exceptions=True)
assert owner_tasks assert owner_tasks
assert closed_tasks == 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 @pytest.mark.asyncio
async def test_close_server_ignores_server_cancelled_error(tmp_path): async def test_close_server_ignores_server_cancelled_error(tmp_path):
loop = _make_loop(tmp_path) provider, _registry = _make_provider()
class _ServerCancelledStack: class _ServerCancelledStack:
async def aclose(self) -> None: async def aclose(self) -> None:
raise asyncio.CancelledError() 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 @pytest.mark.asyncio
async def test_close_mcp_servers_continues_after_server_cancelled_error(tmp_path): async def test_provider_close_continues_after_server_cancelled_error(tmp_path):
loop = _make_loop(tmp_path) provider, _registry = _make_provider()
closed: list[str] = [] closed: list[str] = []
class _ServerCancelledStack: class _ServerCancelledStack:
@@ -247,21 +247,53 @@ async def test_close_mcp_servers_continues_after_server_cancelled_error(tmp_path
async def aclose(self) -> None: async def aclose(self) -> None:
closed.append("second") closed.append("second")
loop._mcp_stacks = { provider._connections = {
"first": _ServerCancelledStack(), "first": _ServerCancelledStack(),
"second": _TrackedStack(), "second": _TrackedStack(),
} }
await mcp_runtime.close_mcp_servers(loop) await provider.aclose()
assert closed == ["second"] 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.asyncio
@pytest.mark.parametrize("close_all", [False, True], ids=["single", "all"]) @pytest.mark.parametrize("close_all", [False, True], ids=["single", "all"])
async def test_mcp_cleanup_re_raises_external_cancellation(tmp_path, close_all: bool): 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() started = asyncio.Event()
class _BlockingStack: class _BlockingStack:
@@ -269,12 +301,12 @@ async def test_mcp_cleanup_re_raises_external_cancellation(tmp_path, close_all:
started.set() started.set()
await asyncio.Event().wait() await asyncio.Event().wait()
loop._mcp_stacks = {"test": _BlockingStack()} provider._connections = {"test": _BlockingStack()}
if close_all: if close_all:
task = asyncio.create_task(mcp_runtime.close_mcp_servers(loop)) task = asyncio.create_task(provider.aclose())
else: 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) await asyncio.wait_for(started.wait(), timeout=1)
task.cancel() task.cancel()
@@ -312,41 +344,38 @@ async def test_reload_mcp_servers_adds_and_removes_tools_without_restart(
return stacks return stacks
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect) 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["ok"] is True
assert added["added"] == ["browserbase"] assert added["added"] == ["browserbase"]
assert loop.tools.has("mcp_browserbase_navigate") assert registry.has("mcp_browserbase_navigate")
assert "browserbase" in loop._mcp_stacks assert provider.connected_server_names == {"browserbase"}
config = load_config() config = load_config()
del config.tools.mcp_servers["browserbase"] del config.tools.mcp_servers["browserbase"]
save_config(config) save_config(config)
removed = await mcp_runtime.reload_servers(loop, loop.tools) removed = await provider.reload()
assert removed["ok"] is True assert removed["ok"] is True
assert removed["removed"] == ["browserbase"] assert removed["removed"] == ["browserbase"]
assert not loop.tools.has("mcp_browserbase_navigate") assert not registry.has("mcp_browserbase_navigate")
assert "browserbase" not in loop._mcp_stacks assert provider.connected_server_names == set()
assert closed == ["browserbase"] assert closed == ["browserbase"]
@pytest.mark.asyncio @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, tmp_path,
monkeypatch: pytest.MonkeyPatch, monkeypatch: pytest.MonkeyPatch,
): ):
config_path = tmp_path / "config.json" browserbase = MCPServerConfig(
monkeypatch.setattr("nanobot.config.loader._current_config_path", config_path)
config = load_config()
config.tools.mcp_servers["browserbase"] = MCPServerConfig(
type="stdio", type="stdio",
command="browserbase-mcp", command="browserbase-mcp",
) )
save_config(config) configured: dict[str, MCPServerConfig] = {"browserbase": browserbase}
closed: list[str] = [] closed: list[str] = []
@@ -364,37 +393,68 @@ async def test_request_mcp_reload_reaches_runtime_control_without_restart(
return stacks return stacks
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect) 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: result = await provider.reload()
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
assert result["ok"] is True assert result["ok"] is True
assert result["added"] == ["browserbase"] assert result["added"] == ["browserbase"]
assert result["requires_restart"] is False assert result["requires_restart"] is False
assert loop.tools.has("mcp_browserbase_navigate") assert registry.has("mcp_browserbase_navigate")
config = load_config() configured = {}
del config.tools.mcp_servers["browserbase"]
save_config(config)
consumer = asyncio.create_task(_handle_one_runtime_control()) result = await provider.reload()
result = await mcp_runtime.request_mcp_reload(loop.bus, timeout=2.0)
await consumer
assert result["ok"] is True assert result["ok"] is True
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 registry.has("mcp_browserbase_navigate")
assert closed == ["browserbase"] 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 @pytest.mark.asyncio
async def test_reload_mcp_servers_retries_configured_server_without_live_stack( async def test_reload_mcp_servers_retries_configured_server_without_live_stack(
tmp_path, tmp_path,
@@ -419,16 +479,18 @@ async def test_reload_mcp_servers_retries_configured_server_without_live_stack(
return stacks return stacks
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect) 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["ok"] is True
assert result["added"] == [] assert result["added"] == []
assert result["changed"] == [] assert result["changed"] == []
assert result["retried"] == ["browserbase"] assert result["retried"] == ["browserbase"]
assert loop.tools.has("mcp_browserbase_navigate") assert registry.has("mcp_browserbase_navigate")
await loop.close_mcp() await provider.aclose()
@pytest.mark.asyncio @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", "nanobot.agent.tools.mcp_oauth.mcp_oauth_has_credentials",
lambda name, _url: name == "linear", 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 attempted == ["linear"]
assert result["ok"] is True assert result["ok"] is True
assert result["failed"] == [] assert result["failed"] == []
assert result["retried"] == [] assert result["retried"] == []
assert result["connected"] == ["linear"] assert result["connected"] == ["linear"]
await loop.close_mcp() await provider.aclose()
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -482,7 +544,9 @@ async def test_mcp_tool_reconnects_after_session_terminated(
tmp_path, tmp_path,
monkeypatch: pytest.MonkeyPatch, 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] = [] closed: list[str] = []
sessions: list[Any] = [] sessions: list[Any] = []
connect_count = 0 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) monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
await loop._connect_mcp() await provider.connect()
old_tool = loop.tools.get("mcp_remote_quote") old_tool = registry.get("mcp_remote_quote")
assert isinstance(old_tool, MCPToolWrapper) assert isinstance(old_tool, MCPToolWrapper)
output = await old_tool.execute(symbol="AAPL") output = await old_tool.execute(symbol="AAPL")
@@ -536,8 +600,8 @@ async def test_mcp_tool_reconnects_after_session_terminated(
assert closed == ["remote"] 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 provider.connected_server_names == {"remote"}
assert loop.tools.get("mcp_remote_quote") is not old_tool assert registry.get("mcp_remote_quote") is not old_tool
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -545,7 +609,9 @@ async def test_mcp_reconnect_handler_uses_sanitized_server_prefix(
tmp_path, tmp_path,
monkeypatch: pytest.MonkeyPatch, 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 connect_count = 0
class _FakeSession: 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) monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
await loop._connect_mcp() await provider.connect()
old_tool = loop.tools.get("mcp_remote_quote") old_tool = registry.get("mcp_remote_quote")
assert isinstance(old_tool, MCPToolWrapper) assert isinstance(old_tool, MCPToolWrapper)
output = await old_tool.execute() output = await old_tool.execute()
assert output == "recovered" assert output == "recovered"
assert connect_count == 2 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 @pytest.mark.asyncio
@@ -594,7 +660,9 @@ async def test_concurrent_mcp_reconnect_reuses_fresh_session(
tmp_path, tmp_path,
monkeypatch: pytest.MonkeyPatch, 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] = [] closed: list[str] = []
connect_count = 0 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) monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
await loop._connect_mcp() await provider.connect()
old_alpha = loop.tools.get("mcp_remote_resource_alpha") old_alpha = registry.get("mcp_remote_resource_alpha")
old_beta = loop.tools.get("mcp_remote_resource_beta") old_beta = registry.get("mcp_remote_resource_beta")
assert isinstance(old_alpha, MCPResourceWrapper) assert isinstance(old_alpha, MCPResourceWrapper)
assert isinstance(old_beta, MCPResourceWrapper) assert isinstance(old_beta, MCPResourceWrapper)
+15 -26
View File
@@ -15,15 +15,13 @@ import asyncio
import multiprocessing import multiprocessing
import socket import socket
import time import time
from unittest.mock import MagicMock
import httpx import httpx
import pytest import pytest
from nanobot.agent.loop import AgentLoop
from nanobot.agent.tools import mcp as mcp_module from nanobot.agent.tools import mcp as mcp_module
from nanobot.agent.tools.mcp import MCPToolWrapper from nanobot.agent.tools.mcp import MCPProvider, MCPToolWrapper
from nanobot.bus.queue import MessageBus from nanobot.agent.tools.registry import ToolRegistry
from nanobot.config.schema import MCPServerConfig from nanobot.config.schema import MCPServerConfig
from nanobot.security import network as security_network from nanobot.security import network as security_network
@@ -113,18 +111,9 @@ def mcp_server_url():
process.join(timeout=2.0) process.join(timeout=2.0)
def _make_loop(tmp_path, *, mcp_servers: dict) -> AgentLoop: def _make_provider(*, mcp_servers: dict) -> tuple[MCPProvider, ToolRegistry]:
bus = MessageBus() registry = ToolRegistry()
provider = MagicMock() return MCPProvider(mcp_servers, registry), registry
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,
)
@pytest.fixture(autouse=True) @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, tool_timeout=_TOOL_TIMEOUT_SECONDS,
enabled_tools=["*"], 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()) await asyncio.create_task(provider.connect())
assert "repro" in loop._mcp_stacks assert provider.connected_server_names == {"repro"}
tool = loop.tools.get("mcp_repro_greet") tool = registry.get("mcp_repro_greet")
assert isinstance(tool, MCPToolWrapper) assert isinstance(tool, MCPToolWrapper)
output = await asyncio.create_task(tool.execute(name="first")) 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")) output = await asyncio.create_task(tool.execute(name="second"))
assert "Hello, second" in output assert "Hello, second" in output
await asyncio.create_task(loop.close_mcp()) await asyncio.create_task(provider.aclose())
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -203,10 +192,10 @@ async def test_mcp_reconnect_during_shutdown_does_not_crash(
tool_timeout=_TOOL_TIMEOUT_SECONDS, tool_timeout=_TOOL_TIMEOUT_SECONDS,
enabled_tools=["*"], 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()) await asyncio.create_task(provider.connect())
tool = loop.tools.get("mcp_repro_greet") tool = registry.get("mcp_repro_greet")
assert isinstance(tool, MCPToolWrapper) assert isinstance(tool, MCPToolWrapper)
await asyncio.create_task(tool.execute(name="first")) 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) 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"))
await asyncio.wait_for(reconnect_started.wait(), timeout=5) 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) await asyncio.sleep(0)
finish_reconnect.set() finish_reconnect.set()
@@ -245,4 +234,4 @@ async def test_mcp_reconnect_during_shutdown_does_not_crash(
unhandled.append(exc) unhandled.append(exc)
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 == {} assert provider.connected_server_names == set()
+6 -1
View File
@@ -7,6 +7,7 @@ from unittest.mock import MagicMock
import pytest import pytest
from nanobot.agent.loop import AgentLoop from nanobot.agent.loop import AgentLoop
from nanobot.agent.tools.registry import ToolRegistry
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
from nanobot.bus.runtime_events import RuntimeModelChanged from nanobot.bus.runtime_events import RuntimeModelChanged
from nanobot.config.errors import ConfigLoadError from nanobot.config.errors import ConfigLoadError
@@ -312,7 +313,11 @@ def test_settings_context_window_refreshes_runtime_state(
def loader(*, preset_name: str | None = None) -> ProviderSnapshot: def loader(*, preset_name: str | None = None) -> ProviderSnapshot:
return load_provider_snapshot(config_path, preset_name=preset_name) return load_provider_snapshot(config_path, preset_name=preset_name)
loop = AgentLoop.from_config(config, provider_snapshot_loader=loader) loop = AgentLoop.from_config(
config,
tool_registry=ToolRegistry(),
provider_snapshot_loader=loader,
)
payload = update_agent_settings({"context_window_tokens": ["262144"]}) payload = update_agent_settings({"context_window_tokens": ["262144"]})
loop.runtime_resolver.invalidate() loop.runtime_resolver.invalidate()
+3 -2
View File
@@ -5,6 +5,7 @@ import pytest
from nanobot.agent.loop import AgentLoop from nanobot.agent.loop import AgentLoop
from nanobot.agent.tools.context import RequestContext, request_context from nanobot.agent.tools.context import RequestContext, request_context
from nanobot.agent.tools.registry import ToolRegistry
from nanobot.agent.tools.runtime_control import AgentRuntimeControl from nanobot.agent.tools.runtime_control import AgentRuntimeControl
from nanobot.agent.tools.self import MyTool from nanobot.agent.tools.self import MyTool
from nanobot.bus.queue import MessageBus from nanobot.bus.queue import MessageBus
@@ -390,7 +391,7 @@ def test_from_config_injects_default_preset(tmp_path) -> None:
}) })
fake_provider = _provider("openai/gpt-4.1") fake_provider = _provider("openai/gpt-4.1")
with patch("nanobot.providers.factory.make_provider", return_value=fake_provider): with patch("nanobot.providers.factory.make_provider", return_value=fake_provider):
loop = AgentLoop.from_config(config) loop = AgentLoop.from_config(config, tool_registry=ToolRegistry())
assert loop.model == "openai/gpt-4.1" assert loop.model == "openai/gpt-4.1"
assert loop.model_preset is None assert loop.model_preset is None
assert "default" in loop.model_presets assert "default" in loop.model_presets
@@ -407,7 +408,7 @@ def test_from_config_static_preset_loader_does_not_enable_hot_reload(tmp_path) -
}) })
fake_provider = _provider("openai/gpt-4.1") fake_provider = _provider("openai/gpt-4.1")
with patch("nanobot.providers.factory.make_provider", return_value=fake_provider): with patch("nanobot.providers.factory.make_provider", return_value=fake_provider):
loop = AgentLoop.from_config(config) loop = AgentLoop.from_config(config, tool_registry=ToolRegistry())
default_runtime = loop.runtime_resolver.runtime default_runtime = loop.runtime_resolver.runtime
resolved = loop.runtime_resolver.resolve_preset("fast") resolved = loop.runtime_resolver.resolve_preset("fast")
assert resolved.model == "openai/gpt-4.1-mini" assert resolved.model == "openai/gpt-4.1-mini"
+11 -14
View File
@@ -56,7 +56,7 @@ class TestHandleStop:
assert "No active task" in out.content assert "No active task" in out.content
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_close_mcp_cancels_active_turn_before_resources(self): async def test_aclose_cancels_active_turn_before_resources(self):
loop, _bus = _make_loop() loop, _bus = _make_loop()
events: list[str] = [] events: list[str] = []
@@ -76,14 +76,13 @@ class TestHandleStop:
loop.subagents.close = close_subagents loop.subagents.close = close_subagents
loop._exec_session_manager.close_all = AsyncMock() loop._exec_session_manager.close_all = AsyncMock()
with patch("nanobot.agent.loop.agent_context.close_mcp", AsyncMock()): await loop.aclose()
await loop.close_mcp()
assert events == ["turn_cancelled", "resources_closed"] assert events == ["turn_cancelled", "resources_closed"]
assert task.cancelled() assert task.cancelled()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_close_mcp_serializes_duplicate_cleanup(self): async def test_aclose_serializes_duplicate_cleanup(self):
loop, _bus = _make_loop() loop, _bus = _make_loop()
entered = asyncio.Event() entered = asyncio.Event()
release = asyncio.Event() release = asyncio.Event()
@@ -100,14 +99,13 @@ class TestHandleStop:
loop.subagents.close = close_subagents loop.subagents.close = close_subagents
loop._exec_session_manager.close_all = AsyncMock() loop._exec_session_manager.close_all = AsyncMock()
with patch("nanobot.agent.loop.agent_context.close_mcp", AsyncMock()): first = asyncio.create_task(loop.aclose())
first = asyncio.create_task(loop.close_mcp()) await entered.wait()
await entered.wait() second = asyncio.create_task(loop.aclose())
second = asyncio.create_task(loop.close_mcp()) await asyncio.sleep(0)
await asyncio.sleep(0) assert not second.done()
assert not second.done() release.set()
release.set() await asyncio.gather(first, second)
await asyncio.gather(first, second)
assert max_concurrent == 1 assert max_concurrent == 1
@@ -172,8 +170,7 @@ class TestDispatch:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_run_logs_and_continues_after_leaked_cancelled_error(self, monkeypatch): async def test_run_logs_and_continues_after_leaked_cancelled_error(self, monkeypatch):
loop, bus = _make_loop() loop, bus = _make_loop()
loop._connect_mcp = AsyncMock() loop.aclose = AsyncMock()
loop.close_mcp = AsyncMock()
loop.auto_compact.check_expired = MagicMock() loop.auto_compact.check_expired = MagicMock()
warnings: list[str] = [] warnings: list[str] = []
calls = 0 calls = 0
-21
View File
@@ -493,20 +493,6 @@ class TestModifyOpen:
assert "Set workspace" in result assert "Set workspace" in result
assert tool._runtime_control.snapshot().workspace == "/new/path" assert tool._runtime_control.snapshot().workspace == "/new/path"
@pytest.mark.asyncio
async def test_modify_mcp_servers_blocked(self):
"""_mcp_servers contains API credentials — must be blocked."""
tool = _make_tool()
result = await tool.execute(action="set", key="_mcp_servers", value={"evil": "leaked"})
assert "protected" in result
@pytest.mark.asyncio
async def test_modify_mcp_stacks_blocked(self):
"""_mcp_stacks holds connection handles — must be blocked."""
tool = _make_tool()
result = await tool.execute(action="set", key="_mcp_stacks", value={})
assert "protected" in result
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_modify_pending_queues_blocked(self): async def test_modify_pending_queues_blocked(self):
"""_pending_queues controls message routing — must be blocked.""" """_pending_queues controls message routing — must be blocked."""
@@ -535,13 +521,6 @@ class TestModifyOpen:
result = await tool.execute(action="set", key="_background_tasks", value=[]) result = await tool.execute(action="set", key="_background_tasks", value=[])
assert "protected" in result assert "protected" in result
@pytest.mark.asyncio
async def test_inspect_mcp_servers_blocked(self):
"""_mcp_servers contains credentials — check must be blocked too."""
tool = _make_tool()
result = await tool.execute(action="check", key="_mcp_servers")
assert "not accessible" in result
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_modify_wrapped_denied(self): async def test_modify_wrapped_denied(self):
"""__wrapped__ allows decorator bypass — must be denied.""" """__wrapped__ allows decorator bypass — must be denied."""
+4 -4
View File
@@ -64,7 +64,7 @@ def test_interactive_agent_routes_a_complete_user_turn(
def __init__(self, bus) -> None: def __init__(self, bus) -> None:
self.bus = bus self.bus = bus
self.stopped = asyncio.Event() self.stopped = asyncio.Event()
self.close_mcp_calls = 0 self.aclose_calls = 0
async def run(self) -> None: async def run(self) -> None:
message = await self.bus.consume_inbound() message = await self.bus.consume_inbound()
@@ -97,8 +97,8 @@ def test_interactive_agent_routes_a_complete_user_turn(
def stop(self) -> None: def stop(self) -> None:
self.stopped.set() self.stopped.set()
async def close_mcp(self) -> None: async def aclose(self) -> None:
self.close_mcp_calls += 1 self.aclose_calls += 1
read_input = AsyncMock(side_effect=["hello nanobot", "exit"]) read_input = AsyncMock(side_effect=["hello nanobot", "exit"])
print_response = MagicMock() print_response = MagicMock()
@@ -136,7 +136,7 @@ def test_interactive_agent_routes_a_complete_user_turn(
assert inbound.metadata == {"_wants_stream": True} assert inbound.metadata == {"_wants_stream": True}
loop = seen["loop"] loop = seen["loop"]
assert isinstance(loop, _AgentLoop) assert isinstance(loop, _AgentLoop)
assert loop.close_mcp_calls == 1 assert loop.aclose_calls == 1
assert len(renderers) == 1 assert len(renderers) == 1
renderer = renderers[0] renderer = renderers[0]
assert isinstance(renderer, _Renderer) assert isinstance(renderer, _Renderer)
+50 -19
View File
@@ -1543,7 +1543,7 @@ def mock_agent_runtime(tmp_path):
agent_loop.process_direct = AsyncMock( agent_loop.process_direct = AsyncMock(
return_value=OutboundMessage(channel="cli", chat_id="direct", content="mock-response"), return_value=OutboundMessage(channel="cli", chat_id="direct", content="mock-response"),
) )
agent_loop.close_mcp = AsyncMock(return_value=None) agent_loop.aclose = AsyncMock(return_value=None)
mock_from_config.return_value = agent_loop mock_from_config.return_value = agent_loop
yield { yield {
@@ -1621,7 +1621,7 @@ def test_agent_config_sets_active_path(monkeypatch, tmp_path: Path) -> None:
async def process_direct(self, *_args, **_kwargs): async def process_direct(self, *_args, **_kwargs):
return OutboundMessage(channel="cli", chat_id="direct", content="ok") return OutboundMessage(channel="cli", chat_id="direct", content="ok")
async def close_mcp(self) -> None: async def aclose(self) -> None:
return None return None
monkeypatch.setattr("nanobot.cli.agent.AgentLoop", _FakeAgentLoop) monkeypatch.setattr("nanobot.cli.agent.AgentLoop", _FakeAgentLoop)
@@ -1662,7 +1662,7 @@ def test_agent_uses_workspace_directory_for_cron_store(monkeypatch, tmp_path: Pa
async def process_direct(self, *_args, **_kwargs): async def process_direct(self, *_args, **_kwargs):
return OutboundMessage(channel="cli", chat_id="direct", content="ok") return OutboundMessage(channel="cli", chat_id="direct", content="ok")
async def close_mcp(self) -> None: async def aclose(self) -> None:
return None return None
monkeypatch.setattr("nanobot.cron.service.CronService", _FakeCron) monkeypatch.setattr("nanobot.cron.service.CronService", _FakeCron)
@@ -1712,7 +1712,7 @@ def test_agent_workspace_override_does_not_migrate_legacy_cron(
async def process_direct(self, *_args, **_kwargs): async def process_direct(self, *_args, **_kwargs):
return OutboundMessage(channel="cli", chat_id="direct", content="ok") return OutboundMessage(channel="cli", chat_id="direct", content="ok")
async def close_mcp(self) -> None: async def aclose(self) -> None:
return None return None
monkeypatch.setattr("nanobot.cron.service.CronService", _FakeCron) monkeypatch.setattr("nanobot.cron.service.CronService", _FakeCron)
@@ -1768,7 +1768,7 @@ def test_agent_custom_config_workspace_does_not_migrate_legacy_cron(
async def process_direct(self, *_args, **_kwargs): async def process_direct(self, *_args, **_kwargs):
return OutboundMessage(channel="cli", chat_id="direct", content="ok") return OutboundMessage(channel="cli", chat_id="direct", content="ok")
async def close_mcp(self) -> None: async def aclose(self) -> None:
return None return None
monkeypatch.setattr("nanobot.cron.service.CronService", _FakeCron) monkeypatch.setattr("nanobot.cron.service.CronService", _FakeCron)
@@ -2062,7 +2062,7 @@ def test_heartbeat_empty_response_still_retains_recent_messages(
async def process_direct(self, *_args, **_kwargs): async def process_direct(self, *_args, **_kwargs):
return SimpleNamespace(content="") return SimpleNamespace(content="")
async def close_mcp(self) -> None: async def aclose(self) -> None:
return None return None
async def run(self) -> None: async def run(self) -> None:
@@ -2738,10 +2738,7 @@ def _patch_serve_runtime(monkeypatch, config: Config, seen: dict[str, object]) -
def __init__(self, **kwargs) -> None: def __init__(self, **kwargs) -> None:
seen["workspace"] = kwargs["workspace"] seen["workspace"] = kwargs["workspace"]
async def _connect_mcp(self) -> None: async def aclose(self) -> None:
return None
async def close_mcp(self) -> None:
return None return None
def _fake_create_app( def _fake_create_app(
@@ -2749,11 +2746,13 @@ def _patch_serve_runtime(monkeypatch, config: Config, seen: dict[str, object]) -
model_name: str, model_name: str,
request_timeout: float, request_timeout: float,
api_key: str = "", api_key: str = "",
prepare_agent=None,
): ):
seen["agent_loop"] = agent_loop seen["agent_loop"] = agent_loop
seen["model_name"] = model_name seen["model_name"] = model_name
seen["request_timeout"] = request_timeout seen["request_timeout"] = request_timeout
seen["api_key"] = api_key seen["api_key"] = api_key
seen["prepare_agent"] = prepare_agent
return _FakeApiApp() return _FakeApiApp()
def _fake_run_app(api_app, host: str, port: int, print): def _fake_run_app(api_app, host: str, port: int, print):
@@ -2914,7 +2913,7 @@ def test_gateway_unbound_agent_cron_is_skipped(
async def submit_cron_turn(self, _msg: InboundMessage): async def submit_cron_turn(self, _msg: InboundMessage):
raise AssertionError("unbound cron job must not run as a bound cron turn") raise AssertionError("unbound cron job must not run as a bound cron turn")
async def close_mcp(self) -> None: async def aclose(self) -> None:
return None return None
async def run(self) -> None: async def run(self) -> None:
@@ -3033,7 +3032,7 @@ def test_gateway_bound_cron_runs_as_session_turn(
content="Checked the repo.", content="Checked the repo.",
) )
async def close_mcp(self) -> None: async def aclose(self) -> None:
return None return None
async def run(self) -> None: async def run(self) -> None:
@@ -3253,7 +3252,7 @@ def test_gateway_local_trigger_queue_submits_agent_turns(
self.runtime_resolver.invalidate.assert_called_once_with() self.runtime_resolver.invalidate.assert_called_once_with()
await asyncio.Event().wait() await asyncio.Event().wait()
async def close_mcp(self) -> None: async def aclose(self) -> None:
return None return None
def stop(self) -> None: def stop(self) -> None:
@@ -3499,7 +3498,7 @@ def test_gateway_health_endpoint_binds_and_serves_expected_responses(
async def run(self) -> None: async def run(self) -> None:
await asyncio.Event().wait() await asyncio.Event().wait()
async def close_mcp(self) -> None: async def aclose(self) -> None:
return None return None
def stop(self) -> None: def stop(self) -> None:
@@ -3668,7 +3667,7 @@ def test_gateway_health_endpoint_binds_and_serves_expected_responses(
assert timed_out_writer.output == b"" assert timed_out_writer.output == b""
def test_gateway_shutdown_lets_agent_task_own_mcp_cleanup( def test_gateway_agent_task_owns_initial_mcp_provider_close(
monkeypatch, monkeypatch,
tmp_path: Path, tmp_path: Path,
) -> None: ) -> None:
@@ -3696,17 +3695,41 @@ def test_gateway_shutdown_lets_agent_task_own_mcp_cleanup(
return None return None
async def run(self) -> None: async def run(self) -> None:
seen["agent_task"] = asyncio.current_task()
try: try:
await asyncio.Event().wait() await asyncio.Event().wait()
finally: finally:
seen["agent_task_cleaned_up"] = True seen["agent_task_cleaned_up"] = True
async def close_mcp(self) -> None: async def aclose(self) -> None:
raise AssertionError("gateway must not close MCP from the outer task") seen["agent_closed"] = True
def stop(self) -> None: def stop(self) -> None:
seen["agent_stopped"] = True seen["agent_stopped"] = True
class _FakeMCPProvider:
def __init__(self) -> None:
self.connect_task: asyncio.Task | None = None
self.close_tasks: list[asyncio.Task | None] = []
@classmethod
def from_config(cls, _config, _registry):
provider = cls()
seen["mcp_provider"] = provider
return provider
async def connect(self) -> None:
self.connect_task = asyncio.current_task()
async def aclose(self) -> None:
self.close_tasks.append(asyncio.current_task())
def runtime_status(self) -> dict[str, str]:
return {}
async def reload(self) -> dict[str, object]:
return {"ok": True}
class _FakeChannelManager: class _FakeChannelManager:
def __init__(self, _config, _bus, **_kwargs) -> None: def __init__(self, _config, _bus, **_kwargs) -> None:
self.enabled_channels = ["telegram"] self.enabled_channels = ["telegram"]
@@ -3753,6 +3776,7 @@ def test_gateway_shutdown_lets_agent_task_own_mcp_cleanup(
session_manager=lambda _workspace: object(), session_manager=lambda _workspace: object(),
) )
monkeypatch.setattr("nanobot.cli.gateway_runtime.AgentLoop", _FakeAgentLoop) monkeypatch.setattr("nanobot.cli.gateway_runtime.AgentLoop", _FakeAgentLoop)
monkeypatch.setattr("nanobot.cli.gateway_runtime.MCPProvider", _FakeMCPProvider)
monkeypatch.setattr("nanobot.channels.manager.ChannelManager", _FakeChannelManager) monkeypatch.setattr("nanobot.channels.manager.ChannelManager", _FakeChannelManager)
monkeypatch.setattr("nanobot.cron.service.CronService", _FakeCronService) monkeypatch.setattr("nanobot.cron.service.CronService", _FakeCronService)
monkeypatch.setattr("asyncio.start_server", _fake_start_server) monkeypatch.setattr("asyncio.start_server", _fake_start_server)
@@ -3761,9 +3785,15 @@ def test_gateway_shutdown_lets_agent_task_own_mcp_cleanup(
assert result.exit_code == 0 assert result.exit_code == 0
assert seen["agent_stopped"] is True assert seen["agent_stopped"] is True
assert seen["agent_closed"] is True
assert seen["agent_task_cleaned_up"] is True assert seen["agent_task_cleaned_up"] is True
assert seen["channels_stopped"] is True assert seen["channels_stopped"] is True
assert seen["cron_stopped"] is True assert seen["cron_stopped"] is True
mcp_provider = seen["mcp_provider"]
assert isinstance(mcp_provider, _FakeMCPProvider)
assert mcp_provider.connect_task is seen["agent_task"]
assert mcp_provider.close_tasks[0] is mcp_provider.connect_task
assert len(mcp_provider.close_tasks) == 2
def test_gateway_shutdown_event_exits_forever_runtime_tasks( def test_gateway_shutdown_event_exits_forever_runtime_tasks(
@@ -3800,8 +3830,8 @@ def test_gateway_shutdown_event_exits_forever_runtime_tasks(
finally: finally:
seen["agent_task_cleaned_up"] = True seen["agent_task_cleaned_up"] = True
async def close_mcp(self) -> None: async def aclose(self) -> None:
raise AssertionError("gateway must not close MCP from the outer task") seen["agent_closed"] = True
def stop(self) -> None: def stop(self) -> None:
seen["agent_stopped"] = True seen["agent_stopped"] = True
@@ -3881,6 +3911,7 @@ def test_gateway_shutdown_event_exits_forever_runtime_tasks(
assert result.exit_code == 0 assert result.exit_code == 0
assert seen["agent_stopped"] is True assert seen["agent_stopped"] is True
assert seen["agent_closed"] is True
assert seen["agent_task_cleaned_up"] is True assert seen["agent_task_cleaned_up"] is True
assert seen["channel_task_cleaned_up"] is True assert seen["channel_task_cleaned_up"] is True
assert seen["channels_stopped"] is True assert seen["channels_stopped"] is True
+43 -11
View File
@@ -22,7 +22,7 @@ class _FakeAgent:
self.raise_on_close = False self.raise_on_close = False
self.background: asyncio.Task[None] | None = None self.background: asyncio.Task[None] | None = None
async def close_mcp(self) -> None: async def aclose(self) -> None:
self.close_calls += 1 self.close_calls += 1
if self.hang_on_close: if self.hang_on_close:
await asyncio.sleep(3600) await asyncio.sleep(3600)
@@ -30,7 +30,7 @@ class _FakeAgent:
raise RuntimeError("cleanup exploded") raise RuntimeError("cleanup exploded")
if self.background is not None: if self.background is not None:
await self.background await self.background
self.events.append("close_mcp") self.events.append("aclose")
class _FakeChannels: class _FakeChannels:
@@ -43,6 +43,16 @@ class _FakeChannels:
self.events.append("channels_stopped") self.events.append("channels_stopped")
class _FakeMCPProvider:
def __init__(self, events: list[str] | None = None) -> None:
self.close_calls = 0
self.events = events if events is not None else []
async def aclose(self) -> None:
self.close_calls += 1
self.events.append("mcp_closed")
async def _cancellable_task(events: list[str]) -> None: async def _cancellable_task(events: list[str]) -> None:
try: try:
await asyncio.sleep(3600) await asyncio.sleep(3600)
@@ -64,13 +74,14 @@ async def _stubborn_task(events: list[str]) -> None:
async def test_runtime_tasks_cancelled_before_resources_closed() -> None: async def test_runtime_tasks_cancelled_before_resources_closed() -> None:
events: list[str] = [] events: list[str] = []
agent = _FakeAgent(events) agent = _FakeAgent(events)
provider = _FakeMCPProvider(events)
channels = _FakeChannels() channels = _FakeChannels()
task = asyncio.create_task(_cancellable_task(events)) task = asyncio.create_task(_cancellable_task(events))
await asyncio.sleep(0) # let the task start (cancellation pre-start skips its body) await asyncio.sleep(0) # let the task start (cancellation pre-start skips its body)
await _close_gateway_runtime(agent, channels, [task], None) await _close_gateway_runtime(agent, provider, channels, [task], None)
assert events == ["cancelled", "close_mcp"] # cancel happens before close assert events == ["cancelled", "aclose", "mcp_closed"]
assert channels.stopped == 1 assert channels.stopped == 1
assert agent.close_calls == 1 assert agent.close_calls == 1
assert task.cancelled() assert task.cancelled()
@@ -78,6 +89,7 @@ async def test_runtime_tasks_cancelled_before_resources_closed() -> None:
async def test_pending_background_work_is_drained_before_close_returns() -> None: async def test_pending_background_work_is_drained_before_close_returns() -> None:
agent = _FakeAgent() agent = _FakeAgent()
provider = _FakeMCPProvider()
channels = _FakeChannels() channels = _FakeChannels()
done: dict[str, bool] = {"done": False} done: dict[str, bool] = {"done": False}
@@ -87,7 +99,7 @@ async def test_pending_background_work_is_drained_before_close_returns() -> None
agent.background = asyncio.create_task(background_work()) agent.background = asyncio.create_task(background_work())
await _close_gateway_runtime(agent, channels, [], None) await _close_gateway_runtime(agent, provider, channels, [], None)
assert done["done"] is True assert done["done"] is True
assert agent.close_calls == 1 assert agent.close_calls == 1
@@ -95,6 +107,7 @@ async def test_pending_background_work_is_drained_before_close_returns() -> None
async def test_stubborn_task_does_not_block_past_wait_timeout() -> None: async def test_stubborn_task_does_not_block_past_wait_timeout() -> None:
agent = _FakeAgent() agent = _FakeAgent()
provider = _FakeMCPProvider()
channels = _FakeChannels() channels = _FakeChannels()
events: list[str] = [] events: list[str] = []
task = asyncio.create_task(_stubborn_task(events)) task = asyncio.create_task(_stubborn_task(events))
@@ -104,6 +117,7 @@ async def test_stubborn_task_does_not_block_past_wait_timeout() -> None:
start = time.monotonic() start = time.monotonic()
await _close_gateway_runtime( await _close_gateway_runtime(
agent, agent,
provider,
channels, channels,
[task], [task],
runtime_tasks, runtime_tasks,
@@ -117,70 +131,88 @@ async def test_stubborn_task_does_not_block_past_wait_timeout() -> None:
assert task.done() # the timed-out task received a second cancellation assert task.done() # the timed-out task received a second cancellation
assert runtime_tasks.done() assert runtime_tasks.done()
assert agent.close_calls == 1 # resources still closed underneath it assert agent.close_calls == 1 # resources still closed underneath it
assert provider.close_calls == 1
assert elapsed < 1.0 # bounded, not held open by the stubborn task assert elapsed < 1.0 # bounded, not held open by the stubborn task
async def test_hanging_close_is_bounded_and_does_not_raise() -> None: async def test_hanging_close_is_bounded_and_does_not_raise() -> None:
agent = _FakeAgent() agent = _FakeAgent()
provider = _FakeMCPProvider()
agent.hang_on_close = True agent.hang_on_close = True
channels = _FakeChannels() channels = _FakeChannels()
start = time.monotonic() start = time.monotonic()
await _close_gateway_runtime(agent, channels, [], None, close_timeout=0.05) await _close_gateway_runtime(
agent,
provider,
channels,
[],
None,
close_timeout=0.05,
)
elapsed = time.monotonic() - start elapsed = time.monotonic() - start
assert agent.close_calls == 1 assert agent.close_calls == 1
assert provider.close_calls == 1
assert channels.stopped == 1 assert channels.stopped == 1
assert elapsed < 1.0 assert elapsed < 1.0
async def test_failing_close_is_logged_but_shutdown_proceeds() -> None: async def test_failing_close_is_logged_but_shutdown_proceeds() -> None:
agent = _FakeAgent() agent = _FakeAgent()
provider = _FakeMCPProvider()
agent.raise_on_close = True agent.raise_on_close = True
channels = _FakeChannels() channels = _FakeChannels()
await _close_gateway_runtime(agent, channels, [], None) await _close_gateway_runtime(agent, provider, channels, [], None)
assert agent.close_calls == 1 assert agent.close_calls == 1
assert provider.close_calls == 1
assert channels.stopped == 1 # teardown continued past the failure assert channels.stopped == 1 # teardown continued past the failure
async def test_duplicate_cleanup_is_idempotent() -> None: async def test_duplicate_cleanup_is_idempotent() -> None:
agent = _FakeAgent() agent = _FakeAgent()
provider = _FakeMCPProvider()
channels = _FakeChannels() channels = _FakeChannels()
task = asyncio.create_task(_cancellable_task([])) task = asyncio.create_task(_cancellable_task([]))
await _close_gateway_runtime(agent, channels, [task], None) await _close_gateway_runtime(agent, provider, channels, [task], None)
await _close_gateway_runtime(agent, channels, [task], None) await _close_gateway_runtime(agent, provider, channels, [task], None)
assert agent.close_calls == 2 # second pass is a clean no-op assert agent.close_calls == 2 # second pass is a clean no-op
assert provider.close_calls == 2
assert channels.stopped == 2 assert channels.stopped == 2
assert task.cancelled() assert task.cancelled()
async def test_finished_runtime_tasks_gather_is_retrieved() -> None: async def test_finished_runtime_tasks_gather_is_retrieved() -> None:
agent = _FakeAgent() agent = _FakeAgent()
provider = _FakeMCPProvider()
channels = _FakeChannels() channels = _FakeChannels()
finished = asyncio.get_running_loop().create_future() finished = asyncio.get_running_loop().create_future()
finished.set_result(None) finished.set_result(None)
runtime_tasks = asyncio.gather(finished) runtime_tasks = asyncio.gather(finished)
await asyncio.sleep(0) # let the gather observe the finished child await asyncio.sleep(0) # let the gather observe the finished child
await _close_gateway_runtime(agent, channels, [], runtime_tasks) await _close_gateway_runtime(agent, provider, channels, [], runtime_tasks)
assert runtime_tasks.done() assert runtime_tasks.done()
assert agent.close_calls == 1 assert agent.close_calls == 1
assert provider.close_calls == 1
async def test_cancelled_runtime_tasks_gather_does_not_raise() -> None: async def test_cancelled_runtime_tasks_gather_does_not_raise() -> None:
agent = _FakeAgent() agent = _FakeAgent()
provider = _FakeMCPProvider()
channels = _FakeChannels() channels = _FakeChannels()
runtime_tasks = asyncio.gather(asyncio.sleep(3600)) runtime_tasks = asyncio.gather(asyncio.sleep(3600))
runtime_tasks.cancel() runtime_tasks.cancel()
await _close_gateway_runtime(agent, channels, [], runtime_tasks) await _close_gateway_runtime(agent, provider, channels, [], runtime_tasks)
with suppress(asyncio.CancelledError): with suppress(asyncio.CancelledError):
await runtime_tasks # settle the cancelled gather without raising await runtime_tasks # settle the cancelled gather without raising
assert runtime_tasks.done() # the cancelled gather was awaited without raising assert runtime_tasks.done() # the cancelled gather was awaited without raising
assert agent.close_calls == 1 assert agent.close_calls == 1
assert provider.close_calls == 1
+1 -2
View File
@@ -32,8 +32,7 @@ AUTH_HEADERS = {"Authorization": f"Bearer {API_KEY}"}
def _make_mock_agent(response_text: str = "mock response") -> MagicMock: def _make_mock_agent(response_text: str = "mock response") -> MagicMock:
agent = MagicMock() agent = MagicMock()
agent.process_direct = AsyncMock(return_value=response_text) agent.process_direct = AsyncMock(return_value=response_text)
agent._connect_mcp = AsyncMock() agent.aclose = AsyncMock()
agent.close_mcp = AsyncMock()
agent._last_usage = {} agent._last_usage = {}
return agent return agent
+8 -16
View File
@@ -64,8 +64,7 @@ def test_sse_done_format() -> None:
def _make_streaming_agent(tokens: list[str]) -> MagicMock: def _make_streaming_agent(tokens: list[str]) -> MagicMock:
"""Create a mock agent that streams tokens via on_stream callback.""" """Create a mock agent that streams tokens via on_stream callback."""
agent = MagicMock() agent = MagicMock()
agent._connect_mcp = AsyncMock() agent.aclose = AsyncMock()
agent.close_mcp = AsyncMock()
async def fake_process_direct(*, content="", media=None, session_key="", async def fake_process_direct(*, content="", media=None, session_key="",
channel="", chat_id="", on_stream=None, channel="", chat_id="", on_stream=None,
@@ -136,8 +135,7 @@ async def test_stream_false_returns_json(aiohttp_client) -> None:
"""stream=false should still return regular JSON response.""" """stream=false should still return regular JSON response."""
agent = MagicMock() agent = MagicMock()
agent.process_direct = AsyncMock(return_value="normal reply") agent.process_direct = AsyncMock(return_value="normal reply")
agent._connect_mcp = AsyncMock() agent.aclose = AsyncMock()
agent.close_mcp = AsyncMock()
agent._last_usage = {} agent._last_usage = {}
app = create_app(agent, model_name="m", api_key=API_KEY) app = create_app(agent, model_name="m", api_key=API_KEY)
@@ -160,8 +158,7 @@ async def test_stream_default_is_false(aiohttp_client) -> None:
"""Omitting stream should behave like stream=false.""" """Omitting stream should behave like stream=false."""
agent = MagicMock() agent = MagicMock()
agent.process_direct = AsyncMock(return_value="default reply") agent.process_direct = AsyncMock(return_value="default reply")
agent._connect_mcp = AsyncMock() agent.aclose = AsyncMock()
agent.close_mcp = AsyncMock()
agent._last_usage = {} agent._last_usage = {}
app = create_app(agent, model_name="m", api_key=API_KEY) app = create_app(agent, model_name="m", api_key=API_KEY)
@@ -217,8 +214,7 @@ async def test_stream_passes_on_stream_callbacks(aiohttp_client) -> None:
agent = MagicMock() agent = MagicMock()
agent.process_direct = fake_process_direct agent.process_direct = fake_process_direct
agent._connect_mcp = AsyncMock() agent.aclose = AsyncMock()
agent.close_mcp = AsyncMock()
agent._last_usage = {} agent._last_usage = {}
app = create_app(agent, model_name="m", api_key=API_KEY) app = create_app(agent, model_name="m", api_key=API_KEY)
@@ -251,8 +247,7 @@ async def test_stream_segment_end_does_not_close_sse(aiohttp_client) -> None:
return "planning final" return "planning final"
agent.process_direct = fake_process_direct agent.process_direct = fake_process_direct
agent._connect_mcp = AsyncMock() agent.aclose = AsyncMock()
agent.close_mcp = AsyncMock()
agent._last_usage = {} agent._last_usage = {}
app = create_app(agent, model_name="m", api_key=API_KEY) app = create_app(agent, model_name="m", api_key=API_KEY)
@@ -291,8 +286,7 @@ async def test_stream_uses_final_response_when_no_deltas(aiohttp_client) -> None
return "plain final" return "plain final"
agent.process_direct = fake_process_direct agent.process_direct = fake_process_direct
agent._connect_mcp = AsyncMock() agent.aclose = AsyncMock()
agent.close_mcp = AsyncMock()
agent._last_usage = {} agent._last_usage = {}
app = create_app(agent, model_name="m", api_key=API_KEY) app = create_app(agent, model_name="m", api_key=API_KEY)
@@ -334,8 +328,7 @@ async def test_stream_with_session_id(aiohttp_client) -> None:
agent = MagicMock() agent = MagicMock()
agent.process_direct = fake_process_direct agent.process_direct = fake_process_direct
agent._connect_mcp = AsyncMock() agent.aclose = AsyncMock()
agent.close_mcp = AsyncMock()
agent._last_usage = {} agent._last_usage = {}
app = create_app(agent, model_name="m", api_key=API_KEY) app = create_app(agent, model_name="m", api_key=API_KEY)
@@ -364,8 +357,7 @@ async def test_streaming_backend_failure_does_not_emit_success_terminator(aiohtt
raise RuntimeError("backend blew up") raise RuntimeError("backend blew up")
agent.process_direct = boom agent.process_direct = boom
agent._connect_mcp = AsyncMock() agent.aclose = AsyncMock()
agent.close_mcp = AsyncMock()
agent._last_usage = {} agent._last_usage = {}
app = create_app(agent, model_name="m", api_key=API_KEY) app = create_app(agent, model_name="m", api_key=API_KEY)
+23 -7
View File
@@ -101,6 +101,19 @@ def test_from_config_creates_instance(tmp_path):
assert bot._loop.workspace == tmp_path assert bot._loop.workspace == tmp_path
def test_from_config_composes_configured_mcp_outside_agent_loop(tmp_path):
config_path = _write_config(
tmp_path,
{"tools": {"mcpServers": {"demo": {"command": "fake-mcp"}}}},
)
bot = Nanobot.from_config(config_path, workspace=tmp_path)
assert bot._mcp_provider is not None
assert bot._mcp_provider.configured_server_names == {"demo"}
assert bot._mcp_provider._registry is bot._loop.tools
def test_from_config_accepts_default_model_override(tmp_path): def test_from_config_accepts_default_model_override(tmp_path):
config_path = _write_config(tmp_path) config_path = _write_config(tmp_path)
@@ -1637,37 +1650,40 @@ async def test_runtime_helpers_expose_model_workspace_and_compact(tmp_path):
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_aclose_delegates_to_loop_close_mcp(tmp_path): async def test_aclose_releases_loop_and_mcp_provider(tmp_path):
config_path = _write_config(tmp_path) config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path)
bot._loop.close_mcp = AsyncMock() bot._loop.aclose = AsyncMock()
assert bot._mcp_provider is not None
bot._mcp_provider.aclose = AsyncMock()
await bot.aclose() await bot.aclose()
bot._loop.close_mcp.assert_awaited_once() bot._loop.aclose.assert_awaited_once()
bot._mcp_provider.aclose.assert_awaited_once()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_context_manager_calls_aclose_on_exit(tmp_path): async def test_context_manager_calls_aclose_on_exit(tmp_path):
config_path = _write_config(tmp_path) config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path)
bot._loop.close_mcp = AsyncMock() bot._loop.aclose = AsyncMock()
async with bot as b: async with bot as b:
assert b is bot assert b is bot
bot._loop.close_mcp.assert_awaited_once() bot._loop.aclose.assert_awaited_once()
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_context_manager_does_not_swallow_exceptions(tmp_path): async def test_context_manager_does_not_swallow_exceptions(tmp_path):
config_path = _write_config(tmp_path) config_path = _write_config(tmp_path)
bot = Nanobot.from_config(config_path, workspace=tmp_path) bot = Nanobot.from_config(config_path, workspace=tmp_path)
bot._loop.close_mcp = AsyncMock() bot._loop.aclose = AsyncMock()
with pytest.raises(ValueError): with pytest.raises(ValueError):
async with bot as b: async with bot as b:
assert b is bot assert b is bot
raise ValueError("boom") raise ValueError("boom")
bot._loop.close_mcp.assert_awaited_once() bot._loop.aclose.assert_awaited_once()
+57 -9
View File
@@ -34,8 +34,7 @@ AUTH_HEADERS = {"Authorization": f"Bearer {API_KEY}"}
def _make_mock_agent(response_text: str = "mock response") -> MagicMock: def _make_mock_agent(response_text: str = "mock response") -> MagicMock:
agent = MagicMock() agent = MagicMock()
agent.process_direct = AsyncMock(return_value=response_text) agent.process_direct = AsyncMock(return_value=response_text)
agent._connect_mcp = AsyncMock() agent.aclose = AsyncMock()
agent.close_mcp = AsyncMock()
agent._last_usage = {"prompt_tokens": 100, "completion_tokens": 50} agent._last_usage = {"prompt_tokens": 100, "completion_tokens": 50}
return agent return agent
@@ -149,6 +148,59 @@ async def test_api_routes_allow_requests_without_configured_api_key(aiohttp_clie
mock_agent.process_direct.assert_called_once() mock_agent.process_direct.assert_called_once()
@pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed")
@pytest.mark.asyncio
async def test_api_prepares_application_resources_before_each_turn(aiohttp_client) -> None:
events: list[str] = []
agent = _make_mock_agent()
async def prepare_agent() -> None:
events.append("prepare")
async def process_direct(**_kwargs):
events.append("process")
return "ready"
agent.process_direct = process_direct
app = create_app(agent, prepare_agent=prepare_agent)
client = await aiohttp_client(app)
response = await client.post(
"/v1/chat/completions",
json={"messages": [{"role": "user", "content": "hello"}]},
)
assert response.status == 200
assert events == ["prepare", "process"]
@pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed")
@pytest.mark.asyncio
async def test_api_preparation_is_bounded_by_request_timeout(aiohttp_client) -> None:
agent = _make_mock_agent()
started = asyncio.Event()
async def prepare_agent() -> None:
started.set()
await asyncio.Event().wait()
app = create_app(
agent,
request_timeout=0.01,
prepare_agent=prepare_agent,
)
client = await aiohttp_client(app)
response = await client.post(
"/v1/chat/completions",
json={"messages": [{"role": "user", "content": "hello"}]},
)
assert started.is_set()
assert response.status == 504
agent.process_direct.assert_not_awaited()
@pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed") @pytest.mark.skipif(not HAS_AIOHTTP, reason="aiohttp not installed")
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_no_user_message_returns_400(aiohttp_client, app) -> None: async def test_no_user_message_returns_400(aiohttp_client, app) -> None:
@@ -275,8 +327,7 @@ async def test_followup_requests_share_same_session_key(aiohttp_client) -> None:
agent = MagicMock() agent = MagicMock()
agent.process_direct = fake_process agent.process_direct = fake_process
agent._connect_mcp = AsyncMock() agent.aclose = AsyncMock()
agent.close_mcp = AsyncMock()
agent._last_usage = {} agent._last_usage = {}
app = create_app(agent, model_name="m", api_key=API_KEY) app = create_app(agent, model_name="m", api_key=API_KEY)
@@ -315,8 +366,7 @@ async def test_fixed_session_requests_are_serialized(aiohttp_client) -> None:
agent = MagicMock() agent = MagicMock()
agent.process_direct = slow_process agent.process_direct = slow_process
agent._connect_mcp = AsyncMock() agent.aclose = AsyncMock()
agent.close_mcp = AsyncMock()
agent._last_usage = {} agent._last_usage = {}
app = create_app(agent, model_name="m", api_key=API_KEY) app = create_app(agent, model_name="m", api_key=API_KEY)
@@ -433,8 +483,7 @@ async def test_empty_response_falls_back_without_retry(aiohttp_client) -> None:
agent = MagicMock() agent = MagicMock()
agent.process_direct = always_empty agent.process_direct = always_empty
agent._connect_mcp = AsyncMock() agent.aclose = AsyncMock()
agent.close_mcp = AsyncMock()
agent._last_usage = {} agent._last_usage = {}
app = create_app(agent, model_name="m", api_key=API_KEY) app = create_app(agent, model_name="m", api_key=API_KEY)
@@ -457,7 +506,6 @@ async def test_process_direct_accepts_media() -> None:
from nanobot.bus.runtime_events import RuntimeEventPublisher from nanobot.bus.runtime_events import RuntimeEventPublisher
loop = AgentLoop.__new__(AgentLoop) loop = AgentLoop.__new__(AgentLoop)
loop._connect_mcp = AsyncMock()
loop._session_locks = {} loop._session_locks = {}
loop.runtime_event_publisher = RuntimeEventPublisher() loop.runtime_event_publisher = RuntimeEventPublisher()
+7 -17
View File
@@ -12,7 +12,6 @@ from unittest.mock import AsyncMock
import pytest import pytest
from nanobot.agent import context as agent_context
from nanobot.agent.loop import AgentLoop from nanobot.agent.loop import AgentLoop
from nanobot.agent.tools.context import RequestContext, bind_request_context, reset_request_context from nanobot.agent.tools.context import RequestContext, bind_request_context, reset_request_context
from nanobot.agent.tools.exec_session import ( from nanobot.agent.tools.exec_session import (
@@ -712,7 +711,7 @@ def test_exec_session_manager_preserves_single_cleanup_error():
asyncio.run(run()) asyncio.run(run())
def test_agent_loop_shutdown_closes_exec_sessions(tmp_path, monkeypatch): def test_agent_loop_shutdown_closes_exec_sessions(tmp_path):
async def run() -> None: async def run() -> None:
manager = ExecSessionManager() manager = ExecSessionManager()
tool = ExecTool(working_dir=str(tmp_path), timeout=30, session_manager=manager) tool = ExecTool(working_dir=str(tmp_path), timeout=30, session_manager=manager)
@@ -723,14 +722,13 @@ def test_agent_loop_shutdown_closes_exec_sessions(tmp_path, monkeypatch):
sid = _session_id(initial) sid = _session_id(initial)
process = manager._sessions[sid].process process = manager._sessions[sid].process
monkeypatch.setattr(agent_context, "close_mcp", lambda _state: asyncio.sleep(0))
loop = object.__new__(AgentLoop) loop = object.__new__(AgentLoop)
loop._background_tasks = set() loop._background_tasks = set()
loop._exec_session_manager = manager loop._exec_session_manager = manager
loop.subagents = SimpleNamespace(close=AsyncMock()) loop.subagents = SimpleNamespace(close=AsyncMock())
await loop.close_mcp() await loop.aclose()
await loop.close_mcp() await loop.aclose()
assert process.returncode is not None assert process.returncode is not None
assert manager._sessions == {} assert manager._sessions == {}
@@ -739,7 +737,7 @@ def test_agent_loop_shutdown_closes_exec_sessions(tmp_path, monkeypatch):
asyncio.run(run()) asyncio.run(run())
def test_agent_loop_shutdown_attempts_all_cleanup_after_errors(monkeypatch): def test_agent_loop_shutdown_attempts_all_cleanup_after_errors():
async def run() -> None: async def run() -> None:
loop = object.__new__(AgentLoop) loop = object.__new__(AgentLoop)
loop._background_tasks = set() loop._background_tasks = set()
@@ -749,16 +747,12 @@ def test_agent_loop_shutdown_attempts_all_cleanup_after_errors(monkeypatch):
loop._exec_session_manager = SimpleNamespace( loop._exec_session_manager = SimpleNamespace(
close_all=AsyncMock(side_effect=OSError("exec cleanup failed")), close_all=AsyncMock(side_effect=OSError("exec cleanup failed")),
) )
close_mcp = AsyncMock()
monkeypatch.setattr(agent_context, "close_mcp", close_mcp)
with pytest.raises(BaseExceptionGroup) as exc_info: with pytest.raises(BaseExceptionGroup) as exc_info:
await loop.close_mcp() await loop.aclose()
assert len(exc_info.value.exceptions) == 2 assert len(exc_info.value.exceptions) == 2
loop.subagents.close.assert_awaited_once() loop.subagents.close.assert_awaited_once()
loop._exec_session_manager.close_all.assert_awaited_once() loop._exec_session_manager.close_all.assert_awaited_once()
close_mcp.assert_awaited_once_with(loop)
asyncio.run(run()) asyncio.run(run())
@@ -892,7 +886,7 @@ def test_terminate_by_owner_skips_sessions_without_owner_key(tmp_path):
asyncio.run(run()) asyncio.run(run())
def test_agent_loop_shutdown_preserves_single_cleanup_error(monkeypatch): def test_agent_loop_shutdown_preserves_single_cleanup_error():
async def run() -> None: async def run() -> None:
loop = object.__new__(AgentLoop) loop = object.__new__(AgentLoop)
loop._background_tasks = set() loop._background_tasks = set()
@@ -900,13 +894,9 @@ def test_agent_loop_shutdown_preserves_single_cleanup_error(monkeypatch):
close=AsyncMock(side_effect=RuntimeError("subagent cleanup failed")), close=AsyncMock(side_effect=RuntimeError("subagent cleanup failed")),
) )
loop._exec_session_manager = SimpleNamespace(close_all=AsyncMock()) loop._exec_session_manager = SimpleNamespace(close_all=AsyncMock())
close_mcp = AsyncMock()
monkeypatch.setattr(agent_context, "close_mcp", close_mcp)
with pytest.raises(RuntimeError, match="subagent cleanup failed"): with pytest.raises(RuntimeError, match="subagent cleanup failed"):
await loop.close_mcp() await loop.aclose()
loop._exec_session_manager.close_all.assert_awaited_once() loop._exec_session_manager.close_all.assert_awaited_once()
close_mcp.assert_awaited_once_with(loop)
asyncio.run(run()) asyncio.run(run())
+76 -26
View File
@@ -13,6 +13,7 @@ import pytest
import nanobot.agent.tools.mcp as mcp_mod import nanobot.agent.tools.mcp as mcp_mod
from nanobot.agent.tools.mcp import ( from nanobot.agent.tools.mcp import (
MCPPromptWrapper, MCPPromptWrapper,
MCPProvider,
MCPResourceWrapper, MCPResourceWrapper,
MCPToolWrapper, MCPToolWrapper,
_normalize_windows_stdio_command, _normalize_windows_stdio_command,
@@ -153,7 +154,7 @@ def _make_wrapper(session: object, *, timeout: float = 0.1) -> MCPToolWrapper:
@pytest.mark.asyncio @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() started = asyncio.Event()
async def connect_mcp_servers(_servers: dict, _registry: ToolRegistry) -> dict: 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) await asyncio.sleep(60)
return {} return {}
class State: provider = MCPProvider(
pass {"test": MCPServerConfig(command="fake")},
ToolRegistry(),
state = State() )
state._mcp_closing = False
state._mcp_servers = {"test": MCPServerConfig(command="fake")}
state._mcp_stacks = {}
state._mcp_connecting = False
monkeypatch.setattr(mcp_mod, "connect_mcp_servers", connect_mcp_servers) 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) await asyncio.wait_for(started.wait(), timeout=1.0)
task.cancel() task.cancel()
with pytest.raises(asyncio.CancelledError): with pytest.raises(asyncio.CancelledError):
await task await task
assert state._mcp_connecting is False assert provider.connected_server_names == set()
assert provider.runtime_status() == {"test": "failed"}
@pytest.mark.asyncio @pytest.mark.asyncio
@@ -223,25 +221,20 @@ async def test_saved_oauth_http_403_projects_failed_runtime_without_details(
rejected_streamable_http, rejected_streamable_http,
) )
class State: provider = MCPProvider(
pass {
"xmind": MCPServerConfig(
state = State()
state._mcp_closing = False
state._mcp_servers = {
"xmind": MCPServerConfig(
type="streamableHttp", type="streamableHttp",
auth="oauth", auth="oauth",
url="https://app.xmind.com/api/mcp", url="https://app.xmind.com/api/mcp",
) )
} },
state._mcp_stacks = {} ToolRegistry(),
state._mcp_runtime_statuses = {} )
state._mcp_connecting = False
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 snapshot == {"xmind": "failed"}
assert "saved-oauth-secret" not in str(snapshot) assert "saved-oauth-secret" not in str(snapshot)
assert "app.xmind.com" 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) 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 @pytest.mark.asyncio
async def test_connect_mcp_servers_streamable_http_uses_finite_timeout( async def test_connect_mcp_servers_streamable_http_uses_finite_timeout(
fake_mcp_runtime: dict[str, object | None], 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 len(wrapper.name) == 64
assert not wrapper.name.startswith(mcp_mod._tool_prefix(server_name)) 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 wrapper._reconnect is not None
assert other_wrapper._reconnect is None assert other_wrapper._reconnect is None
+31 -1
View File
@@ -1,7 +1,8 @@
from __future__ import annotations from __future__ import annotations
import asyncio
import json import json
from collections.abc import Callable, Mapping from collections.abc import Awaitable, Callable, Mapping
from pathlib import Path from pathlib import Path
from types import SimpleNamespace from types import SimpleNamespace
from unittest.mock import ANY, AsyncMock, MagicMock from unittest.mock import ANY, AsyncMock, MagicMock
@@ -22,6 +23,7 @@ def _router(
authorized: bool = True, authorized: bool = True,
config_path: Path | None = None, config_path: Path | None = None,
mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None, mcp_runtime_status: Callable[[], Mapping[str, str]] | None = None,
mcp_reload: Callable[[], Awaitable[dict[str, object]]] | None = None,
) -> WebUISettingsRouter: ) -> WebUISettingsRouter:
return WebUISettingsRouter( return WebUISettingsRouter(
settings=WebUISettingsServices.create(config_path or get_config_path()), settings=WebUISettingsServices.create(config_path or get_config_path()),
@@ -37,6 +39,7 @@ def _router(
runtime_surface="browser", runtime_surface="browser",
runtime_capabilities={}, runtime_capabilities={},
mcp_runtime_status=mcp_runtime_status, mcp_runtime_status=mcp_runtime_status,
mcp_reload=mcp_reload,
mcp_oauth_redirect_uri=lambda _request: "https://gateway.example/auth/mcp/callback", mcp_oauth_redirect_uri=lambda _request: "https://gateway.example/auth/mcp/callback",
) )
@@ -89,6 +92,33 @@ async def test_mcp_list_serializes_local_runtime_failure_snapshot(tmp_path) -> N
assert snapshot_calls == 1 assert snapshot_calls == 1
@pytest.mark.asyncio
async def test_mcp_reload_callback_is_bounded(
monkeypatch: pytest.MonkeyPatch,
) -> None:
started = asyncio.Event()
async def reload_mcp() -> dict[str, object]:
started.set()
await asyncio.Event().wait()
return {"ok": True}
monkeypatch.setattr(
"nanobot.webui.settings_routes._MCP_RELOAD_TIMEOUT_SECONDS",
0.01,
)
router = _router(mcp_reload=reload_mcp)
result = await router._reload_mcp_runtime()
assert started.is_set()
assert result == {
"ok": False,
"message": "MCP hot reload timed out. Restart nanobot to pick up changes.",
"requires_restart": True,
}
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_mcp_oauth_start_uses_gateway_callback_and_requires_api_auth(monkeypatch) -> None: async def test_mcp_oauth_start_uses_gateway_callback_and_requires_api_auth(monkeypatch) -> None:
config = SimpleNamespace( config = SimpleNamespace(