diff --git a/nanobot/cli/gateway_runtime.py b/nanobot/cli/gateway_runtime.py index 8d5d0f3a0..f2315ab08 100644 --- a/nanobot/cli/gateway_runtime.py +++ b/nanobot/cli/gateway_runtime.py @@ -12,6 +12,7 @@ from loguru import logger from rich.console import Console from nanobot import __logo__, __version__ +from nanobot.agent.hook import AgentHook, AgentRunHookContext from nanobot.agent.hooks import create_file_edit_activity_hook from nanobot.agent.loop import AgentLoop from nanobot.agent.tools.mcp import MCPProvider @@ -46,6 +47,17 @@ __all__ = ["_run_gateway"] console = Console() +class _MCPReadinessHook(AgentHook): + """Retry application-owned MCP connections before the runner reads tools.""" + + def __init__(self, provider: MCPProvider) -> None: + super().__init__() + self._provider = provider + + async def before_run(self, context: AgentRunHookContext) -> None: + await self._provider.connect() + + def _http_endpoint_responding(url: str, *, timeout_s: float = 0.25) -> bool: """Return whether an HTTP endpoint responds, including with an auth error.""" import urllib.error @@ -445,6 +457,7 @@ def _run_gateway( turn_delivery_factory=turn_delivery_factory, provider_signature=provider_snapshot.signature, local_trigger_store=trigger_store, + hooks=[_MCPReadinessHook(mcp_provider)], hook_factories=[create_file_edit_activity_hook], tool_registry=tools, recovery_admission=recovery, diff --git a/tests/cli/test_commands.py b/tests/cli/test_commands.py index 4f1f4fb14..ea565b7bf 100644 --- a/tests/cli/test_commands.py +++ b/tests/cli/test_commands.py @@ -3697,6 +3697,7 @@ def test_gateway_agent_task_owns_initial_mcp_provider_close( return cls(**extra) def __init__(self, **_kwargs) -> None: + seen["hooks"] = _kwargs.get("hooks") self.model = "test-model" self.provider = object() self.sessions = _FakeSessionManager() @@ -3809,6 +3810,11 @@ def test_gateway_agent_task_owns_initial_mcp_provider_close( 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 + hooks = seen["hooks"] + assert isinstance(hooks, list) + assert len(hooks) == 1 + hook = hooks[0] + assert isinstance(hook, cli_gateway_runtime._MCPReadinessHook) def test_gateway_shutdown_event_exits_forever_runtime_tasks( diff --git a/tests/cli/test_gateway_runtime.py b/tests/cli/test_gateway_runtime.py index d5fb0c512..38adf624f 100644 --- a/tests/cli/test_gateway_runtime.py +++ b/tests/cli/test_gateway_runtime.py @@ -11,7 +11,10 @@ import asyncio import time from contextlib import suppress -from nanobot.cli.gateway_runtime import _close_gateway_runtime +from nanobot.agent.hook import AgentRunHookContext +from nanobot.agent.tools.mcp import MCPProvider +from nanobot.agent.tools.registry import ToolRegistry +from nanobot.cli.gateway_runtime import _close_gateway_runtime, _MCPReadinessHook class _FakeAgent: @@ -53,6 +56,24 @@ class _FakeMCPProvider: self.events.append("mcp_closed") +class _TrackingMCPProvider(MCPProvider): + def __init__(self) -> None: + super().__init__({}, ToolRegistry()) + self.connect_calls = 0 + + async def connect(self) -> None: + self.connect_calls += 1 + + +async def test_mcp_readiness_hook_delegates_to_application_provider() -> None: + provider = _TrackingMCPProvider() + hook = _MCPReadinessHook(provider) + + await hook.before_run(AgentRunHookContext(messages=[])) + + assert provider.connect_calls == 1 + + async def _cancellable_task(events: list[str]) -> None: try: await asyncio.sleep(3600)