diff --git a/tools/mcp_schema_cache.py b/tools/mcp_schema_cache.py index 558099e8d9..ba22dcb177 100644 --- a/tools/mcp_schema_cache.py +++ b/tools/mcp_schema_cache.py @@ -21,7 +21,6 @@ _cache_lock = threading.Lock() def _cache_path() -> Path: from hermes_constants import get_hermes_home - return get_hermes_home() / "cache" / _CACHE_FILENAME @@ -53,7 +52,6 @@ def _load_all() -> Dict[str, Any]: def _save_all(data: Dict[str, Any]) -> None: from utils import atomic_json_write - # 0o600: the cache file is trusted input on the lazy registration path, keep it user-only. atomic_json_write(_cache_path(), data, mode=0o600) diff --git a/tools/mcp_tool.py b/tools/mcp_tool.py index 7dc0c8ebf3..f29a01bcbd 100644 --- a/tools/mcp_tool.py +++ b/tools/mcp_tool.py @@ -113,16 +113,6 @@ ClientSession: Any = None _MCP_SDK_IMPORT_ATTEMPTED = False _MCP_SDK_IMPORT_LOCK = threading.Lock() -# SDK symbols bound by _ensure_mcp_sdk(). Module __getattr__ (PEP 562) imports the SDK on -# first external access, so mock.patch("tools.mcp_tool.stdio_client") sees a real original -# and the mock is never clobbered (_ensure is idempotent). -_MCP_SDK_LAZY_SYMBOLS = frozenset({ - "StdioServerParameters", "stdio_client", "streamablehttp_client", "streamable_http_client", - "CreateMessageResult", "CreateMessageResultWithTools", "ErrorData", "SamplingCapability", - "SamplingToolsCapability", "TextContent", "ToolUseContent", "ElicitRequestParams", - "ElicitResult", "ServerNotification", "ToolListChangedNotification", - "PromptListChangedNotification", "ResourceListChangedNotification"}) - # Optional SDK type families (module, names, debug message when absent) bound in this order to # _MCP_SAMPLING_TYPES / _MCP_ELICITATION_TYPES / _MCP_NOTIFICATION_TYPES. Each is gated # separately so an older SDK only loses that feature, not MCP. @@ -136,6 +126,12 @@ _OPTIONAL_TYPE_FAMILIES = ( "ResourceListChangedNotification"), "MCP notification types not available -- dynamic tool discovery disabled"), ) +# SDK symbols bound by _ensure_mcp_sdk(). Module __getattr__ (PEP 562) imports the SDK on +# first external access, so mock.patch("tools.mcp_tool.stdio_client") sees a real original +# and the mock is never clobbered (_ensure is idempotent). +_MCP_SDK_LAZY_SYMBOLS = frozenset( + {"StdioServerParameters", "stdio_client", "streamablehttp_client", "streamable_http_client"} + | {n for _mod, names, _msg in _OPTIONAL_TYPE_FAMILIES for n in names}) def __getattr__(name: str): @@ -173,7 +169,6 @@ def _ensure_mcp_sdk() -> bool: global _MCP_SAMPLING_TYPES, _MCP_NOTIFICATION_TYPES, _MCP_ELICITATION_TYPES, sse_client global _MCP_MESSAGE_HANDLER_SUPPORTED, _MCP_LOGGING_CALLBACK_SUPPORTED, LATEST_HANDSHAKE_VERSION global _JSONRPC_METHOD_NOT_FOUND - if not _MCP_AVAILABLE: return False if _MCP_SDK_IMPORT_ATTEMPTED or ClientSession is not None: @@ -201,13 +196,11 @@ def _ensure_mcp_sdk() -> bool: _import_sdk_names(*family) for family in _OPTIONAL_TYPE_FAMILIES] else: logger.debug("mcp package not installed -- MCP tool support disabled") - if _MCP_AVAILABLE: try: _JSONRPC_METHOD_NOT_FOUND = importlib.import_module("mcp.types").METHOD_NOT_FOUND except Exception: # pragma: no cover — SDK without the constant pass - _MCP_MESSAGE_HANDLER_SUPPORTED = _client_session_accepts("message_handler") if _MCP_AVAILABLE and not _MCP_MESSAGE_HANDLER_SUPPORTED: logger.debug("MCP SDK does not support message_handler -- dynamic tool discovery disabled") @@ -236,15 +229,13 @@ def sdk_httpx(): _SDK_HTTPX_MOD = getattr(_transport, "httpx2", None) or getattr(_transport, "httpx", None) except ImportError: _SDK_HTTPX_MOD = None - if _SDK_HTTPX_MOD is None: + for fallback in ("httpx2", "httpx"): + if _SDK_HTTPX_MOD is not None: + break try: - import httpx2 as _fallback + _SDK_HTTPX_MOD = importlib.import_module(fallback) except ImportError: - try: - import httpx as _fallback # type: ignore[no-redef] - except ImportError: - return None - _SDK_HTTPX_MOD = _fallback + pass return _SDK_HTTPX_MOD @@ -314,21 +305,16 @@ async def _paginate_full_list(list_method, items_attr: str, server_name: str, # mcp 2.0 takes params=PaginatedRequestParams, 1.x takes cursor=. try: import mcp.types as _types # late: keeps the SDK import lazy - _params_cls = getattr(_types, "PaginatedRequestParams", None) - if _params_cls is not None: - result = await list_method(params=_params_cls(cursor=cursor)) - else: - result = await list_method(cursor=cursor) + result = await (list_method(params=_params_cls(cursor=cursor)) if _params_cls is not None + else list_method(cursor=cursor)) except TypeError: result = await list_method(cursor=cursor) if cache_meta_out is not None and not items: - _ttl = mcp_field(result, "ttl_ms", "ttlMs") - _scope = mcp_field(result, "cache_scope", "cacheScope") - if _ttl is not None: - cache_meta_out["ttl_ms"] = _ttl - if _scope is not None: - cache_meta_out["cache_scope"] = _scope + for key, snake, camel in (("ttl_ms", "ttl_ms", "ttlMs"), ("cache_scope", "cache_scope", "cacheScope")): + hint = mcp_field(result, snake, camel) + if hint is not None: + cache_meta_out[key] = hint items.extend(getattr(result, items_attr, None) or []) cursor = mcp_field(result, "next_cursor", "nextCursor") # Cursor is an opaque string; anything else (incl. mocks) = last page. @@ -517,11 +503,9 @@ def _mcp_registry_scope() -> Optional[str]: """Registry scope for MCP registrations: under a profile multiplexer each profile's MCP tools live in its own registry overlay; single-profile processes stay global (None).""" from agent.secret_scope import is_multiplex_active - if not is_multiplex_active(): return None from tools.registry import registry - return registry.current_scope_key() diff --git a/tools/mcp_tool_common.py b/tools/mcp_tool_common.py index 1e2f38425c..e120ea4a24 100644 --- a/tools/mcp_tool_common.py +++ b/tools/mcp_tool_common.py @@ -21,7 +21,6 @@ class _OriginProxy: def __getattr__(self, name: str): from tools import mcp_tool - return getattr(mcp_tool, name) @@ -51,7 +50,6 @@ def _resolve_tool_timeout(config: dict) -> float: return per_server try: from agent.deadline import resolve_timeout - resolved = resolve_timeout("mcp.tool_call", default=_DEFAULT_TOOL_TIMEOUT) if resolved is not None: return resolved @@ -151,9 +149,8 @@ def _get_lifecycle_seconds(config: dict, key: str) -> Optional[float]: """Optional positive lifecycle timeout from top-level/nested ``lifecycle`` config (``0`` disables; negatives and non-numbers are warned about and ignored).""" raw = config.get(key) - lifecycle = config.get("lifecycle") - if raw is None and isinstance(lifecycle, dict): - raw = lifecycle.get(key) + if raw is None and isinstance(config.get("lifecycle"), dict): + raw = config["lifecycle"].get(key) if raw is None: return None try: diff --git a/tools/mcp_tool_discovery.py b/tools/mcp_tool_discovery.py index 2c0de1781f..65cdd321f0 100644 --- a/tools/mcp_tool_discovery.py +++ b/tools/mcp_tool_discovery.py @@ -72,9 +72,7 @@ async def _connect_server(name: str, config: dict) -> _core.MCPServerTask: def _request_lazy_reconnect(server_name: str, server: _core.MCPServerTask) -> bool: """Wake a recycled stdio server and wait briefly for a fresh session.""" - if not server._is_recycled_stdio(): - return False - loop = _core._running_loop() + loop = _core._running_loop() if server._is_recycled_stdio() else None if loop is None: return False @@ -122,6 +120,13 @@ def _note_connect_success(name: str) -> None: _core._clear_connect_failure(name) +def _adopt_server(name: str, server: _core.MCPServerTask) -> None: + """Publish *server* into ``_servers`` with its owning registry scope (under ``_lock``).""" + with _core._lock: + _core._servers[name] = server + _core._server_scope_keys[name] = _core._mcp_registry_scope() + + def _ensure_lazy_server_connected(server_name: str) -> bool: """Connect a lazily-registered server on demand (sync; blocks the caller). @@ -139,21 +144,15 @@ def _ensure_lazy_server_connected(server_name: str) -> bool: return False _core._server_connecting.add(server_name) _core._server_connect_errors.pop(server_name, None) - logger.info("MCP server '%s': lazy start on first use", server_name) _core._ensure_mcp_loop() connect_timeout = config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT) - - async def _connect(): - return await _core._discover_and_register_server(server_name, config) - try: - _core._run_on_mcp_loop(_connect, timeout=float(connect_timeout) + 30.0) + _core._run_on_mcp_loop(lambda: _core._discover_and_register_server(server_name, config), + timeout=float(connect_timeout) + 30.0) except BaseException as exc: - message = _note_connect_failure(server_name, exc) - logger.warning("Lazy MCP connect failed for '%s': %s", server_name, message) + logger.warning("Lazy MCP connect failed for '%s': %s", server_name, _note_connect_failure(server_name, exc)) return False - _note_connect_success(server_name) with _core._lock: _core._lazy_server_configs.pop(server_name, None) @@ -165,7 +164,6 @@ def _ensure_lazy_server_connected(server_name: str) -> bool: phantom_names = [n for n in cached_names if n not in live_names] if phantom_names: from tools.registry import registry - for tool_name in phantom_names: registry.deregister(tool_name, scope=_core._server_registry_scope(server_name)) _core._forget_mcp_tool_server(tool_name) @@ -184,25 +182,23 @@ def _get_connected_server_for_call(server_name: str) -> Optional[_core.MCPServer is_lazy = server_name in _core._lazy_server_configs if is_lazy and (server is None or server.session is None): _core._ensure_lazy_server_connected(server_name) - with _core._lock: - server = _core._servers.get(server_name) - return server - if server is not None and server.session is None and server._is_recycled_stdio(): + elif server is not None and server.session is None and server._is_recycled_stdio(): _core._request_lazy_reconnect(server_name, server) - with _core._lock: - server = _core._servers.get(server_name) - return server + else: + return server + with _core._lock: + return _core._servers.get(server_name) async def _discover_and_register_server(name: str, config: dict) -> List[str]: """Connect one server, register its tools; return the registered names.""" - connect_timeout = config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT) # The claim callback runs inside _connect_server while this frame is suspended; a list # append avoids a nonlocal rebind. claimed: List[_core.MCPServerTask] = [] claim_token = _core._connect_server_claim.set(claimed.append) try: - server = await asyncio.wait_for(_core._connect_server(name, config), timeout=connect_timeout) + server = await asyncio.wait_for(_core._connect_server(name, config), + timeout=config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT)) except BaseException: server = claimed[0] if claimed else None task = server._task if server is not None else None @@ -211,21 +207,16 @@ async def _discover_and_register_server(name: str, config: dict) -> List[str]: and not task.done() and not task_cancelling): # Recoverable park: the run task stays alive to self-probe, so adopt it for # shutdown/revival. - with _core._lock: - _core._servers[name] = server - _core._server_scope_keys[name] = _core._mcp_registry_scope() + _adopt_server(name, server) elif server is not None: await server.shutdown() raise finally: _core._connect_server_claim.reset(claim_token) - with _core._lock: _core._server_connecting.discard(name) _core._server_connect_errors.pop(name, None) - _core._servers[name] = server - _core._server_scope_keys[name] = _core._mcp_registry_scope() - + _adopt_server(name, server) registered_names = _core._register_server_tools(name, server, config) server._registered_tool_names = list(registered_names) logger.info("MCP server '%s' (%s): registered %d tool(s): %s", name, @@ -258,7 +249,6 @@ def _select_new_servers(servers: Dict[str, dict]) -> Dict[str, dict]: _core._parallel_safe_servers.add(srv_name) else: _core._parallel_safe_servers.discard(srv_name) - for srv in stale_cached: _core._signal_reconnect(srv) return new_servers @@ -366,23 +356,19 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: if not _core._ensure_mcp_sdk(): logger.debug("MCP SDK not available -- skipping explicit MCP registration") return [] - servers = _core._filter_suspicious_mcp_servers(servers) if not servers: logger.debug("No explicit MCP servers provided") return [] - new_servers = _select_new_servers(servers) if not new_servers: return _core._existing_tool_names() - new_servers, lazy_registered, lazy_server_count = _register_lazy_from_cache(new_servers) if not new_servers: if lazy_registered: logger.info("MCP: registered %d lazy tool(s) from schema cache (no processes spawned)", lazy_registered) return _core._existing_tool_names() - _core._ensure_mcp_loop() _run_discovery_pass(new_servers) _log_summary("MCP: registered", new_servers, lazy_tools=lazy_registered, lazy_servers=lazy_server_count) @@ -402,7 +388,6 @@ def _acquire_discovery_lock_with_retry(): cookie = _core._try_acquire_mcp_discovery_lock() if cookie is not None: break - if cookie is None: logger.warning("MCP discovery lock still held after %d retries -- running discovery unguarded", _core._MCP_DISCOVERY_LOCK_MAX_RETRIES) @@ -419,19 +404,16 @@ def discover_mcp_tools() -> List[str]: if not servers: logger.debug("No MCP servers configured") return [] - # SDK import deferred to here so a config without servers never pays it. if not _core._ensure_mcp_sdk(): logger.debug("MCP SDK not available -- skipping MCP tool discovery") return [] - cookie = _acquire_discovery_lock_with_retry() try: with _core._lock: connecting = set(_core._server_connecting) new_server_names = [name for name, cfg in servers.items() if name not in _core._servers and name not in connecting and _enabled(cfg)] - tool_names = _core.register_mcp_servers(servers) if new_server_names: _log_summary(" MCP:", new_server_names) @@ -461,7 +443,6 @@ def get_mcp_status() -> List[dict]: configured = _core._load_mcp_config() if not configured: return [] - with _core._lock: active_servers = dict(_core._servers) connecting = set(_core._server_connecting) @@ -474,7 +455,6 @@ def get_mcp_status() -> List[dict]: result: List[dict] = [] for name, cfg in configured.items(): transport = cfg.get("transport", "http") if "url" in cfg else "stdio" - enabled = _enabled(cfg) server = active_servers.get(name) if server and server.session is not None: entry = _entry(name, transport, "connected", connected=True) @@ -482,7 +462,7 @@ def get_mcp_status() -> List[dict]: else len(server._tools)) if server._sampling: entry["sampling"] = dict(server._sampling.metrics) - elif not enabled: + elif not _enabled(cfg): entry = _entry(name, transport, "disabled") elif name in connecting: entry = _entry(name, transport, "connecting") @@ -491,7 +471,6 @@ def get_mcp_status() -> List[dict]: else: entry = _entry(name, transport, "configured") result.append(entry) - return result @@ -500,17 +479,10 @@ def probe_mcp_server_tools() -> Dict[str, List[tuple]]: without registering anything. Failed servers are omitted.""" if not _core._ensure_mcp_sdk(): return {} - - servers_config = _core._load_mcp_config() - if not servers_config: - return {} - - enabled = {k: v for k, v in servers_config.items() if _enabled(v)} + enabled = {k: v for k, v in (_core._load_mcp_config() or {}).items() if _enabled(v)} if not enabled: return {} - _core._ensure_mcp_loop() - result: Dict[str, List[tuple]] = {} probed_servers: List[_core.MCPServerTask] = [] @@ -533,7 +505,6 @@ def probe_mcp_server_tools() -> Dict[str, List[tuple]]: logger.debug("MCP probe failed: %s", exc) finally: _core._stop_mcp_loop_if_idle() - return result diff --git a/tools/mcp_tool_handlers.py b/tools/mcp_tool_handlers.py index b06457768c..133e796b20 100644 --- a/tools/mcp_tool_handlers.py +++ b/tools/mcp_tool_handlers.py @@ -33,7 +33,6 @@ def _trust_gate_check(server_name: str, tool_name: str) -> Optional[str]: # Lazy import: tools.approval routes the prompt to whichever surface owns the session. try: from tools.approval import request_elicitation_consent - answer = request_elicitation_consent( f"MCP tool '{tool_name}' on UNTRUSTED server '{server_name}' wants to run. This " f"tool is write-capable (no readOnlyHint=true annotation) and may modify external state.", @@ -113,8 +112,7 @@ def _strike(server_name: str, message: str, **extra) -> str: def _mcp_loop_running() -> bool: - loop = _core._mcp_loop - return loop is not None and loop.is_running() + return _core._mcp_loop is not None and _core._mcp_loop.is_running() def _lookup_reconnectable_server(server_name: str, require_loop: bool = False): @@ -153,12 +151,8 @@ def _handle_auth_error_and_retry(server_name: str, exc: BaseException, retry_cal return None from tools.mcp_oauth_manager import get_manager manager = get_manager() - - async def _recover(): - return await manager.handle_401(server_name, None) - try: - recovered = _core._run_on_mcp_loop(_recover, timeout=10) + recovered = _core._run_on_mcp_loop(lambda: manager.handle_401(server_name, None), timeout=10) except Exception as rec_exc: logger.warning("MCP OAuth '%s': recovery attempt failed: %s", server_name, rec_exc) recovered = False @@ -223,7 +217,6 @@ def _handle_stdio_child_exited_and_retry(server_name: str, exc: Exception, retry # No MCP loop to wait on (non-async adapters, tests) — still request the respawn # so the next call lands on a live transport. _core._signal_reconnect(srv) - if not reconnected: return _strike( server_name, @@ -277,9 +270,8 @@ def _invoke_with_recovery(server_name: str, call_once: Callable[[], str], op: st def _mark_server_call_started(server: Any) -> None: """Record a user-visible MCP operation when the server supports it.""" - mark_tool_call = getattr(server, "mark_tool_call", None) - if callable(mark_tool_call): - mark_tool_call() + if callable(getattr(server, "mark_tool_call", None)): + server.mark_tool_call() @asynccontextmanager @@ -320,7 +312,6 @@ async def _call_tool_racing_stdio_death(server, server_name: str, tool_name: str if not (_watch_children is not None and inspect.iscoroutinefunction(_watch_children) and asyncio.iscoroutine(_call_coro)): # Stubbed sessions return a non-awaitable, or there is no child-watcher to race: plain await. return await _call_coro if asyncio.iscoroutine(_call_coro) else _call_coro - rpc_task = asyncio.ensure_future(_call_coro) watch_task = asyncio.ensure_future(_watch_children()) try: @@ -388,7 +379,6 @@ def _render_call_tool_result(result, server_name: str) -> str: ``.is_error`` is ``.isError`` before mcp 2.0.""" if mcp_field(result, "is_error", "isError", False): return tool_error(_sanitize_error(_truncate_mcp_text_result(_error_result_text(result) or "MCP tool returned an error"))) - text_result = _render_content_blocks(result, server_name) structured = _capped_structured_content(result) meta = _strip_reserved_meta_keys(mcp_field(result, "meta", "meta")) @@ -436,9 +426,8 @@ def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float): finally: server._pending_call_context = None # Round-trip completed: transport is healthy even if the tool returned isError. - _mark_proven = getattr(server, "_mark_session_proven", None) - if _mark_proven is not None: - _mark_proven() + if getattr(server, "_mark_session_proven", None) is not None: + server._mark_session_proven() return _render_call_tool_result(result, server_name) def _on_failure(exc): @@ -502,9 +491,9 @@ def _render_resource_list(all_resources, server_name: str) -> dict: if "uri" in entry: entry["uri"] = str(entry["uri"]) # Key stays camelCase — this is the tool's own JSON output shape. - _mime = mcp_field(r, "mime_type", "mimeType") - if _mime: - entry["mimeType"] = _mime + mime = mcp_field(r, "mime_type", "mimeType") + if mime: + entry["mimeType"] = mime resources.append(entry) return {"resources": resources} @@ -539,8 +528,7 @@ def _render_get_prompt(result, server_name: str) -> dict: for msg in getattr(result, "messages", []): entry = _pick(msg, ("role", "role")) if hasattr(msg, "content"): - content = msg.content - entry["content"] = strip_unicode_tags(content.text if hasattr(content, "text") else str(content)) + entry["content"] = strip_unicode_tags(msg.content.text if hasattr(msg.content, "text") else str(msg.content)) messages.append(entry) resp = {"messages": messages} if getattr(result, "description", None): @@ -552,6 +540,7 @@ def _utility_factory(op: str, log_label: str, rpc, render, required: Optional[st """``(server_name, tool_timeout) -> sync handler`` for one utility tool.""" def _factory(server_name: str, tool_timeout: float): return _make_utility_handler(server_name, tool_timeout, op, log_label, rpc, render, required) + return _factory diff --git a/tools/mcp_tool_registration.py b/tools/mcp_tool_registration.py index 9f49adbd23..c67b196a12 100644 --- a/tools/mcp_tool_registration.py +++ b/tools/mcp_tool_registration.py @@ -154,13 +154,9 @@ class _CachedMCPTool: @classmethod def from_cache_dicts(cls, raws: Iterable[Any]) -> List["_CachedMCPTool"]: """Cached rows -> stand-ins; rows that are not dicts or lack a name are dropped.""" - out = [] - for raw in raws: - if isinstance(raw, dict) and raw.get("name"): - schema = raw.get("inputSchema") - out.append(cls(raw["name"], raw.get("description") or "", - schema if isinstance(schema, dict) else {}, raw.get("annotations"))) - return out + return [cls(raw["name"], raw.get("description") or "", + raw["inputSchema"] if isinstance(raw.get("inputSchema"), dict) else {}, raw.get("annotations")) + for raw in raws if isinstance(raw, dict) and raw.get("name")] @dataclass @@ -224,7 +220,6 @@ def _resolve_name_collisions(name: str, candidates: List[_Candidate]) -> List[_C seen.add((c.registry_name, c.origin)) unique.append(c) origins_by_name.setdefault(c.registry_name, set()).add(c.origin) - ambiguous: Dict[str, List[str]] = {} shadowed: set[tuple[str, str]] = set() for registry_name, origins in origins_by_name.items(): @@ -273,7 +268,6 @@ def _register_candidates(name: str, candidates: List[_Candidate], *, check_fn: C ``ToolRegistry.register()`` is the atomic ownership gate and its verdict is re-read after every call.""" from tools.registry import registry - toolset_name = f"mcp-{name}" registered: List[str] = [] for c in candidates: @@ -301,7 +295,6 @@ def _write_schema_cache(name: str, server: "MCPServerTask", config: dict, should lazily without spawning it. Never raises.""" try: from tools.mcp_schema_cache import config_fingerprint, write_cache_entry - tools_payload = [] for t in server._tools: if should_register(t.name): @@ -345,7 +338,6 @@ def _register_from_cache_sync(name: str, config: dict, entry: dict) -> List[str] ``_ensure_lazy_server_connected``. Trust metadata is recorded first so the call-time gate is identical whether the server was spawned live or registered from cache.""" from tools.mcp_schema_cache import config_fingerprint, tools_from_cache_entry, utility_tools_from_cache_entry - tool_timeout = _resolve_tool_timeout(config) cached_tools = _CachedMCPTool.from_cache_dicts(tools_from_cache_entry(entry)) _record_tool_trust_metadata(name, config, cached_tools) diff --git a/tools/mcp_tool_schema.py b/tools/mcp_tool_schema.py index eea73b3c11..be7b6040cf 100644 --- a/tools/mcp_tool_schema.py +++ b/tools/mcp_tool_schema.py @@ -101,12 +101,10 @@ def _normalize_mcp_input_schema(schema: dict | None) -> dict: if not schema: return dict(_EMPTY_OBJECT_SCHEMA) from tools.schema_sanitizer import collapse_const_unions, strip_nullable_unions - normalized = _rewrite_local_refs(schema) normalized = strip_nullable_unions(normalized, keep_nullable_hint=True) normalized = collapse_const_unions(normalized) normalized = _repair_object_shape(normalized) - if not isinstance(normalized, dict): return dict(_EMPTY_OBJECT_SCHEMA) if normalized.get("type") == "object" and "properties" not in normalized: diff --git a/tools/mcp_tool_server_run.py b/tools/mcp_tool_server_run.py index 430651a24c..e380d65d47 100644 --- a/tools/mcp_tool_server_run.py +++ b/tools/mcp_tool_server_run.py @@ -35,6 +35,11 @@ class MCPServerRunMixin: except (asyncio.CancelledError, Exception): pass + def _event_waiters(self) -> tuple: + """Fresh ``(shutdown, reconnect)`` wait tasks; cancel them via ``_cancel_waiters``.""" + return (asyncio.ensure_future(self._shutdown_event.wait()), + asyncio.ensure_future(self._reconnect_event.wait())) + def _recycle_if_due(self) -> bool: """Latch a stdio idle/lifetime recycle when its deadline has passed.""" recycle_reason = self._stdio_recycle_reason() @@ -56,26 +61,21 @@ class MCPServerRunMixin: keepalive_interval = max( _core._MIN_KEEPALIVE_INTERVAL, float(self._config.get("keepalive_interval", _core._DEFAULT_KEEPALIVE_INTERVAL))) - - shutdown_task = asyncio.create_task(self._shutdown_event.wait()) - reconnect_task = asyncio.create_task(self._reconnect_event.wait()) + shutdown_task, reconnect_task = self._event_waiters() try: while True: if self._recycle_if_due(): return "recycle" - timeout = keepalive_interval recycle_deadline = self._next_stdio_recycle_deadline() if recycle_deadline is not None: timeout = max(0.0, min(timeout, recycle_deadline - time.monotonic())) - done, _pending = await asyncio.wait( {shutdown_task, reconnect_task}, timeout=timeout, return_when=asyncio.FIRST_COMPLETED) if done: break if self._recycle_if_due(): return "recycle" - # Timeout: probe for a stale session — but NEVER while an RPC is in flight (a # concurrent ping can wedge the single stdio stream, and a busy server is # provably alive anyway). @@ -98,7 +98,6 @@ class MCPServerRunMixin: self._mark_session_proven() finally: await self._cancel_waiters(shutdown_task, reconnect_task) - if self._shutdown_event.is_set(): self._fail_inflight_calls("shutdown") return "shutdown" @@ -112,8 +111,7 @@ class MCPServerRunMixin: """Wait, while parked, for a reconnect request or shutdown. Returns ``"shutdown"`` or ``"reconnect"`` (explicit request or, with ``timeout``, the periodic self-probe); the reconnect event is cleared first. Shutdown wins a tie.""" - shutdown_task = asyncio.ensure_future(self._shutdown_event.wait()) - reconnect_task = asyncio.ensure_future(self._reconnect_event.wait()) + shutdown_task, reconnect_task = self._event_waiters() try: await asyncio.wait({shutdown_task, reconnect_task}, return_when=asyncio.FIRST_COMPLETED, timeout=timeout) finally: @@ -154,10 +152,8 @@ class MCPServerRunMixin: self._auth_type = (config.get("auth") or "").lower().strip() self._idle_timeout_seconds = _core._get_lifecycle_seconds(config, "idle_timeout_seconds") self._max_lifetime_seconds = _core._get_lifecycle_seconds(config, "max_lifetime_seconds") - # The _MCP_*_TYPES flags are False until the lazy SDK import runs. _core._ensure_mcp_sdk() - sampling_config = config.get("sampling", {}) self._sampling = (_core.SamplingHandler(self.name, sampling_config) if sampling_config.get("enabled", True) and _core._MCP_SAMPLING_TYPES else None) @@ -166,12 +162,10 @@ class MCPServerRunMixin: elicitation_config = config.get("elicitation", {}) self._elicitation = (_core.ElicitationHandler(self.name, elicitation_config, owner=self) if elicitation_config.get("enabled", True) and _core._MCP_ELICITATION_TYPES else None) - if "url" in config and "command" in config: logger.warning("MCP server '%s' has both 'url' and 'command' in config. " "Using HTTP transport ('url'). Remove 'command' to silence " "this warning.", self.name) - if not self._is_http(): return True try: @@ -206,17 +200,12 @@ class MCPServerRunMixin: """ if not await self._prepare_run(config): return - self._reconnect_retries = 0 budget = _RetryBudget() - while True: try: - if self._is_http(): - lifecycle_reason = await self._run_http(config) - else: - lifecycle_reason = await self._run_stdio(config) - if not await self._on_clean_return(lifecycle_reason, budget): + run_transport = self._run_http if self._is_http() else self._run_stdio + if not await self._on_clean_return(await run_transport(config), budget): break except asyncio.CancelledError: # Not a connection failure: re-raise so cancellation reaches asyncio and @@ -253,11 +242,9 @@ class MCPServerRunMixin: logger.info("MCP server '%s': reconnect after teardown race " "(in-flight calls were failed); not charging the " "rapid-drop budget", self.name) - self._teardown_race = False - budget.backoff = 1.0 + self._teardown_race, budget.backoff = False, 1.0 elif self._session_proven: - self._reconnect_retries = 0 - budget.backoff = 1.0 + self._reconnect_retries, budget.backoff = 0, 1.0 else: self._reconnect_retries += 1 if self._reconnect_retries > _core._MAX_RECONNECT_RETRIES: @@ -280,8 +267,7 @@ class MCPServerRunMixin: again instead of burning 5 rapid retries. False on shutdown.""" if await self._park(revival_reason): return False - self._reconnect_retries = _core._MAX_RECONNECT_RETRIES - budget.backoff = 1.0 + self._reconnect_retries, budget.backoff = _core._MAX_RECONNECT_RETRIES, 1.0 return True async def _park_initial_failure(self, exc: Exception, revival_reason: str, budget: "_RetryBudget") -> bool: @@ -291,8 +277,7 @@ class MCPServerRunMixin: self._ready.set() if await self._park(revival_reason): return False - budget.initial_retries = 0 - self._reconnect_retries = 0 + budget.initial_retries = self._reconnect_retries = 0 budget.backoff = 1.0 self._error = None self._ready.clear() @@ -314,20 +299,16 @@ class MCPServerRunMixin: "failed, marking unavailable while retrying: %s: %s", self.name, type(root).__name__, root) self._recycled_reason = None - # Initial-connect ladder: a transient blip at startup must not kill the server. Gated # on _ever_connected (never cleared), not _ready (cleared every reconnect cycle). if not self._ever_connected: return await self._on_initial_connect_error(exc, root, failure_class, budget) - if self._shutdown_event.is_set(): logger.debug("MCP server '%s' disconnected during shutdown: %s: %s", self.name, type(root).__name__, root) return False - if failure_class == "permanent": return await self._on_permanent_error(root, budget) - self._reconnect_retries += 1 if self._reconnect_retries > _core._MAX_RECONNECT_RETRIES: logger.warning( @@ -337,7 +318,6 @@ class MCPServerRunMixin: self.name, _core._MAX_RECONNECT_RETRIES, _core._PARKED_RETRY_INTERVAL, type(root).__name__, root) return await self._park_and_rearm("from parked state", budget) - logger.debug("MCP server '%s' connection lost (attempt %d/%d), " "reconnecting in %.0fs: %s: %s", self.name, self._reconnect_retries, _core._MAX_RECONNECT_RETRIES, @@ -351,19 +331,12 @@ class MCPServerRunMixin: # Deterministic failure (bad command, non-MCP URL, 401/403): park at once instead # of burning the ladder. Auth failures park rather than return so the task stays # alive to pick up fresh tokens later. - if _core._is_auth_error(root): - logger.warning( - "MCP server '%s' failed initial authentication, " - "parking until credentials change; re-authenticate with `hermes mcp login %s` " - "(state: connecting → parked): %s: %s", - self.name, self.name, type(root).__name__, root) - else: - logger.warning( - "MCP server '%s' failed initial connection with a " - "permanent error, parking without retries (state: connecting → parked): %s: %s", - self.name, type(root).__name__, root) + detail = (f"authentication, parking until credentials change; re-authenticate with " + f"`hermes mcp login {self.name}`" if _core._is_auth_error(root) + else "connection with a permanent error, parking without retries") + logger.warning("MCP server '%s' failed initial %s (state: connecting → parked): %s: %s", + self.name, detail, type(root).__name__, root) return await self._park_initial_failure(exc, "after permanent initial failure", budget) - budget.initial_retries += 1 if budget.initial_retries > _core._MAX_INITIAL_CONNECT_RETRIES: logger.warning( @@ -372,7 +345,6 @@ class MCPServerRunMixin: "requested (state: connecting → parked): %s: %s", self.name, _core._MAX_INITIAL_CONNECT_RETRIES, type(root).__name__, root) return await self._park_initial_failure(exc, "after initial connection failures", budget) - logger.debug( "MCP server '%s' initial connection failed (attempt %d/%d), retrying in %.0fs: %s: %s", self.name, budget.initial_retries, _core._MAX_INITIAL_CONNECT_RETRIES, budget.backoff, @@ -381,8 +353,7 @@ class MCPServerRunMixin: if self._shutdown_event.is_set(): self._error = exc self._ready.set() - return False - return True + return not self._shutdown_event.is_set() async def _on_permanent_error(self, root: BaseException, budget: "_RetryBudget") -> bool: # An auth failure on a PROVEN session is often a corrupt OAuth lock from a raced @@ -395,8 +366,7 @@ class MCPServerRunMixin: "healthy session — marking suspect and forcing " "one reconnect instead of parking (state: connected → suspect): %s: %s", self.name, type(root).__name__, root) - self._reconnect_retries = 0 - budget.backoff = 1.0 + self._reconnect_retries, budget.backoff = 0, 1.0 await asyncio.sleep(_core._jittered(1.0)) return not self._shutdown_event.is_set() # Deterministic failure on a working server: park now. @@ -426,9 +396,8 @@ class MCPServerRunMixin: async def shutdown(self): """Signal the Task to exit and wait for clean resource teardown.""" self._shutdown_event.set() - # _shutdown_event alone unblocks _wait_for_lifecycle_event (it checks shutdown first), - # but setting reconnect too closes any race where the helper misses the shutdown flag - # after returning "reconnect". + # Also set reconnect: closes any race where _wait_for_lifecycle_event misses the + # shutdown flag after returning "reconnect". self._reconnect_event.set() if self._task and not self._task.done(): try: @@ -453,7 +422,6 @@ class MCPServerRunMixin: AND when the reconnect budget is exhausted, so a dead server never leaves phantom tool definitions bloating the prompt cache and producing "not connected" errors.""" from tools.registry import registry - for tool_name in list(getattr(self, "_registered_tool_names", [])): registry.deregister(tool_name, scope=_core._server_registry_scope(self.name)) _core._forget_mcp_tool_server(tool_name)