Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 3 additions & 8 deletions nanobot/agent/context.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,21 +31,16 @@ def runtime_lines(state: Any, msg: Any, workspace: Path, *, skip: bool = False)
"""Return model-visible runtime annotations for turn-attached capabilities."""
return [
*cli_app_utils.runtime_lines(msg, workspace, skip=skip),
*mcp_tools.runtime_lines(
msg,
configured_server_names=set(state._mcp_servers),
connected_server_names=set(state._mcp_stacks),
skip=skip,
),
*state.mcp_provider.runtime_lines(msg, skip=skip),
]


async def connect_mcp(state: Any, tools: ToolRegistry) -> None:
await mcp_tools.connect_missing_servers(state, tools)
await state.mcp_provider.connect(tools)


async def handle_runtime_control(state: Any, msg: InboundMessage, tools: ToolRegistry) -> bool:
return await mcp_tools.handle_runtime_control(state, msg, tools)
return await state.mcp_provider.handle_runtime_control(msg, tools)


class ContextBuilder:
Expand Down
14 changes: 4 additions & 10 deletions nanobot/agent/loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
import os
import time
from collections.abc import Mapping
from contextlib import AsyncExitStack, nullcontext, suppress
from contextlib import nullcontext, suppress
from dataclasses import dataclass, field
from enum import Enum, auto
from functools import partial
Expand All @@ -29,6 +29,7 @@
from nanobot.agent.subagent import SubagentManager
from nanobot.agent.tools.context import RequestContext, bind_request_context, reset_request_context
from nanobot.agent.tools.file_state import FileStateStore, bind_file_states, reset_file_states
from nanobot.agent.tools.mcp import MCPToolProvider
from nanobot.agent.tools.message import MessageTool
from nanobot.agent.tools.registry import ToolRegistry
from nanobot.agent.tools.self import MyTool
Expand Down Expand Up @@ -359,9 +360,7 @@ def __init__(
)
self._unified_session = unified_session
self._running = False
self._mcp_servers = mcp_servers or {}
self._mcp_stacks: dict[str, AsyncExitStack] = {}
self._mcp_connecting = False
self.mcp_provider = MCPToolProvider(mcp_servers)
self._active_tasks: dict[str, list[asyncio.Task]] = {} # session_key -> tasks
self._background_tasks: list[asyncio.Task] = []
self._session_locks: dict[str, asyncio.Lock] = {}
Expand Down Expand Up @@ -1177,12 +1176,7 @@ async def close_mcp(self) -> None:
if self._background_tasks:
await asyncio.gather(*self._background_tasks, return_exceptions=True)
self._background_tasks.clear()
for name, stack in self._mcp_stacks.items():
try:
await stack.aclose()
except (RuntimeError, BaseExceptionGroup):
logger.debug("MCP server '{}' cleanup error (can be ignored)", name)
self._mcp_stacks.clear()
await self.mcp_provider.close()

def _schedule_background(self, coro) -> None:
"""Schedule a coroutine as a tracked background task (drained on shutdown)."""
Expand Down
46 changes: 46 additions & 0 deletions nanobot/agent/tools/mcp.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,6 +54,52 @@
_ReconnectCallback = Callable[[str, str, Tool], Awaitable[Tool | None]]


class MCPToolProvider:
"""Own the MCP dynamic tool provider lifecycle and runtime state."""

def __init__(self, servers: Mapping[str, Any] | None = None) -> None:
self._mcp_servers = dict(servers or {})
self._mcp_stacks: dict[str, AsyncExitStack] = {}
self._mcp_connecting = False

@property
def configured_server_names(self) -> set[str]:
return set(self._mcp_servers)

@property
def connected_server_names(self) -> set[str]:
return set(self._mcp_stacks)

async def connect(self, registry: ToolRegistry) -> None:
await connect_missing_servers(self, registry)

async def close(self) -> None:
for name, stack in self._mcp_stacks.items():
try:
await stack.aclose()
except (RuntimeError, BaseExceptionGroup):
logger.debug("MCP server '{}' cleanup error (can be ignored)", name)
self._mcp_stacks.clear()

async def reload(self, registry: ToolRegistry) -> dict[str, Any]:
return await reload_servers(self, registry)

