mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-31 08:13:11 +03:00
fix(gateway): retry MCP readiness before turns (#5535)
This commit is contained in:
@@ -12,6 +12,7 @@ from loguru import logger
|
|||||||
from rich.console import Console
|
from rich.console import Console
|
||||||
|
|
||||||
from nanobot import __logo__, __version__
|
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.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.mcp import MCPProvider
|
||||||
@@ -46,6 +47,17 @@ __all__ = ["_run_gateway"]
|
|||||||
console = Console()
|
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:
|
def _http_endpoint_responding(url: str, *, timeout_s: float = 0.25) -> bool:
|
||||||
"""Return whether an HTTP endpoint responds, including with an auth error."""
|
"""Return whether an HTTP endpoint responds, including with an auth error."""
|
||||||
import urllib.error
|
import urllib.error
|
||||||
@@ -445,6 +457,7 @@ def _run_gateway(
|
|||||||
turn_delivery_factory=turn_delivery_factory,
|
turn_delivery_factory=turn_delivery_factory,
|
||||||
provider_signature=provider_snapshot.signature,
|
provider_signature=provider_snapshot.signature,
|
||||||
local_trigger_store=trigger_store,
|
local_trigger_store=trigger_store,
|
||||||
|
hooks=[_MCPReadinessHook(mcp_provider)],
|
||||||
hook_factories=[create_file_edit_activity_hook],
|
hook_factories=[create_file_edit_activity_hook],
|
||||||
tool_registry=tools,
|
tool_registry=tools,
|
||||||
recovery_admission=recovery,
|
recovery_admission=recovery,
|
||||||
|
|||||||
@@ -3697,6 +3697,7 @@ def test_gateway_agent_task_owns_initial_mcp_provider_close(
|
|||||||
return cls(**extra)
|
return cls(**extra)
|
||||||
|
|
||||||
def __init__(self, **_kwargs) -> None:
|
def __init__(self, **_kwargs) -> None:
|
||||||
|
seen["hooks"] = _kwargs.get("hooks")
|
||||||
self.model = "test-model"
|
self.model = "test-model"
|
||||||
self.provider = object()
|
self.provider = object()
|
||||||
self.sessions = _FakeSessionManager()
|
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.connect_task is seen["agent_task"]
|
||||||
assert mcp_provider.close_tasks[0] is mcp_provider.connect_task
|
assert mcp_provider.close_tasks[0] is mcp_provider.connect_task
|
||||||
assert len(mcp_provider.close_tasks) == 2
|
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(
|
def test_gateway_shutdown_event_exits_forever_runtime_tasks(
|
||||||
|
|||||||
@@ -11,7 +11,10 @@ import asyncio
|
|||||||
import time
|
import time
|
||||||
from contextlib import suppress
|
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:
|
class _FakeAgent:
|
||||||
@@ -53,6 +56,24 @@ class _FakeMCPProvider:
|
|||||||
self.events.append("mcp_closed")
|
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:
|
async def _cancellable_task(events: list[str]) -> None:
|
||||||
try:
|
try:
|
||||||
await asyncio.sleep(3600)
|
await asyncio.sleep(3600)
|
||||||
|
|||||||
Reference in New Issue
Block a user