async def handle_runtime_control(
self,
msg: InboundMessage,
registry: ToolRegistry,
) -> bool:
return await handle_runtime_control(self, msg, registry)

def runtime_lines(self, message: Any, *, skip: bool = False) -> list[str]:
return runtime_lines(
message,
configured_server_names=self.configured_server_names,
connected_server_names=self.connected_server_names,
skip=skip,
)


def _is_malformed_mcp_progress_notification(message: Any) -> bool:
payload = _mcp_jsonrpc_payload(message)
if _payload_value(payload, "method") != "notifications/progress":
Expand Down
2 changes: 1 addition & 1 deletion nanobot/agent/tools/self.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,7 +64,7 @@ def enabled(cls, ctx: Any) -> bool:
"runner", "sessions", "consolidator",
"dream", "auto_compact", "context", "commands",
# Sensitive runtime state (credentials, message routing, task tracking)
"_mcp_servers", "_mcp_stacks", "_pending_queues",
"_mcp_servers", "_mcp_stacks", "mcp_provider", "_pending_queues",
"_session_locks", "_active_tasks", "_background_tasks",
# Security boundaries (inspect + modify both blocked)
"restrict_to_workspace", "channels_config",
Expand Down
18 changes: 9 additions & 9 deletions tests/agent/test_mcp_connection.py
Original file line number Diff line number Diff line change
Expand Up @@ -131,7 +131,7 @@ async def _fake_connect(_servers, _registry):
await loop._connect_mcp()

assert attempts == 2
assert loop._mcp_stacks == {}
assert loop.mcp_provider._mcp_stacks == {}


@pytest.mark.asyncio
Expand Down Expand Up @@ -168,7 +168,7 @@ async def _fake_connect(servers, _registry):

assert owner_tasks
assert closed_tasks == owner_tasks
assert loop._mcp_stacks == {}
assert loop.mcp_provider._mcp_stacks == {}


@pytest.mark.asyncio
Expand Down Expand Up @@ -203,23 +203,23 @@ async def _fake_connect(servers, registry):
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
loop = _make_loop(tmp_path, mcp_servers={})

added = await mcp_runtime.reload_servers(loop, loop.tools)
added = await loop.mcp_provider.reload(loop.tools)

assert added["ok"] is True
assert added["added"] == ["browserbase"]
assert loop.tools.has("mcp_browserbase_navigate")
assert "browserbase" in loop._mcp_stacks
assert "browserbase" in loop.mcp_provider._mcp_stacks

config = load_config()
del config.tools.mcp_servers["browserbase"]
save_config(config)

removed = await mcp_runtime.reload_servers(loop, loop.tools)
removed = await loop.mcp_provider.reload(loop.tools)

assert removed["ok"] is True
assert removed["removed"] == ["browserbase"]
assert not loop.tools.has("mcp_browserbase_navigate")
assert "browserbase" not in loop._mcp_stacks
assert "browserbase" not in loop.mcp_provider._mcp_stacks
assert closed == ["browserbase"]


Expand Down Expand Up @@ -257,7 +257,7 @@ async def _fake_connect(servers, registry):

async def _handle_one_runtime_control() -> None:
msg = await loop.bus.consume_inbound()
handled = await mcp_runtime.handle_runtime_control(loop, msg, loop.tools)
handled = await loop.mcp_provider.handle_runtime_control(msg, loop.tools)
assert handled is True

consumer = asyncio.create_task(_handle_one_runtime_control())
Expand Down Expand Up @@ -310,7 +310,7 @@ async def _fake_connect(servers, registry):
monkeypatch.setattr("nanobot.agent.tools.mcp.connect_mcp_servers", _fake_connect)
loop = _make_loop(tmp_path, mcp_servers={"browserbase": config.tools.mcp_servers["browserbase"]})

result = await mcp_runtime.reload_servers(loop, loop.tools)
result = await loop.mcp_provider.reload(loop.tools)

assert result["ok"] is True
assert result["added"] == []
Expand Down Expand Up @@ -379,7 +379,7 @@ async def _fake_connect(servers, registry):
assert closed == ["remote"]
assert sessions[0].call_count == 1
assert sessions[1].call_count == 1
assert "remote" in loop._mcp_stacks
assert "remote" in loop.mcp_provider._mcp_stacks
assert loop.tools.get("mcp_remote_quote") is not old_tool


Expand Down