From ee81b1abdd1f86ff2e620fa0c826883e8de05640 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 22:45:12 -0700 Subject: [PATCH] refactor(tools): compact MCP oauth/config/health/lifecycle/agent/content modules Dead-code removal (_make_hermes_provider_class factory -> direct class def, _same_endpoint/_context_var_value inlined), image/audio cache unified into _cache_mcp_media_block, match->isinstance chain, closer hugging, docstring compaction keeping every invariant. Tool schemas byte-identical. --- tools/mcp_oauth_manager.py | 289 +++++++++++++----------------------- tools/mcp_oauth_provider.py | 57 +++---- tools/mcp_stdio_watchdog.py | 38 ++--- tools/mcp_tool_agent.py | 138 ++++++++--------- tools/mcp_tool_config.py | 110 ++++++-------- tools/mcp_tool_content.py | 149 ++++++++----------- tools/mcp_tool_health.py | 80 +++++----- tools/mcp_tool_lifecycle.py | 119 ++++++--------- tools/setup_mcp_tool.py | 48 ++---- 9 files changed, 403 insertions(+), 625 deletions(-) diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index 99891d21e3..ae3945ca71 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -1,16 +1,11 @@ -#!/usr/bin/env python3 """Central manager for per-server MCP OAuth state (one instance per process). -Holds per-server provider instances and coordinates cross-process token reload (mtime-based -disk watch, so tokens refreshed by cron/another CLI are picked up without a restart), 401 -deduplication (N concurrent tool calls hitting 401 with the same access_token trigger one -recovery attempt) and reconnect signalling (``MCPServerTask`` in ``mcp_tool.py`` drives the -reconnect; the manager decides when it is warranted). - -This module is the ONLY place that instantiates the SDK's ``OAuthClientProvider`` for runtime -use; other code paths go through ``get_manager()``. We lean on the SDK's lazy refresh rather -than refreshing before every op: one ``stat()`` per tool call is cheaper than an await + -refresh round-trip. +Holds per-server providers and coordinates cross-process token reload (mtime-based disk watch, +so tokens refreshed by cron/another CLI are picked up without a restart), 401 deduplication +(N concurrent 401s with the same access_token trigger one recovery) and reconnect signalling +(``MCPServerTask`` drives the reconnect; the manager decides when). The ONLY place that +instantiates the SDK's ``OAuthClientProvider`` for runtime use. We rely on the SDK's lazy +refresh: one ``stat()`` per tool call is cheaper than an await + refresh round-trip. """ from __future__ import annotations @@ -27,25 +22,18 @@ from tools.mcp_oauth_provider import HermesProviderMixin logger = logging.getLogger(__name__) - -def _same_endpoint(a: str, b: str) -> bool: - """True if two URLs target the same endpoint: scheme, host (case-insensitive) and path, - ignoring query/fragment. Confirms a rejected response actually came from the OAuth token - endpoint before we act on an ``invalid_client`` body.""" - from urllib.parse import urlsplit - - try: - pa, pb = urlsplit(a), urlsplit(b) - except ValueError: # pragma: no cover — malformed URL - return False - return pa.scheme == pb.scheme and pa.netloc.lower() == pb.netloc.lower() and pa.path.rstrip("/") == pb.path.rstrip("/") +try: + from mcp.client.auth.oauth2 import OAuthClientProvider as _SDKOAuthClientProvider + _SDK_BASES: tuple = (_SDKOAuthClientProvider,) +except ImportError: # pragma: no cover — SDK required in CI; module must still import + _SDK_BASES = () @dataclass class _ProviderEntry: - """Per-server OAuth state. ``last_mtime_ns`` is the last-seen tokens-file mtime (0 = never - read) for external-refresh detection; ``lock`` binds to whichever asyncio loop first awaits - it (the MCP event loop); ``pending_401`` dedupes thundering-herd 401s by failed access_token.""" + """Per-server OAuth state. ``last_mtime_ns``: last-seen tokens-file mtime (0 = never read) + for external-refresh detection; ``lock`` binds to the first asyncio loop awaiting it (the MCP + loop); ``pending_401`` dedupes thundering-herd 401s by failed access_token.""" server_url: str oauth_config: Optional[dict] @@ -55,12 +43,11 @@ class _ProviderEntry: pending_401: dict[str, "asyncio.Future[bool]"] = field(default_factory=dict) -# -- HermesMCPOAuthProvider — OAuthClientProvider subclass with disk-watch ---- -class _HermesRuntimeProviderMixin: - """Runtime-only provider behaviour layered over ``HermesProviderMixin``: pre-flow disk-mtime - reload, expiry seeding on cold load, pre-flight metadata discovery, dead-client-registration - detection and the bidirectional ``async_auth_flow`` bridge. Must precede the SDK class in - the MRO.""" +class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): + """OAuthClientProvider with pre-flow disk-mtime reload (external refreshes become visible to + a running session), expiry seeding on cold load, pre-flight metadata discovery, dead-client + registration detection and the bidirectional ``async_auth_flow`` bridge. Token-endpoint + fixes come from ``HermesProviderMixin``. Only usable when the SDK's OAuth module imported.""" _hermes_logger = logger @@ -71,19 +58,16 @@ class _HermesRuntimeProviderMixin: # generator from another task. A binary semaphore keeps mutual exclusion without task # ownership; async_auth_flow narrows its scope around resource I/O. import anyio - self.context.lock = anyio.Semaphore(1, max_value=1) self._hermes_server_name = server_name self._hermes_home = "" - # A config-supplied (pre-registered) client_id rejected as invalid_client means the - # *config* is wrong — re-registration can't help, so only dynamically-registered - # clients auto-heal. + # A config-supplied client_id rejected as invalid_client means the *config* is wrong — + # re-registration can't help, so only dynamically-registered clients auto-heal. self._hermes_preregistered = preregistered def _hermes_storage(self): """The context storage when it is a ``HermesTokenStorage``, else None.""" from tools.mcp_oauth import HermesTokenStorage - storage = self.context.storage return storage if isinstance(storage, HermesTokenStorage) else None @@ -93,56 +77,44 @@ class _HermesRuntimeProviderMixin: async def _initialize(self) -> None: """Load stored state, seed ``token_expiry_time``, restore/prefetch metadata. - The SDK's ``_initialize`` populates ``current_tokens`` but never calls - ``update_token_expiry``, so ``is_token_valid()`` is True for any loaded token regardless - of age and a restarted process ships stale Bearer tokens (some providers answer 200 with - an app-level auth error the transport can't see). Seeding the expiry makes the SDK take - ``can_refresh_token()`` and refresh first; ``HermesTokenStorage`` persists absolute - ``expires_at`` so the TTL reflects wall-clock age. Metadata is restored from disk, else - discovered pre-flight when we hold tokens but no metadata: otherwise ``_refresh_token`` - guesses ``{server_url}/token`` (wrong for split-origin providers such as BetterStack), - 404s, and we fall through to browser reauth. - """ + The SDK's ``_initialize`` never calls ``update_token_expiry``, so ``is_token_valid()`` is + True for any loaded token regardless of age and a restarted process ships stale Bearer + tokens (some providers answer 200 with an app-level auth error). Seeding the expiry makes + the SDK refresh first; ``HermesTokenStorage`` persists absolute ``expires_at`` so the TTL + reflects wall-clock age. Metadata is restored from disk, else discovered pre-flight when + we hold tokens but no metadata: otherwise ``_refresh_token`` guesses ``{server_url}/token`` + (wrong for split-origin providers), 404s, and we fall through to browser reauth.""" await super()._initialize() tokens = self.context.current_tokens if tokens is not None and tokens.expires_in is not None: self.context.update_token_expiry(tokens) - storage = self._hermes_storage() if storage is not None and self.context.oauth_metadata is None: meta = storage.load_oauth_metadata() if meta is not None: self.context.oauth_metadata = meta - logger.debug( - "MCP OAuth '%s': restored metadata from disk (token_endpoint=%s)", - self._hermes_server_name, meta.token_endpoint, - ) - + logger.debug("MCP OAuth '%s': restored metadata from disk (token_endpoint=%s)", + self._hermes_server_name, meta.token_endpoint) if tokens is not None and self.context.oauth_metadata is None: try: await self._prefetch_oauth_metadata() - except Exception as exc: # pragma: no cover — non-fatal: the SDK's 401-branch discovery runs next request + except Exception as exc: # pragma: no cover — the SDK's 401-branch discovery runs next request self._log_nonfatal("pre-flight metadata discovery", exc) async def _prefetch_oauth_metadata(self) -> None: - """Fetch PRM + ASM from the well-known endpoints and cache on context. Mirrors the SDK's - 401-branch discovery but runs before the first request, using the SDK's own URL - builders/response handlers so we track whatever the pinned SDK version expects.""" + """Fetch PRM + ASM from the well-known endpoints before the first request, using the + SDK's own URL builders/response handlers so we track whatever the pinned SDK expects.""" # The SDK's httpx flavour, not Hermes' — mcp 2.0 builds on httpx2 and - # `create_oauth_metadata_request` returns *its* Request objects, which only its own - # AsyncClient can send (tools.mcp_tool.sdk_httpx). + # `create_oauth_metadata_request` returns *its* Request objects. from tools.mcp_tool import sdk_httpx httpx = sdk_httpx() if httpx is None: # pragma: no cover — SDK import would have failed return from mcp.client.auth.utils import ( build_oauth_authorization_server_metadata_discovery_urls, - build_protected_resource_metadata_discovery_urls, - create_oauth_metadata_request, - handle_auth_metadata_response, - handle_protected_resource_response, + build_protected_resource_metadata_discovery_urls, create_oauth_metadata_request, + handle_auth_metadata_response, handle_protected_resource_response, ) - server_url = self.context.server_url async def _send(client, url: str, label: str): @@ -153,7 +125,7 @@ class _HermesRuntimeProviderMixin: return None async with httpx.AsyncClient(timeout=10.0) as client: - # Step 1: PRM discovery to learn the authorization_server URL. + # PRM discovery to learn the authorization_server URL. for url in build_protected_resource_metadata_discovery_urls(None, server_url): resp = await _send(client, url, "PRM") prm = await handle_protected_resource_response(resp) if resp is not None else None @@ -162,8 +134,7 @@ class _HermesRuntimeProviderMixin: if prm.authorization_servers: self.context.auth_server_url = str(prm.authorization_servers[0]) break - - # Step 2: ASM discovery against auth_server_url (server_url fallback for legacy providers). + # ASM discovery against auth_server_url (server_url fallback for legacy providers). for url in build_oauth_authorization_server_metadata_discovery_urls(self.context.auth_server_url, server_url): resp = await _send(client, url, "ASM") if resp is None: @@ -173,19 +144,15 @@ class _HermesRuntimeProviderMixin: break if asm: self.context.oauth_metadata = asm - # Persist now so a later cold-load skips discovery. - storage = self._hermes_storage() + storage = self._hermes_storage() # persist now so a later cold-load skips discovery if storage is not None: storage.save_oauth_metadata(asm) - logger.debug( - "MCP OAuth '%s': pre-flight ASM discovered token_endpoint=%s", - self._hermes_server_name, asm.token_endpoint, - ) + logger.debug("MCP OAuth '%s': pre-flight ASM discovered token_endpoint=%s", + self._hermes_server_name, asm.token_endpoint) break def _persist_oauth_metadata_if_changed(self) -> None: - """Save metadata the SDK discovered lazily (401 branch) for future restarts; no-op when - absent, not our storage, or unchanged.""" + """Save metadata the SDK discovered lazily (401 branch); no-op when absent/unchanged.""" meta = self.context.oauth_metadata storage = self._hermes_storage() if meta is None or storage is None: @@ -195,14 +162,23 @@ class _HermesRuntimeProviderMixin: storage.save_oauth_metadata(meta) async def _is_invalid_client_at_token_endpoint(self, response: Any) -> bool: - """True when *response* is the token endpoint rejecting our client_id with - ``invalid_client`` (whole word, so RFC 7591's ``invalid_client_metadata`` does not trip - it). The body is read only after the endpoint matches.""" + """True when *response* is the token endpoint (same scheme/host/path, query ignored) + rejecting our client_id with ``invalid_client`` — whole word, so RFC 7591's + ``invalid_client_metadata`` does not trip it. The body is read only after the endpoint + matches.""" + from urllib.parse import urlsplit meta = getattr(self.context, "oauth_metadata", None) token_endpoint = str(meta.token_endpoint) if meta is not None and getattr(meta, "token_endpoint", None) else None req = getattr(response, "request", None) req_url = str(req.url) if req is not None else None - if not token_endpoint or not req_url or not _same_endpoint(req_url, token_endpoint): + if not token_endpoint or not req_url: + return False + try: + pa, pb = urlsplit(req_url), urlsplit(token_endpoint) + except ValueError: # pragma: no cover — malformed URL + return False + if not (pa.scheme == pb.scheme and pa.netloc.lower() == pb.netloc.lower() + and pa.path.rstrip("/") == pb.path.rstrip("/")): return False body = await response.aread() return re.search(rb"\binvalid_client\b", body.lower()) is not None @@ -210,46 +186,37 @@ class _HermesRuntimeProviderMixin: async def _maybe_flag_poisoned_client(self, response: Any) -> None: """Detect a dead client registration and force re-registration. - An ``invalid_client`` rejection of our ``client_id`` at the token endpoint (exchange or - refresh) proves the cached registration is dead server-side; delete ``client.json`` - (+ stale metadata) so the SDK re-runs DCR next flow. The browser-side "Redirect URI - Mismatch" case has no HTTP signal and is left to ``hermes mcp reauth``. - - Conservative by construction — acts ONLY when status is 400/401, the request hit the - discovered ``token_endpoint`` (the only request carrying our ``client_id``), and the body - carries ``invalid_client``. Pre-registered clients are never poisoned. Best-effort: any - failure is swallowed so a miss never breaks the live flow. If ``token_endpoint`` was - never discovered the guard returns early. - """ + An ``invalid_client`` rejection of our ``client_id`` at the token endpoint proves the + cached registration is dead server-side; delete ``client.json`` (+ stale metadata) so the + SDK re-runs DCR next flow. Conservative: acts ONLY on status 400/401 at the discovered + ``token_endpoint`` (the only request carrying our ``client_id``) with ``invalid_client`` + in the body; pre-registered clients are never poisoned; any failure is swallowed so a + miss never breaks the live flow. The browser-side "Redirect URI Mismatch" case has no + HTTP signal and is left to ``hermes mcp reauth``.""" try: if self._hermes_preregistered or getattr(response, "status_code", None) not in (400, 401): return if not await self._is_invalid_client_at_token_endpoint(response): return - storage = self._hermes_storage() # If the rejected client_id was our CIMD URL, re-presenting it would loop (the # server already fetched and refused it). Drop the URL so the retry takes DCR, and # mark it on disk so the next process doesn't walk back into the same refusal # (`hermes mcp login` clears the marker). cimd_url = getattr(self.context, "client_metadata_url", None) - rejected_id = getattr(self.context.client_info, "client_id", None) - if cimd_url and rejected_id == cimd_url: - logger.warning( - "MCP OAuth '%s': authorization server rejected our Client ID Metadata Document (%s) " - "with invalid_client — falling back to dynamic client registration.", - self._hermes_server_name, cimd_url, - ) + if cimd_url and getattr(self.context.client_info, "client_id", None) == cimd_url: + logger.warning("MCP OAuth '%s': authorization server rejected our Client ID Metadata Document (%s) " + "with invalid_client — falling back to dynamic client registration.", + self._hermes_server_name, cimd_url) self.context.client_metadata_url = None if storage is not None: storage.mark_cimd_rejected() - if storage is not None: storage.poison_client_registration() # Drop the in-memory client so the SDK re-registers next flow. self.context.client_info = None self._initialized = False - except Exception as exc: # pragma: no cover — defensive, must not throw + except Exception as exc: # pragma: no cover — must not throw self._log_nonfatal("invalid_client detection", exc) async def async_auth_flow(self, request): # type: ignore[override] @@ -260,9 +227,8 @@ class _HermesRuntimeProviderMixin: self._log_nonfatal("pre-flow disk-watch", exc) # Bridge the bidirectional generator protocol by hand: httpx feeds responses back via - # ``auth_flow.asend(response)``. A naive ``async for item in inner: yield item`` - # DISCARDS those values, so the SDK's ``response = yield request`` sees None and - # crashes on ``response.status_code``. + # ``auth_flow.asend(response)``. A naive ``async for item in inner: yield item`` DISCARDS + # those values, so the SDK's ``response = yield request`` sees None and crashes. inner = super().async_auth_flow(request) resource_lock_released = False sent_access_token = None @@ -286,29 +252,22 @@ class _HermesRuntimeProviderMixin: # retry with that token instead of a duplicate OAuth transition from the stale # 401/403. tokens = self.context.current_tokens - if ( - getattr(incoming, "status_code", None) in (401, 403) - and self.context.is_token_valid() - and tokens is not None - and tokens.access_token != sent_access_token - ): + if (getattr(incoming, "status_code", None) in (401, 403) and self.context.is_token_valid() + and tokens is not None and tokens.access_token != sent_access_token): self._add_auth_header(request) await inner.aclose() retry_after_concurrent_auth = True break - # Sniff for a dead-client-registration signal (best-effort). await self._maybe_flag_poisoned_client(incoming) outgoing = await inner.asend(incoming) except StopAsyncIteration: - # Persist metadata discovered lazily in the 401 branch. - self._persist_oauth_metadata_if_changed() + self._persist_oauth_metadata_if_changed() # metadata discovered lazily in the 401 branch return finally: if resource_lock_released: # Balance the SDK's surrounding ``async with`` even when HTTPX cancels/closes # the flow mid-request. Shield only this local bookkeeping. import anyio - with anyio.CancelScope(shield=True): await self.context.lock.acquire() @@ -318,32 +277,14 @@ class _HermesRuntimeProviderMixin: return -def _make_hermes_provider_class() -> Optional[type]: - """Lazy-import the SDK base class and return our subclass (None if the SDK's OAuth module is - unavailable, so this module still imports).""" - try: - from mcp.client.auth.oauth2 import OAuthClientProvider - except ImportError: # pragma: no cover — SDK required in CI - return None - - class HermesMCPOAuthProvider(_HermesRuntimeProviderMixin, HermesProviderMixin, OAuthClientProvider): - """OAuthClientProvider with pre-flow disk-mtime reload: before every ``async_auth_flow`` - the manager checks whether the tokens file changed on disk and, if so, resets - ``_initialized`` so the next flow re-reads storage — making external refreshes visible - to a running session. Token-endpoint fixes come from ``HermesProviderMixin``.""" - - return HermesMCPOAuthProvider +# Cached at import time; None when the SDK's OAuth module is unavailable. +_HERMES_PROVIDER_CLS: Optional[type] = HermesMCPOAuthProvider if _SDK_BASES else None -# Cached at import time. Tested and used by :class:`MCPOAuthManager`. -_HERMES_PROVIDER_CLS: Optional[type] = _make_hermes_provider_class() - - -# -- Manager ----------------------------------------------------------------- class MCPOAuthManager: - """Single source of truth for per-server MCP OAuth state. Thread-safe: the ``_entries`` dict - is guarded by ``_entries_lock`` for get-or-create semantics; per-entry state is guarded by - the entry's own ``asyncio.Lock`` (used from the MCP event loop thread).""" + """Single source of truth for per-server MCP OAuth state. Thread-safe: ``_entries`` is + guarded by ``_entries_lock`` for get-or-create; per-entry state by the entry's own + ``asyncio.Lock`` (used from the MCP event loop thread).""" def __init__(self) -> None: self._entries: dict[tuple[str, str], _ProviderEntry] = {} @@ -352,12 +293,10 @@ class MCPOAuthManager: # mid-run and leave `await pending` hanging forever. self._inflight_tasks: set[asyncio.Task] = set() - # -- Provider construction / caching -- def get_or_build_provider(self, server_name: str, server_url: str, oauth_config: Optional[dict]) -> Optional[Any]: - """Return a cached OAuth provider for ``server_name`` or build one. Idempotent: repeat - calls with the same name return the same instance; if ``server_url`` changes for a - given name the cached entry is discarded and a fresh provider is built. None if the MCP - SDK's OAuth support is unavailable.""" + """Cached OAuth provider for ``server_name``, built on first use. If ``server_url`` + changes for a name the cached entry is discarded and rebuilt. None if the MCP SDK's + OAuth support is unavailable.""" key = self._key(server_name) with self._entries_lock: entry = self._entries.get(key) @@ -365,8 +304,7 @@ class MCPOAuthManager: logger.info("MCP OAuth '%s': URL changed from %s to %s, discarding cache", server_name, entry.server_url, server_url) entry = None if entry is None: - entry = _ProviderEntry(server_url=server_url, oauth_config=oauth_config) - self._entries[key] = entry + entry = self._entries[key] = _ProviderEntry(server_url=server_url, oauth_config=oauth_config) if entry.provider is None: entry.provider = self._build_provider(server_name, entry) if entry.provider is not None: @@ -376,13 +314,11 @@ class MCPOAuthManager: @staticmethod def _key(server_name: str, hermes_home: str | Path | None = None) -> tuple[str, str]: from hermes_constants import get_hermes_home - home = Path(hermes_home) if hermes_home is not None else get_hermes_home() return (str(home.expanduser().resolve(strict=False)), server_name) def _build_provider(self, server_name: str, entry: _ProviderEntry) -> Optional[Any]: - """Build a :class:`HermesMCPOAuthProvider` from the shared ``tools.mcp_oauth`` helpers; - None if the SDK's OAuth support is unavailable.""" + """Build a ``HermesMCPOAuthProvider``; None if the SDK's OAuth support is unavailable.""" if _HERMES_PROVIDER_CLS is None: logger.warning("MCP OAuth '%s': SDK auth module unavailable", server_name) return None @@ -390,34 +326,27 @@ class MCPOAuthManager: from tools.mcp_dashboard_oauth import get_dashboard_oauth_flow from tools.mcp_oauth import _OAUTH_AVAILABLE, OAuthNonInteractiveError, _is_interactive from tools.mcp_oauth_provider import build_provider_kwargs, prepare_oauth_config - if not _OAUTH_AVAILABLE: return None cfg, storage = prepare_oauth_config(server_name, entry.server_url, entry.oauth_config) if get_dashboard_oauth_flow() is None and not _is_interactive() and not storage.has_cached_tokens(): raise OAuthNonInteractiveError( f"MCP OAuth for '{server_name}': non-interactive environment and no cached tokens found. " - f"Run `hermes mcp login {server_name}` interactively first to complete initial authorization." - ) + f"Run `hermes mcp login {server_name}` interactively first to complete initial authorization.") return _HERMES_PROVIDER_CLS( - server_name=server_name, - preregistered=bool(cfg.get("client_id")), - server_url=entry.server_url, - **build_provider_kwargs(cfg, storage, ssh_proxy_hint=False), - ) + server_name=server_name, preregistered=bool(cfg.get("client_id")), server_url=entry.server_url, + **build_provider_kwargs(cfg, storage, ssh_proxy_hint=False)) def remove(self, server_name: str, *, hermes_home: str | Path | None = None) -> _ProviderEntry | None: - """Evict the provider from cache AND delete tokens from disk (``hermes mcp remove`` and, - indirectly, ``hermes mcp login`` during forced re-auth).""" + """Evict the provider from cache AND delete tokens from disk (``hermes mcp remove``, + and ``hermes mcp login`` during forced re-auth).""" entry = self.evict(server_name, hermes_home=hermes_home) from tools.mcp_oauth import remove_oauth_tokens remove_oauth_tokens(server_name, hermes_home=hermes_home) logger.info("MCP OAuth '%s': evicted from cache and removed from disk", server_name) return entry - def restore_entry( - self, server_name: str, entry: _ProviderEntry | None, *, hermes_home: str | Path | None = None - ) -> None: + def restore_entry(self, server_name: str, entry: _ProviderEntry | None, *, hermes_home: str | Path | None = None) -> None: """Restore a provider entry removed for a failed reauthorization.""" if entry is None: return @@ -429,13 +358,10 @@ class MCPOAuthManager: with self._entries_lock: return self._entries.pop(self._key(server_name, hermes_home), None) - # -- Disk watch -- async def invalidate_if_disk_changed(self, server_name: str, *, hermes_home: str | Path | None = None) -> bool: """Force the SDK provider to reload when the tokens file mtime changed; True if - invalidated. This is the external-refresh fix: a cron job writes fresh tokens and the - next tool call picks them up.""" + invalidated. A cron job writes fresh tokens and the next tool call picks them up.""" from tools.mcp_oauth import _get_token_dir, _safe_filename - entry = self._entries.get(self._key(server_name, hermes_home)) if entry is None or entry.provider is None: return False @@ -447,8 +373,7 @@ class MCPOAuthManager: return False if mtime_ns == entry.last_mtime_ns: return False - old = entry.last_mtime_ns - entry.last_mtime_ns = mtime_ns + old, entry.last_mtime_ns = entry.last_mtime_ns, mtime_ns # `_initialized` is private SDK API but stable across the versions we pin # (>=1.26.0); resetting it forces a reload. if hasattr(entry.provider, "_initialized"): @@ -456,22 +381,18 @@ class MCPOAuthManager: logger.info("MCP OAuth '%s': tokens file changed (mtime %d -> %d), forcing reload", server_name, old, mtime_ns) return True - # -- 401 handler (dedup'd) -- async def _recover_401(self, server_name: str, entry: _ProviderEntry, key: str, pending: asyncio.Future) -> None: """Single recovery attempt behind *pending*; always clears the dedup slot.""" try: - # Step 1: Did disk change? Picks up external refresh. - if await self.invalidate_if_disk_changed(server_name): - if not pending.done(): - pending.set_result(True) - return - # Step 2: No disk change — if the SDK can refresh in place, let the caller retry - # (the httpx.Auth flow refreshes on the next request). - can_refresh_fn = getattr(getattr(entry.provider, "context", None), "can_refresh_token", None) - try: - can_refresh = bool(can_refresh_fn()) if callable(can_refresh_fn) else False - except Exception: - can_refresh = False + # Disk changed (external refresh)? Else: if the SDK can refresh in place, let the + # caller retry (the httpx.Auth flow refreshes on the next request). + can_refresh = True + if not await self.invalidate_if_disk_changed(server_name): + can_refresh_fn = getattr(getattr(entry.provider, "context", None), "can_refresh_token", None) + try: + can_refresh = bool(can_refresh_fn()) if callable(can_refresh_fn) else False + except Exception: + can_refresh = False if not pending.done(): pending.set_result(can_refresh) except Exception as exc: # pragma: no cover — defensive @@ -483,10 +404,10 @@ class MCPOAuthManager: async def handle_401(self, server_name: str, failed_access_token: Optional[str] = None) -> bool: """Handle a 401 from a tool call, deduplicated across concurrent callers. True: a - (possibly new) access token is available — caller should reconnect and retry. False: no - recovery path — caller should surface a ``needs_reauth`` error so the model stops - hallucinating manual refreshes. Thundering-herd protection: N concurrent 401s with the - same ``failed_access_token`` fire one recovery attempt; the rest await its future.""" + (possibly new) access token is available — reconnect and retry. False: no recovery + path — surface a ``needs_reauth`` error so the model stops hallucinating manual + refreshes. N concurrent 401s with the same ``failed_access_token`` fire one recovery + attempt; the rest await its future.""" entry = self._entries.get(self._key(server_name)) if entry is None or entry.provider is None: return False @@ -495,8 +416,7 @@ class MCPOAuthManager: async with entry.lock: pending = entry.pending_401.get(key) if pending is None: - pending = loop.create_future() - entry.pending_401[key] = pending + pending = entry.pending_401[key] = loop.create_future() task = asyncio.create_task(self._recover_401(server_name, entry, key, pending)) self._inflight_tasks.add(task) task.add_done_callback(self._inflight_tasks.discard) @@ -507,7 +427,6 @@ class MCPOAuthManager: return False -# -- Module-level singleton --------------------------------------------------- _MANAGER: Optional[MCPOAuthManager] = None _MANAGER_LOCK = threading.Lock() diff --git a/tools/mcp_oauth_provider.py b/tools/mcp_oauth_provider.py index 60b34b8607..b8d3cea01e 100644 --- a/tools/mcp_oauth_provider.py +++ b/tools/mcp_oauth_provider.py @@ -1,10 +1,9 @@ """Shared ``OAuthClientProvider`` customizations for Hermes MCP OAuth. -Two code paths build an SDK provider — ``tools.mcp_oauth.build_oauth_auth`` -(legacy public API) and ``tools.mcp_oauth_manager.MCPOAuthManager`` — and both -need the same real-world fixes and the same config → constructor-kwargs -plumbing. This module holds that shared core once; the origin modules keep -their own subclass (logger name, disk-watch hooks) on top of it. +Two code paths build an SDK provider — ``tools.mcp_oauth.build_oauth_auth`` (legacy public +API) and ``tools.mcp_oauth_manager.MCPOAuthManager`` — and both need the same real-world +fixes and config → constructor-kwargs plumbing. This module holds that core once; the origin +modules keep their own subclass (logger name, disk-watch hooks) on top of it. """ from __future__ import annotations @@ -22,17 +21,16 @@ class HermesProviderMixin: """Token-endpoint fixes layered over the SDK's ``OAuthClientProvider``. - Supabase-style dynamic registration returns a ``client_secret`` but omits - ``token_endpoint_auth_method``; the SDK then treats the client as public, - omits the secret, and the token endpoint rejects the exchange (looping the - browser authorization page). Coerce the in-memory client info to - ``client_secret_post`` right before token and refresh requests. - - ``token_user_agent`` (``oauth.user_agent``) is stamped onto token-endpoint - requests only — some authorization servers and WAFs reject httpx's default. - - Any 2xx token/refresh response is accepted, and token bodies never leak - into exception text or log output. + ``token_endpoint_auth_method``; the SDK then treats the client as public, omits the + secret, and the token endpoint rejects the exchange (looping the browser authorization + page). Coerce the in-memory client info to ``client_secret_post`` before token requests. + - ``token_user_agent`` (``oauth.user_agent``) is stamped onto token-endpoint requests + only — some authorization servers and WAFs reject httpx's default. + - Any 2xx token/refresh response is accepted, and token bodies never leak into + exception text or log output. - Must precede the SDK class in the MRO. Subclasses set ``_hermes_logger`` so - warnings keep their origin module's logger name. + Must precede the SDK class in the MRO. Subclasses set ``_hermes_logger`` so warnings + keep their origin module's logger name. """ _hermes_logger: logging.Logger = logger @@ -49,8 +47,8 @@ class HermesProviderMixin: return request def _coerce_client_secret_post(self) -> None: - """Same rule as ``HermesTokenStorage._coerce_secret_auth_method``, applied - to the in-memory client info right before a token-endpoint request.""" + """Same rule as ``HermesTokenStorage._coerce_secret_auth_method``, applied to the + in-memory client info BEFORE the SDK builds a token-endpoint request from it.""" info = self.context.client_info if not info: return @@ -110,11 +108,9 @@ class HermesProviderMixin: def prepare_oauth_config(server_name: str, server_url: str, oauth_config: dict | None) -> tuple[dict, "HermesTokenStorage"]: - """Copy the ``oauth:`` block, apply provider defaults, open its token storage. - - The copy matters: later steps record ``_resolved_port`` / ``_cimd_url`` in - the dict, which must never leak back into the caller's config. - """ + """Copy the ``oauth:`` block, apply provider defaults, open its token storage. The copy + matters: later steps record ``_resolved_port`` / ``_cimd_url`` in the dict, which must + never leak back into the caller's config.""" from tools import mcp_oauth as mo cfg = dict(oauth_config or {}) @@ -125,12 +121,10 @@ def prepare_oauth_config(server_name: str, server_url: str, oauth_config: dict | def build_provider_kwargs(cfg: dict, storage: "HermesTokenStorage", *, ssh_proxy_hint: bool) -> dict[str, Any]: """Resolve the callback port and return the shared provider constructor kwargs. - Runs the port → client-metadata → pre-registration sequence (order matters: - metadata needs the resolved port, pre-registration needs the metadata). - ``ssh_proxy_hint`` lets the redirect handler tailor its remote-session hint - to a configured proxy ``redirect_uri``. Helpers are looked up on - ``tools.mcp_oauth`` at call time so tests can patch them there. - """ + Runs port → client-metadata → pre-registration (order matters: metadata needs the + resolved port, pre-registration needs the metadata). ``ssh_proxy_hint`` lets the redirect + handler tailor its remote-session hint to a configured proxy ``redirect_uri``. Helpers are + looked up on ``tools.mcp_oauth`` at call time so tests can patch them there.""" from tools import mcp_oauth as mo port = mo._configure_callback_port(cfg, storage) @@ -143,9 +137,6 @@ def build_provider_kwargs(cfg: dict, storage: "HermesTokenStorage", *, ssh_proxy "redirect_handler": mo._make_redirect_handler(port, redirect_uri=redirect_uri), # mcp 2.0 dropped OAuthClientProvider's own `timeout`; the configured # `oauth.timeout` bounds the callback waiter's poll loop instead. - "callback_handler": mo._make_callback_waiter( - port, cfg.get("_cimd_url"), timeout=float(cfg.get("timeout", 300)) - ), + "callback_handler": mo._make_callback_waiter(port, cfg.get("_cimd_url"), timeout=float(cfg.get("timeout", 300))), "token_user_agent": mo.token_request_user_agent(cfg), - **mo.cimd_provider_kwargs(cfg), - } + **mo.cimd_provider_kwargs(cfg)} diff --git a/tools/mcp_stdio_watchdog.py b/tools/mcp_stdio_watchdog.py index e7b29a105b..7e20d2c6f7 100644 --- a/tools/mcp_stdio_watchdog.py +++ b/tools/mcp_stdio_watchdog.py @@ -1,21 +1,14 @@ #!/usr/bin/env python3 """Parent-death watchdog supervisor for stdio MCP subprocesses. -If Hermes dies hard (kill -9, crash, force-quit) its graceful teardown never -runs and the stdio child plus its descendants are orphaned (macOS has no -``PR_SET_PDEATHSIG``); piled-up orphans then race the legitimate new connection -for the same upstream session. So the MCP command is spawned via this -supervisor, which (1) runs the real command as its own child in a new process -group so the whole tree can be killpg'd, (2) passes stdin/stdout/stderr straight -through — the MCP stdio protocol talks over those pipes, so this must be a -no-op relay, not a proxy — and (3) polls ``getppid()`` against the recorded -parent PID and, once the parent is gone, SIGTERMs the child's group, waits, -then SIGKILLs. Standard-library only so it starts fast and cannot itself leak. - -Usage (see ``_wrap_command_with_watchdog``):: - - python3 -m tools.mcp_stdio_watchdog \\ - --ppid -- ... +If Hermes dies hard (kill -9, crash) its graceful teardown never runs and the stdio child +plus its descendants are orphaned (macOS has no ``PR_SET_PDEATHSIG``); piled-up orphans then +race the new connection for the same upstream session. So the MCP command is spawned via this +supervisor, which (1) runs the real command in a new process group so the whole tree can be +killpg'd, (2) passes stdin/stdout/stderr straight through — the MCP stdio protocol talks over +those pipes, so this is a no-op relay, not a proxy — and (3) polls ``getppid()`` and, once the +parent is gone, SIGTERMs the child's group, waits, then SIGKILLs. Stdlib only so it starts +fast and cannot itself leak. Usage: ``mcp_stdio_watchdog.py --ppid -- `` """ from __future__ import annotations @@ -38,9 +31,8 @@ def _is_orphaned(original_ppid: int, getppid=os.getppid) -> bool: def _terminate_process_group(proc: subprocess.Popen) -> None: - """Best-effort SIGTERM-then-SIGKILL of the child's process group; guards the - POSIX-only primitives so an accidental Windows run degrades to a plain child - kill instead of AttributeError.""" + """Best-effort SIGTERM-then-SIGKILL of the child's process group; guards the POSIX-only + primitives so an accidental Windows run degrades to a plain child kill.""" killpg = getattr(os, "killpg", None) if killpg is None: # windows-footgun: ok — non-POSIX fallback try: @@ -87,13 +79,13 @@ def main(argv: list[str] | None = None) -> int: print("mcp_stdio_watchdog: no command given after '--'", file=sys.stderr) return 2 - # New process group: killpg() reaches the whole tree the real command may - # spawn without touching our own group or the original parent's. + # New process group: killpg() reaches the whole tree the real command may spawn without + # touching our own group or the original parent's. proc = subprocess.Popen(real_argv, stdin=sys.stdin, stdout=sys.stdout, stderr=sys.stderr, start_new_session=True) - # The server lives in its OWN group, so the parent's shutdown killpg of *our* - # group no longer reaches it: forward SIGTERM/SIGINT to the child's group so - # graceful teardown still kills a wedged server that ignores stdin EOF. + # The server lives in its OWN group, so the parent's shutdown killpg of *our* group no + # longer reaches it: forward SIGTERM/SIGINT to the child's group so graceful teardown + # still kills a wedged server that ignores stdin EOF. def _forward_shutdown(signum, frame): # noqa: ARG001 _terminate_process_group(proc) sys.exit(128 + signum) diff --git a/tools/mcp_tool_agent.py b/tools/mcp_tool_agent.py index 697cd45118..ac3407e4e9 100644 --- a/tools/mcp_tool_agent.py +++ b/tools/mcp_tool_agent.py @@ -25,9 +25,9 @@ def _agent_tool_defs(agent) -> list: def _resolve_refresh_toolsets(agent, enabled_override, disabled_override): - """Explicit reloads pass freshly-resolved toolsets (so a server just ENABLED - in config is picked up) and the agent's selection is updated to match; - automatic paths pass nothing and reuse the build-time selection.""" + """Explicit reloads pass freshly-resolved toolsets (so a server just ENABLED in config is + picked up) and the agent's selection is updated to match; automatic paths pass nothing + and reuse the build-time selection.""" enabled = getattr(agent, "enabled_toolsets", None) disabled = getattr(agent, "disabled_toolsets", None) if enabled_override is not None or disabled_override is not None: @@ -39,8 +39,8 @@ def _resolve_refresh_toolsets(agent, enabled_override, disabled_override): def _tool_defs_content_changed(agent, new_defs: list) -> bool: - """Byte-level diff of the serialized tool arrays (dynamic schemas change - CONTENT under stable names); False if either side fails to serialize.""" + """Byte-level diff of the serialized tool arrays (dynamic schemas change CONTENT under + stable names); False if either side fails to serialize.""" try: dump = lambda defs: json.dumps(defs, sort_keys=True, separators=(",", ":"), default=str) # noqa: E731 return dump(_agent_tool_defs(agent)) != dump(new_defs) @@ -52,13 +52,11 @@ def _publish_tool_snapshot( agent, new_defs: list, new_names: set, *, snapshot_generation: int, staged_engine_names: set, content_aware: bool, prefix_registered: Optional[set], ) -> Optional[set]: - """Single atomic read-diff-publish under ``_agent_tools_lock`` so ``added`` - matches what was published and a stale (older-generation) rebuild can't - overwrite a newer one. Returns the added names, or None when nothing was - published (unchanged, or a newer snapshot already won).""" + """Single atomic read-diff-publish under ``_agent_tools_lock`` so ``added`` matches what + was published and a stale (older-generation) rebuild can't overwrite a newer one. Returns + the added names, or None when nothing was published (unchanged, or a newer snapshot won).""" with _agent_tools_lock: - # Tolerate an agent that never set the generation (or a non-int mock) - # rather than failing the whole refresh on the comparison. + # Tolerate an agent that never set the generation (or a non-int mock). published_gen = getattr(agent, "_tool_snapshot_generation", -1) published_gen = published_gen if isinstance(published_gen, int) else -1 if snapshot_generation < published_gen: @@ -84,53 +82,41 @@ def _publish_tool_snapshot( def refresh_agent_mcp_tools( - agent, - *, - enabled_override=None, - disabled_override=None, - quiet_mode: bool = True, - content_aware: bool = False, - preserve_prefix: bool = False, -) -> set: - """Re-derive an already-built agent's tool snapshot from the live registry; - returns the newly-added tool names (empty when unchanged). + agent, *, enabled_override=None, disabled_override=None, quiet_mode: bool = True, + content_aware: bool = False, preserve_prefix: bool = False) -> set: + """Re-derive an already-built agent's tool snapshot from the live registry; returns the + newly-added tool names (empty when unchanged). - The agent snapshots ``agent.tools`` once at build time, so servers that - connect later (slow OAuth server, ``/reload-mcp``) are invisible until - rebuilt. Single shared rebuild for the TUI RPC, gateway reload, late-binding - thread and between-turns refresh: respects the agent's toolset filter, diffs - by tool NAME (a count compare misses an equal-size swap), re-injects the - memory-provider / context-engine (``lcm_*``) tools ``agent_init`` appends - after ``get_tool_definitions`` (a naive rebuild would drop them), and - publishes ``(tools, valid_tool_names)`` together under ``_agent_tools_lock``. + The agent snapshots ``agent.tools`` once at build time, so servers that connect later + (slow OAuth server, ``/reload-mcp``) are invisible until rebuilt. Single shared rebuild for + the TUI RPC, gateway reload, late-binding thread and between-turns refresh: respects the + agent's toolset filter, diffs by tool NAME (a count compare misses an equal-size swap), + re-injects the memory-provider / context-engine (``lcm_*``) tools ``agent_init`` appends + after ``get_tool_definitions``, and publishes ``(tools, valid_tool_names)`` together + under ``_agent_tools_lock``. - ``preserve_prefix`` is for rebuilds inside a live conversation, where the - tool array is a cached request prefix and any moved byte re-prefills the - whole history: existing tools keep their slot (schemas still refresh), a - still-registered tool whose ``check_fn`` merely flapped is carried forward, - a deregistered tool is dropped, new tools append at the tail. Carrying an - unavailable tool forward is safe: ``check_fn`` gates exposure, never - invocation. The caller owns the prompt-cache contract (turn-boundary policy - differs per caller).""" + ``preserve_prefix`` is for rebuilds inside a live conversation, where the tool array is a + cached request prefix and any moved byte re-prefills the whole history: existing tools + keep their slot (schemas still refresh), a still-registered tool whose ``check_fn`` merely + flapped is carried forward (safe: ``check_fn`` gates exposure, never invocation), a + deregistered tool is dropped, new tools append at the tail. The caller owns the + prompt-cache contract (turn-boundary policy differs per caller).""" from model_tools import get_tool_definitions from tools.registry import registry enabled, disabled = _resolve_refresh_toolsets(agent, enabled_override, disabled_override) - # Capture the registry generation BEFORE the slow get_tool_definitions call; - # at publish time a slower caller holding an OLDER set must not clobber a - # newer set another caller already published. + # Capture the registry generation BEFORE the slow get_tool_definitions call; at publish + # time a slower caller holding an OLDER set must not clobber a newer set already published. snapshot_generation = registry._generation - # Computed OUTSIDE the lock (can be slow); diff + publish happen together in - # one critical section so concurrent callers can't torn-publish. - new_defs = list( - get_tool_definitions(enabled_toolsets=enabled, disabled_toolsets=disabled, quiet_mode=quiet_mode) or [] - ) + # Computed OUTSIDE the lock (can be slow); diff + publish happen together in one critical + # section so concurrent callers can't torn-publish. + new_defs = list(get_tool_definitions(enabled_toolsets=enabled, disabled_toolsets=disabled, quiet_mode=quiet_mode) or []) new_names = {_def_name(t) for t in new_defs} - # Re-append the post-build families on LOCALS only; live agent attributes - # are untouched until the single atomic publish. + # Re-append the post-build families on LOCALS only; live agent attributes are untouched + # until the single atomic publish. staged_engine_names = _core._reinject_post_build_tools(agent, new_defs, new_names) - # Registry membership is read OUTSIDE ``_agent_tools_lock``: taking - # ``registry._lock`` under the tools lock would be the first nesting of the two. + # Registry membership is read OUTSIDE ``_agent_tools_lock``: taking ``registry._lock`` + # under the tools lock would be the first nesting of the two. prefix_registered: Optional[set] = None if preserve_prefix: try: @@ -139,20 +125,19 @@ def refresh_agent_mcp_tools( pass # fail open to the plain rebuild added = _publish_tool_snapshot( agent, new_defs, new_names, snapshot_generation=snapshot_generation, - staged_engine_names=staged_engine_names, content_aware=content_aware, prefix_registered=prefix_registered, - ) + staged_engine_names=staged_engine_names, content_aware=content_aware, prefix_registered=prefix_registered) if added is None: return set() - # Re-pin the session's tool order so a rebuild-for-existing-session - # (gateway agent-cache eviction) restores exactly these names. + # Re-pin the session's tool order so a rebuild-for-existing-session (gateway agent-cache + # eviction) restores exactly these names. persist_agent_tool_names(agent) return added def reprobe_tool_availability() -> None: - """Explicit ``/reload-mcp`` hatch out of the tools[] freeze: drop the - ``check_fn`` verdict cache AND the ``get_tool_definitions`` memo (keyed on - registry generation, so it would otherwise replay the stale verdicts).""" + """Explicit ``/reload-mcp`` hatch out of the tools[] freeze: drop the ``check_fn`` verdict + cache AND the ``get_tool_definitions`` memo (keyed on registry generation, so it would + otherwise replay the stale verdicts).""" from model_tools import _clear_tool_defs_cache from tools.registry import invalidate_check_fn_cache @@ -173,12 +158,11 @@ def persist_agent_tool_names(agent) -> None: def restore_agent_tool_prefix(agent, saved_names: list) -> bool: - """Fold a freshly built agent's ``tools`` onto the session's saved order; - True if changed. The gateway rebuilds a NEW AIAgent for an existing session - after agent-cache eviction, with no predecessor to preserve, so the saved - name list stands in: a saved tool still registered but failing its probe is - carried forward from the registry schema, a deregistered one is dropped, new - tools append at the tail (same rule as ``_merge_preserving_prefix``).""" + """Fold a freshly built agent's ``tools`` onto the session's saved order; True if changed. + The gateway rebuilds a NEW AIAgent for an existing session after agent-cache eviction, with + no predecessor to preserve, so the saved name list stands in (same merge rule as + ``_merge_preserving_prefix``; a saved tool still registered but failing its probe is + carried forward from the registry schema).""" if not saved_names: return False from tools.registry import registry @@ -206,11 +190,10 @@ def restore_agent_tool_prefix(agent, saved_names: list) -> bool: def _merge_preserving_prefix(current_defs: list, new_defs: list, registered_names: set) -> tuple[list, set]: - """Fold a fresh tool snapshot into a live one without moving existing bytes. - Ordered by ``current_defs`` (the cached request prefix): a name in both keeps - its slot but takes the fresh schema; a name only in the live list is kept if - still registered (``check_fn`` flapped), else dropped; a name only in the - fresh list is appended at the tail.""" + """Fold a fresh tool snapshot into a live one without moving existing bytes. Ordered by + ``current_defs`` (the cached request prefix): a name in both keeps its slot but takes the + fresh schema; a name only in the live list is kept if still registered (``check_fn`` + flapped), else dropped; a name only in the fresh list is appended at the tail.""" fresh = {_def_name(entry): entry for entry in new_defs if _def_name(entry)} merged = [] for entry in current_defs: @@ -225,11 +208,10 @@ def _merge_preserving_prefix(current_defs: list, new_defs: list, registered_name def _reinject_post_build_tools(agent, tools_list: list, name_set: set) -> set: - """Append memory-provider and context-engine tools onto the caller's staged - ``tools_list`` / ``name_set`` (never the live agent attributes), mirroring - ``agent_init``'s post-build injection. Idempotent and fail-soft. Returns the - context-engine routing names THIS rebuild appended: a name already owned by - a registry/plugin tool is not claimed, matching agent_init.""" + """Append memory-provider and context-engine tools onto the caller's staged ``tools_list`` + / ``name_set`` (never the live agent attributes), mirroring ``agent_init``'s post-build + injection. Idempotent and fail-soft. Returns the context-engine routing names THIS rebuild + appended: a name already owned by a registry/plugin tool is not claimed, matching agent_init.""" def _add(schema) -> bool: name = schema.get("name", "") if isinstance(schema, dict) else "" if not name or name in name_set: @@ -247,19 +229,17 @@ def _reinject_post_build_tools(agent, tools_list: list, name_set: set) -> set: try: get_mem_schemas = _schema_getter("_memory_manager", "get_all_tool_schemas") if get_mem_schemas is not None: - # Same toolset gate inject_memory_provider_tools uses. - from agent.memory_manager import memory_provider_tools_enabled + from agent.memory_manager import memory_provider_tools_enabled # same gate inject_memory_provider_tools uses if memory_provider_tools_enabled( - enabled, getattr(agent, "disabled_toolsets", None), memory_tool_present="memory" in name_set, - ): + enabled, getattr(agent, "disabled_toolsets", None), memory_tool_present="memory" in name_set): for schema in get_mem_schemas(): _add(schema) except Exception: logger.debug("Memory-provider tool re-injection skipped", exc_info=True) - # The `context_engine` toolset is intentionally empty, so lcm_* tools exist - # only via this append. Honor the enabled_toolsets gate agent_init uses, or a - # restricted-toolset platform would re-leak tools the build excluded. + # The `context_engine` toolset is intentionally empty, so lcm_* tools exist only via this + # append. Honor the enabled_toolsets gate agent_init uses, or a restricted-toolset platform + # would re-leak tools the build excluded. staged_engine_names: set = set() try: get_schemas = _schema_getter("context_compressor", "get_tool_schemas") diff --git a/tools/mcp_tool_config.py b/tools/mcp_tool_config.py index feede5259d..72b1e98d9d 100644 --- a/tools/mcp_tool_config.py +++ b/tools/mcp_tool_config.py @@ -19,9 +19,9 @@ _mcp_stderr_log_lock = threading.Lock() def _get_mcp_stderr_log() -> Any: - """Shared append-mode handle for MCP subprocess stderr, opened once per - process. Must expose a real fd (``fileno()``) because asyncio wires the - child's stderr directly to it. Falls back to ``/dev/null``, then real stderr.""" + """Shared append-mode handle for MCP subprocess stderr, opened once per process. Must + expose a real fd (``fileno()``) because asyncio wires the child's stderr directly to it. + Falls back to ``/dev/null``, then real stderr.""" global _mcp_stderr_log_fh with _mcp_stderr_log_lock: if _mcp_stderr_log_fh is not None: @@ -30,8 +30,7 @@ def _get_mcp_stderr_log() -> Any: from hermes_constants import get_hermes_home log_dir = get_hermes_home() / "logs" log_dir.mkdir(parents=True, exist_ok=True) - # Line-buffered so output lands promptly; errors="replace" tolerates - # garbled binary from misbehaving servers. + # Line-buffered so output lands promptly; errors="replace" tolerates garbled binary. fh = open(log_dir / "mcp-stderr.log", "a", encoding="utf-8", errors="replace", buffering=1) fh.fileno() # confirm a real fd before committing _mcp_stderr_log_fh = fh @@ -45,8 +44,8 @@ def _get_mcp_stderr_log() -> Any: def _write_stderr_log_header(server_name: str) -> None: - """Write a session marker so operators can find each server's output in the - shared log without per-line prefixes (which would need a pipe + reader thread).""" + """Session marker so operators can find each server's output in the shared log + (per-line prefixes would need a pipe + reader thread).""" fh = _core._get_mcp_stderr_log() try: ts = datetime.now().strftime("%Y-%m-%d %H:%M:%S") @@ -67,16 +66,15 @@ _SAFE_ENV_KEYS_CASE_INSENSITIVE = frozenset({ "LOCALAPPDATA", "NUMBER_OF_PROCESSORS", "OS", "PATHEXT", "PROCESSOR_ARCHITECTURE", "PROGRAMDATA", "PROGRAMFILES", "PROGRAMFILES(X86)", "PROGRAMW6432", "PUBLIC", "SYSTEMDRIVE", "SYSTEMROOT", "TEMP", "TMP", "USERDOMAIN", "USERNAME", - "USERPROFILE", "WINDIR", -}) + "USERPROFILE", "WINDIR"}) # ${VAR_NAME} interpolation; any non-} chars allowed so MY-VAR / my.var work. _ENV_VAR_PATTERN = re.compile(r"\$\{([^}]+)\}") def _workspace_folder() -> str: - """Absolute workspace root for ``${workspaceFolder}``: the session's - authoritative root (terminal cwd / task override / $TERMINAL_CWD), else cwd.""" + """Absolute workspace root for ``${workspaceFolder}``: the session's authoritative root + (terminal cwd / task override / $TERMINAL_CWD), else cwd.""" try: from tools.file_tools import _authoritative_workspace_root @@ -99,34 +97,21 @@ _CONTEXT_VAR_RESOLVERS = { "workspaceFolder": lambda: _core._workspace_folder(), "workspaceFolderBasename": _workspace_basename, "pathSeparator": lambda: os.sep, - "/": lambda: os.sep, -} - - -def _context_var_value(ref: str) -> Optional[str]: - """Resolve a Cursor context var; None for anything else so it falls through - to env-var lookup.""" - resolver = _CONTEXT_VAR_RESOLVERS.get(ref) - return resolver() if resolver else None + "/": lambda: os.sep} def _build_safe_env(user_env: Optional[dict]) -> dict: - """Filtered env for stdio subprocesses so API keys/tokens don't leak: only - the safe baseline keys, ``XDG_*``, vars injected by an external secret - source (users configured that backend precisely so subprocesses can consume - them), plus the server config's own ``env``.""" + """Filtered env for stdio subprocesses so API keys/tokens don't leak: the safe baseline + keys, ``XDG_*``, vars injected by an external secret source (users configured that backend + precisely so subprocesses can consume them), plus the server config's own ``env``.""" try: from hermes_cli.env_loader import get_secret_source except Exception: # pragma: no cover — early bootstrap/import fallback get_secret_source = None env = { - key: value - for key, value in os.environ.items() - if key in _SAFE_ENV_KEYS - or key.upper() in _SAFE_ENV_KEYS_CASE_INSENSITIVE - or key.startswith("XDG_") - or (get_secret_source is not None and get_secret_source(key)) - } + key: value for key, value in os.environ.items() + if key in _SAFE_ENV_KEYS or key.upper() in _SAFE_ENV_KEYS_CASE_INSENSITIVE + or key.startswith("XDG_") or (get_secret_source is not None and get_secret_source(key))} if user_env: env.update(user_env) return env @@ -150,19 +135,17 @@ def _which_with_config_pathext(command: str, path_arg, env: dict): def _node_fallback(command: str) -> str: - """Well-known Node install locations for bare ``npx``/``npm``/``node`` when - PATH lookup failed; returns *command* unchanged when none is executable.""" + """Well-known Node install locations for bare ``npx``/``npm``/``node`` when PATH lookup + failed; *command* unchanged when none is executable.""" home = os.path.expanduser("~") hermes_home = os.path.expanduser(os.getenv("HERMES_HOME", os.path.join(home, ".hermes"))) candidates = [ os.path.join(hermes_home, "node", "bin", command), os.path.join(home, ".local", "bin", command), - # Canonical Node location for from-source Linux builds, the Hermes Docker - # image and Intel Homebrew. Needed when a user's hand-authored env.PATH - # omits it: npx's shebang re-execs /usr/bin/env node, so a symlink - # workaround fails one layer deeper. - os.path.join(os.sep, "usr", "local", "bin", command), - ] + # Canonical Node location for from-source Linux builds, the Hermes Docker image and + # Intel Homebrew. Needed when a hand-authored env.PATH omits it: npx's shebang re-execs + # /usr/bin/env node, so a symlink workaround fails one layer deeper. + os.path.join(os.sep, "usr", "local", "bin", command)] for candidate in candidates: if os.path.isfile(candidate) and os.access(candidate, os.X_OK): return candidate @@ -192,10 +175,9 @@ def _resolve_stdio_command(command: str, env: dict) -> tuple[str, dict]: def _wrap_command_with_watchdog(command: str, args: list) -> tuple[str, list]: - """Wrap a stdio command in the parent-death watchdog (POSIX only — it relies - on process groups, same scope as the killpg-based orphan cleanup; the - watchdog polls ``getppid()`` against our PID). Unchanged on non-POSIX or if - the PID cannot be read — watchdog bookkeeping must never block a connection.""" + """Wrap a stdio command in the parent-death watchdog (POSIX only — it relies on process + groups, same scope as the killpg-based orphan cleanup). Unchanged on non-POSIX or if the + PID cannot be read — watchdog bookkeeping must never block a connection.""" if os.name != "posix": return command, args try: @@ -207,18 +189,17 @@ def _wrap_command_with_watchdog(command: str, args: list) -> tuple[str, list]: def _interpolate_env_vars(value): - """Recursively resolve ``${VAR}`` / Cursor ``${env:VAR}`` placeholders plus - the Cursor context vars (``_context_var_value``). Env refs resolve from the - active profile's secret scope when multiplexing (so ``${API_KEY}`` picks up - the routed profile's value, not another profile's in ``os.environ``). Unset - vars keep the literal placeholder.""" + """Recursively resolve ``${VAR}`` / Cursor ``${env:VAR}`` placeholders plus the Cursor + context vars. Env refs resolve from the active profile's secret scope when multiplexing + (so ``${API_KEY}`` picks up the routed profile's value, not another profile's in + ``os.environ``). Unset vars keep the literal placeholder.""" from agent.secret_scope import get_secret as _get_secret if isinstance(value, str): def _replace(m): - ctx = _context_var_value(m.group(1).strip()) - if ctx is not None: - return ctx + resolver = _CONTEXT_VAR_RESOLVERS.get(m.group(1).strip()) + if resolver is not None: + return resolver() return _get_secret(_env_ref_name(m.group(1)), m.group(0)) or m.group(0) return _ENV_VAR_PATTERN.sub(_replace, value) if isinstance(value, dict): @@ -234,11 +215,10 @@ _whitespace_warned: Set[Tuple[str, str]] = set() def _warn_hidden_whitespace(server_name: str, config: dict) -> List[str]: - """Warn once per (server, key path) about string values with leading/trailing - whitespace — a pasted newline or leading space causes opaque auth/connect - failures and is invisible in config.yaml. Advisory only: values are never - mutated (whitespace could be intentional) and never logged (often secrets). - Returns the flagged key paths.""" + """Warn once per (server, key path) about string values with leading/trailing whitespace — + a pasted newline causes opaque auth/connect failures and is invisible in config.yaml. + Advisory only: values are never mutated (whitespace could be intentional) and never + logged (often secrets). Returns the flagged key paths.""" flagged: List[str] = [] def _walk(value: Any, path: str) -> None: @@ -263,8 +243,7 @@ def _warn_hidden_whitespace(server_name: str, config: dict) -> List[str]: "trailing whitespace — this often causes authentication or " "connection failures. Check for stray spaces/newlines in " "config.yaml (or the referenced env var).", - server_name, key_path, - ) + server_name, key_path) return flagged @@ -286,8 +265,8 @@ def _filter_suspicious_mcp_servers(servers: Dict[str, dict]) -> Dict[str, dict]: def _portable_mcp_servers(safe_servers: Dict[str, dict]) -> None: - """Merge plugin-provided (portable) MCP servers into *safe_servers*; native - config wins on a name clash. Never raises.""" + """Merge plugin-provided (portable) MCP servers into *safe_servers*; native config wins + on a name clash. Never raises.""" try: from hermes_cli.plugins import discover_plugins, get_plugin_manager @@ -303,10 +282,10 @@ def _portable_mcp_servers(safe_servers: Dict[str, dict]) -> None: def _load_mcp_config() -> Dict[str, dict]: - """Read ``mcp_servers`` from config.yaml as ``{name: config}`` (empty on error - or in safe mode). Entries carry ``command``/``args``/``env`` (stdio) or - ``url``/``headers`` (HTTP) plus optional timeout/auth keys; ``${VAR}`` - placeholders are interpolated after ``.env`` is loaded.""" + """Read ``mcp_servers`` from config.yaml as ``{name: config}`` (empty on error or in safe + mode). Entries carry ``command``/``args``/``env`` (stdio) or ``url``/``headers`` (HTTP) + plus optional timeout/auth keys; ``${VAR}`` placeholders are interpolated after ``.env`` + is loaded.""" try: from hermes_cli.config import load_config from utils import env_var_enabled as _env_enabled @@ -316,8 +295,7 @@ def _load_mcp_config() -> Dict[str, dict]: servers = load_config().get("mcp_servers") if not isinstance(servers, dict): servers = {} - # Ensure .env vars are available for interpolation - try: + try: # ensure .env vars are available for interpolation from hermes_cli.env_loader import load_hermes_dotenv load_hermes_dotenv() except Exception: diff --git a/tools/mcp_tool_content.py b/tools/mcp_tool_content.py index 7b284638f8..3eeed70a7a 100644 --- a/tools/mcp_tool_content.py +++ b/tools/mcp_tool_content.py @@ -13,25 +13,21 @@ from tools.mcp_tool_schema import mcp_prefixed_tool_name logger = logging.getLogger("tools.mcp_tool") -# Hard allocation ceiling for one MCP text payload (chars): the first line of -# defense against a multi-megabyte flood being JSON-encoded and handed -# downstream. Deliberately far ABOVE the budget layer's 50K spillover threshold -# so ordinary large results reach spillover intact; only pathological floods -# are lossy-truncated here. +# Hard allocation ceiling for one MCP text payload (chars): the first line of defense against +# a multi-megabyte flood being JSON-encoded and handed downstream. Deliberately far ABOVE the +# budget layer's 50K spillover threshold so ordinary large results reach spillover intact. _MCP_HARD_RESULT_CAP_CHARS = 2_000_000 -# Hard cap on decoded resource bytes from one block, so a misbehaving server -# can't fill the cache disk. +# Hard cap on decoded resource bytes from one block, so a misbehaving server can't fill the +# cache disk. Base64 expands ~4/3; oversized payloads are rejected BEFORE decoding so a +# multi-GB blob string is never transiently doubled in memory. _MCP_RESOURCE_MAX_BYTES = 50 * 1024 * 1024 - -# Base64 expands ~4/3; reject oversized payloads BEFORE decoding so a multi-GB -# blob string is never transiently doubled in memory. _MCP_RESOURCE_MAX_B64_CHARS = _MCP_RESOURCE_MAX_BYTES * 4 // 3 + 4 def _truncate_mcp_text_result(text: str, max_chars: int = _MCP_HARD_RESULT_CAP_CHARS) -> str: - """Pass text at or under ``max_chars`` unchanged; otherwise keep a 40% head / - 60% tail split with an omission notice between.""" + """Pass text at or under ``max_chars`` unchanged; otherwise keep a 40% head / 60% tail + split with an omission notice between.""" if len(text) <= max_chars: return text head_chars = int(max_chars * 0.4) @@ -41,27 +37,23 @@ def _truncate_mcp_text_result(text: str, max_chars: int = _MCP_HARD_RESULT_CAP_C text[:head_chars] + f"\n\n... [MCP RESULT TRUNCATED - {omitted:,} chars omitted " f"out of {len(text):,} total] ...\n\n" - + text[-tail_chars:] - ) + + text[-tail_chars:]) def _is_reserved_mcp_meta_key(key: str) -> bool: - """True if an MCP ``_meta`` key uses a protocol-reserved prefix: a - ``modelcontextprotocol`` or ``mcp`` label followed by at least one more - label. A trailing one (``com.example.mcp/...``) is a vendor namespace.""" + """True if an MCP ``_meta`` key uses a protocol-reserved prefix: a ``modelcontextprotocol`` + or ``mcp`` label followed by at least one more label. A trailing one + (``com.example.mcp/...``) is a vendor namespace.""" slash = key.find("/") if slash <= 0: return False labels = key[:slash].split(".") - return any( - label in ("modelcontextprotocol", "mcp") and i < len(labels) - 1 - for i, label in enumerate(labels) - ) + return any(label in ("modelcontextprotocol", "mcp") and i < len(labels) - 1 for i, label in enumerate(labels)) def _strip_reserved_meta_keys(meta) -> Optional[Dict[str, Any]]: - """Drop protocol-reserved keys from ``_meta``; None if nothing model-facing - remains or the input wasn't a mapping.""" + """Drop protocol-reserved keys from ``_meta``; None if nothing model-facing remains or the + input wasn't a mapping.""" if not isinstance(meta, dict): return None out = {k: v for k, v in meta.items() if isinstance(k, str) and (not _is_reserved_mcp_meta_key(k))} @@ -83,12 +75,9 @@ def _mcp_image_extension_for_mime_type(mime_type: str) -> str: def _decode_block_b64(data, what: str, label: str, *, cap_what: Optional[str] = None, cap_suffix: str = "", decode_fail: str = "") -> Tuple[Optional[bytes], str]: - """Base64-decode one block payload: ``(bytes, "")`` or ``(None, inline_marker)``. - - With ``cap_what`` the payload is rejected on b64 length BEFORE decoding (a - multi-GB string must never be transiently doubled) and on decoded size after. - Decode failures warn and return ``decode_fail`` ("" = drop the block). - """ + """Base64-decode one block payload: ``(bytes, "")`` or ``(None, inline_marker)``. With + ``cap_what`` the payload is rejected on b64 length BEFORE decoding and on decoded size + after. Decode failures warn and return ``decode_fail`` ("" = drop the block).""" if cap_what and len(data) > _core._MCP_RESOURCE_MAX_B64_CHARS: return None, f"[MCP {cap_what} too large to cache: ~{len(data) * 3 // 4} bytes{cap_suffix}]" try: @@ -103,12 +92,10 @@ def _decode_block_b64(data, what: str, label: str, *, cap_what: Optional[str] = def _write_block_cache(writer: str, what: str, skip_label: str, *args, unavailable: str = "", failed: str = "", **kwargs) -> Tuple[Optional[str], str]: - """Call ``gateway.platforms.base.(*args, **kwargs)``: ``(path, "")`` or ``(None, marker)``. - - Fail-open: gateway deps missing (e.g. cron without gateway) → debug log + - ``unavailable``; any other cache error → warning + ``failed``. One bad - block must never kill the tool result. - """ + """Call ``gateway.platforms.base.(*args, **kwargs)``: ``(path, "")`` or ``(None, + marker)``. Fail-open: gateway deps missing (e.g. cron without gateway) → debug log + + ``unavailable``; any other cache error → warning + ``failed``. One bad block must never + kill the tool result.""" try: import gateway.platforms.base as _base @@ -121,49 +108,41 @@ def _write_block_cache(writer: str, what: str, skip_label: str, *args, return None, failed -def _cache_mcp_image_block(block) -> str: - """Cache an ``ImageContent`` block and return a ``MEDIA:`` tag. - - "" (logging, not raising) when the block isn't an image, the base64 is - malformed, or the cache rejects the bytes: the caller falls through to any - text blocks. - """ - data = getattr(block, "data", None) - mime = _base_mime(mcp_field(block, "mime_type", "mimeType")) - if data is None or not mime.startswith("image/"): - return "" - raw_bytes, err = _decode_block_b64(data, "image block", mime) - if raw_bytes is None: - return err - path, err = _write_block_cache( - "cache_image_from_bytes", "image block", "image", - raw_bytes, ext=_mcp_image_extension_for_mime_type(mime), - ) - return err if path is None else f"MEDIA:{path}" - - _WAV_MIME_EXT = {"audio/wav": ".wav", "audio/x-wav": ".wav", "audio/wave": ".wav"} -def _cache_mcp_audio_block(block) -> str: - """Cache an ``AudioContent`` block and return a ``MEDIA:`` tag; "" when not - audio or on any failure (same fail-open contract as the image path).""" +def _cache_mcp_media_block(block, kind: str, writer: str, ext_for, *, cap_what: Optional[str] = None) -> str: + """Cache an image/audio block and return a ``MEDIA:`` tag. "" (logging, not raising) + when the block isn't ``kind`` media, the base64 is malformed, or the cache rejects the + bytes: the caller falls through to any text blocks.""" data = getattr(block, "data", None) mime = _base_mime(mcp_field(block, "mime_type", "mimeType")) - if data is None or not mime.startswith("audio/"): + if data is None or not mime.startswith(f"{kind}/"): return "" - raw_bytes, err = _decode_block_b64(data, "audio block", mime, cap_what="audio resource") + raw_bytes, err = _decode_block_b64(data, f"{kind} block", mime, cap_what=cap_what) if raw_bytes is None: return err - ext = _WAV_MIME_EXT.get(mime) or mimetypes.guess_extension(mime) or ".ogg" - path, err = _write_block_cache("cache_audio_from_bytes", "audio block", "audio", raw_bytes, ext=ext) + path, err = _write_block_cache(writer, f"{kind} block", kind, raw_bytes, ext=ext_for(mime)) return err if path is None else f"MEDIA:{path}" +def _cache_mcp_image_block(block) -> str: + """Cache an ``ImageContent`` block and return a ``MEDIA:`` tag ("" on any failure).""" + return _cache_mcp_media_block(block, "image", "cache_image_from_bytes", _mcp_image_extension_for_mime_type) + + +def _cache_mcp_audio_block(block) -> str: + """Cache an ``AudioContent`` block and return a ``MEDIA:`` tag ("" on any failure).""" + return _cache_mcp_media_block( + block, "audio", "cache_audio_from_bytes", + lambda mime: _WAV_MIME_EXT.get(mime) or mimetypes.guess_extension(mime) or ".ogg", + cap_what="audio resource") + + def _mcp_resource_filename(uri: str, mime_type: str) -> str: - """Safe display filename from the URI's last path segment, used only as a - name hint: ``cache_document_from_bytes`` re-sanitizes and prefixes it, so - remote path components can't steer the cache location.""" + """Safe display filename from the URI's last path segment, used only as a name hint: + ``cache_document_from_bytes`` re-sanitizes and prefixes it, so remote path components + can't steer the cache location.""" import re as _re from pathlib import Path from urllib.parse import urlparse, unquote @@ -174,8 +153,8 @@ def _mcp_resource_filename(uri: str, mime_type: str) -> str: name = Path(unquote(urlparse(str(uri)).path or "")).name except (ValueError, TypeError): name = "" - # Strip control chars (hostile URIs could inject newlines/ANSI into the - # filename and transcript marker) and cap length, preserving the extension. + # Strip control chars (hostile URIs could inject newlines/ANSI into the filename and + # transcript marker) and cap length, preserving the extension. name = _re.sub(r"[\x00-\x1f\x7f]", "", name).strip() if len(name) > 150: stem, dot, ext = name.rpartition(".") @@ -190,19 +169,15 @@ def _mcp_resource_filename(uri: str, mime_type: str) -> str: def _render_mcp_resource_block(block, server_name: str = "") -> str: - """Render a ``ResourceLink`` or ``EmbeddedResource`` block as text. - - Embedded text → the text; embedded blob → decoded (size-capped) into the - document cache with a path marker; link → the URI plus a pointer at the - server's read_resource tool (no fetch here — links are only readable via - the originating session). "" for non-resource blocks; failures are - reported inline rather than silently dropped. - """ + """Render a ``ResourceLink`` or ``EmbeddedResource`` block as text: embedded text → the + text; embedded blob → decoded (size-capped) into the document cache with a path marker; + link → the URI plus a pointer at the server's read_resource tool (no fetch here — links + are only readable via the originating session). "" for non-resource blocks; failures are + reported inline rather than silently dropped.""" block_type = getattr(block, "type", "") if block_type == "resource_link" or ( - hasattr(block, "uri") and not hasattr(block, "resource") and block_type != "text" - ): + hasattr(block, "uri") and not hasattr(block, "resource") and block_type != "text"): uri = getattr(block, "uri", None) if not uri: return "" @@ -213,11 +188,7 @@ def _render_mcp_resource_block(block, server_name: str = "") -> str: details += f", name={name}" if mime: details += f", mimeType={mime}" - reader = ( - mcp_prefixed_tool_name(server_name, "read_resource") - if server_name - else "the MCP server's read_resource tool" - ) + reader = mcp_prefixed_tool_name(server_name, "read_resource") if server_name else "the MCP server's read_resource tool" return f"[MCP resource link: {details} — fetch it with {reader}]" resource = getattr(block, "resource", None) @@ -233,19 +204,15 @@ def _render_mcp_resource_block(block, server_name: str = "") -> str: uri = str(getattr(resource, "uri", "") or "") mime = str(mcp_field(resource, "mime_type", "mimeType", "") or "") raw_bytes, err = _decode_block_b64( - blob, "embedded resource", mime or uri, cap_what="embedded resource", - cap_suffix=f", uri={uri}", - decode_fail=f"[MCP embedded resource could not be decoded: {mime or uri}]", - ) + blob, "embedded resource", mime or uri, cap_what="embedded resource", cap_suffix=f", uri={uri}", + decode_fail=f"[MCP embedded resource could not be decoded: {mime or uri}]") if raw_bytes is None: return err kind = mime or "unknown type" path, err = _write_block_cache( - "cache_document_from_bytes", "embedded resource", "resource", - raw_bytes, _mcp_resource_filename(uri, mime), + "cache_document_from_bytes", "embedded resource", "resource", raw_bytes, _mcp_resource_filename(uri, mime), unavailable=f"[MCP embedded resource received ({len(raw_bytes)} bytes, {kind}) but document cache unavailable in this process]", - failed=f"[MCP embedded resource could not be cached: {mime or uri}]", - ) + failed=f"[MCP embedded resource could not be cached: {mime or uri}]") if path is None: return err return f"[MCP resource saved to {path} ({kind}, {len(raw_bytes)} bytes) — read it with read_file or terminal tools]" diff --git a/tools/mcp_tool_health.py b/tools/mcp_tool_health.py index a2dbd242ae..75cf1ce2fe 100644 --- a/tools/mcp_tool_health.py +++ b/tools/mcp_tool_health.py @@ -1,4 +1,6 @@ -"""Session health for MCPServerTask: dynamic tool refresh on list_changed notifications, server log forwarding, keepalive probes, suspect-mark / lazy-verify, in-flight call fail-fast, stdio child liveness and stdio idle/lifetime recycling. Split from tools/mcp_tool.py.""" +"""Session health for MCPServerTask: dynamic tool refresh on list_changed notifications, server +log forwarding, keepalive probes, suspect-mark / lazy-verify, in-flight call fail-fast, stdio +child liveness and stdio idle/lifetime recycling. Split from tools/mcp_tool.py.""" import asyncio import json @@ -53,7 +55,7 @@ class MCPServerHealthMixin: self._lifecycle_started_at = self._last_tool_call_at = time.monotonic() self._recycled_reason = None - # ------------------------------------------------------- stdio recycling + # -- stdio recycling -- def _stdio_recycle_deadlines(self): """``[(deadline, reason), ...]`` for the configured lifetime/idle limits; empty for HTTP @@ -73,7 +75,6 @@ class MCPServerHealthMixin: return next((reason for deadline, reason in self._stdio_recycle_deadlines() if now >= deadline), None) def _next_stdio_recycle_deadline(self) -> Optional[float]: - """The next monotonic recycle deadline for stdio, if any.""" deadlines = self._stdio_recycle_deadlines() return min(d for d, _ in deadlines) if deadlines else None @@ -82,10 +83,9 @@ class MCPServerHealthMixin: self._recycled_reason = reason self.session = None - # -------------------------------------------------- notifications / logs + # -- notifications / logs -- async def _refresh_tools_task(self): - """Run a dynamic tool refresh and log failures from background tasks.""" try: await self._refresh_tools() except Exception: @@ -110,8 +110,7 @@ class MCPServerHealthMixin: data = json.dumps(data, ensure_ascii=False, default=str) except (TypeError, ValueError): data = str(data) - # Cap payloads so a chatty server can't flood agent.log. - if len(data) > 2000: + if len(data) > 2000: # cap payloads so a chatty server can't flood agent.log data = data[:2000] + "... [truncated]" logger_name = getattr(params, "logger", None) origin = f"{self.name}/{logger_name}" if logger_name else self.name @@ -128,34 +127,32 @@ class MCPServerHealthMixin: if isinstance(message, Exception): logger.debug("MCP message handler (%s): exception: %s", self.name, message) return - if _core._MCP_NOTIFICATION_TYPES and isinstance(message, _core.ServerNotification): - # mcp 2.0 made ServerNotification a plain union (payload IS the message) - # instead of a RootModel (payload under ``.root``). ``isinstance`` accepts - # both; only the unwrap differs — without it ``.root`` raises into the - # catch-all and refreshes stop. - match getattr(message, "root", message): - case _core.ToolListChangedNotification(): - logger.info("MCP server '%s': received tools/list_changed notification", self.name) - # Refresh in a separate task: some servers emit list_changed right - # after initialize while another request is in flight, and refreshing - # synchronously inside the handler can wedge the stdio JSON-RPC stream. - self._schedule_tools_refresh() - # Yield one tick so short-lived notification contexts (and tests) - # can observe the scheduled refresh. - await asyncio.sleep(0) - case _core.PromptListChangedNotification(): - logger.debug("MCP server '%s': prompts/list_changed (ignored)", self.name) - case _core.ResourceListChangedNotification(): - logger.debug("MCP server '%s': resources/list_changed (ignored)", self.name) - case _: - pass + if not (_core._MCP_NOTIFICATION_TYPES and isinstance(message, _core.ServerNotification)): + return + # mcp 2.0 made ServerNotification a plain union (payload IS the message) instead + # of a RootModel (payload under ``.root``). ``isinstance`` accepts both; only the + # unwrap differs — without it ``.root`` raises into the catch-all and refreshes stop. + payload = getattr(message, "root", message) + if isinstance(payload, _core.ToolListChangedNotification): + logger.info("MCP server '%s': received tools/list_changed notification", self.name) + # Refresh in a separate task: some servers emit list_changed right after + # initialize while another request is in flight, and refreshing synchronously + # inside the handler can wedge the stdio JSON-RPC stream. + self._schedule_tools_refresh() + # Yield one tick so short-lived notification contexts (and tests) can observe + # the scheduled refresh. + await asyncio.sleep(0) + elif isinstance(payload, _core.PromptListChangedNotification): + logger.debug("MCP server '%s': prompts/list_changed (ignored)", self.name) + elif isinstance(payload, _core.ResourceListChangedNotification): + logger.debug("MCP server '%s': resources/list_changed (ignored)", self.name) except Exception: logger.exception("Error in MCP message handler for '%s'", self.name) return _handler def _deregister_owned(self, tool_names: Iterable[str]) -> None: - """Deregister *tool_names* that this server's toolset still owns. Never removes a - colliding name currently owned by another server.""" + """Deregister *tool_names* this server's toolset still owns. Never removes a colliding + name currently owned by another server.""" from tools.registry import registry toolset_name = f"mcp-{self.name}" @@ -186,7 +183,6 @@ class MCPServerHealthMixin: registered_names = _core._register_server_tools(self.name, self, self._config) self._deregister_owned(old_tool_names - set(registered_names)) self._registered_tool_names = registered_names - # Log what changed (user-visible). new_tool_names = set(registered_names) changes = [f"{label}: {', '.join(sorted(names))}" for label, names in (("added", new_tool_names - old_tool_names), ("removed", old_tool_names - new_tool_names)) if names] @@ -197,7 +193,7 @@ class MCPServerHealthMixin: logger.info("MCP server '%s': dynamically refreshed %d tool(s) (no changes)", self.name, len(self._registered_tool_names)) - # ------------------------------------------------------ keepalive / health + # -- keepalive / health -- async def _keepalive_probe(self) -> None: """Exercise the session; raise on a genuine connection failure. ``ping`` first (cheap, @@ -210,8 +206,7 @@ class MCPServerHealthMixin: return except Exception as exc: if _is_method_not_found_error(exc): - # Ping is definitively unsupported. - if not self._advertises_tools(): + if not self._advertises_tools(): # ping definitively unsupported, nothing to fall back to raise self._ping_unsupported = True logger.info("MCP server '%s': does not implement the optional 'ping' utility (-32601); " @@ -224,21 +219,18 @@ class MCPServerHealthMixin: await asyncio.wait_for(self.session.list_tools(), timeout=_KEEPALIVE_RPC_TIMEOUT) except Exception: raise exc from None - # Transport alive; latch so later keepalives skip the 30s wait. - self._ping_unsupported = True + self._ping_unsupported = True # latch so later keepalives skip the 30s wait logger.info("MCP server '%s': ping timed out but list_tools succeeded — server " "silently drops ping; using 'list_tools' for keepalive on this connection.", self.name) return else: raise # closed transport, expired session, etc. — real failure - # Fallback probe for servers without ping support. await asyncio.wait_for(self.session.list_tools(), timeout=_KEEPALIVE_RPC_TIMEOUT) def _mark_session_proven(self) -> None: """Record that the session demonstrated real health (keepalive or tool-call success). - Only then is the reconnect budget cleared: a handshake that drops moments later must - keep consuming ``_reconnect_retries`` so a flapping transport still reaches the park - instead of respawning forever.""" + Only then is the reconnect budget cleared: a handshake that drops moments later must keep + consuming ``_reconnect_retries`` so a flapping transport still reaches the park.""" if self._session_proven: return self._session_proven = True @@ -266,8 +258,7 @@ class MCPServerHealthMixin: reason = self._suspect_reason if not reason: return True - if self.session is None: - # Nothing to verify — the reconnect path owns recovery now. + if self.session is None: # nothing to verify — the reconnect path owns recovery now self._suspect_reason = None self._reconnect_event.set() return False @@ -292,8 +283,8 @@ class MCPServerHealthMixin: def _fail_inflight_calls(self, reason: str) -> None: """Cancel every in-flight RPC on this connection. Called from lifecycle exits BEFORE the - transport unwinds: the SDK does not always fail pending requests when streams close, so - a call would otherwise wait out the full tool timeout. Cancelling anything flags + transport unwinds: the SDK does not always fail pending requests when streams close, so a + call would otherwise wait out the full tool timeout. Cancelling anything flags ``_teardown_race`` so run() treats the next reconnect as recovery rather than charging the rapid-drop budget.""" victims = [t for t in self._inflight_tasks if not t.done()] @@ -306,7 +297,6 @@ class MCPServerHealthMixin: task.cancel() def _stdio_children_dead(self) -> bool: - """True when every stdio child we spawned has exited (see :func:`_stdio_children_dead_impl`).""" return _stdio_children_dead_impl(getattr(self, "_stdio_child_pids", None), self._is_http()) async def _watch_stdio_children(self) -> None: diff --git a/tools/mcp_tool_lifecycle.py b/tools/mcp_tool_lifecycle.py index 71a338c443..7543035ad2 100644 --- a/tools/mcp_tool_lifecycle.py +++ b/tools/mcp_tool_lifecycle.py @@ -31,14 +31,13 @@ _stdio_pgids: Dict[int, int] = {} def _snapshot_child_pids() -> set: """Current direct-child PIDs: /proc on Linux, else psutil, else empty set.""" my_pid = os.getpid() - # /proc//task//children is per-THREAD, and stdio_client() spawns - # from the MCP loop thread, so union every task's children — reading only - # the main thread's file returns an empty set on every Linux install. + # /proc//task//children is per-THREAD, and stdio_client() spawns from the MCP + # loop thread, so union every task's children — reading only the main thread's file + # returns an empty set on every Linux install. try: task_dir = f"/proc/{my_pid}/task" - tids = os.listdir(task_dir) found: set = set() - for tid in tids: + for tid in os.listdir(task_dir): try: with open(f"{task_dir}/{tid}/children", encoding="utf-8") as f: found.update(int(p) for p in f.read().split() if p.strip()) @@ -63,8 +62,7 @@ _NON_MCP_CHILD_CMDLINE_MARKERS: tuple[str, ...] = ( "tui_gateway.entry", "-dorg.eclipse.equinox.launcher", # jdtls (legacy arg style) "eclipse.jdt.ls", - "org.eclipse.equinox.launcher_", -) + "org.eclipse.equinox.launcher_") def _filter_mcp_children(pids: set) -> set: @@ -106,50 +104,40 @@ def shutdown_mcp_servers(*, scope: Optional[str] = None): selected = [name for name in _core._servers if scope is None or _core._server_scope_keys.get(name) == scope] servers_snapshot = [_core._servers[name] for name in selected] - # Fast path: nothing to shut down. Still clear the connect-cooldown maps — - # a server that failed to connect is never in ``_servers``, so this is the - # most likely state for stale backoff entries; a restart must retry at once. - if not servers_snapshot: + if servers_snapshot: + async def _shutdown(): + results = await asyncio.gather(*(server.shutdown() for server in servers_snapshot), return_exceptions=True) + for server, result in zip(servers_snapshot, results): + if isinstance(result, Exception): + logger.debug("Error closing MCP server '%s': %s", server.name, result) + with _core._lock: + for name in selected: + _core._servers.pop(name, None) + _core._server_scope_keys.pop(name, None) + _clear_connect_cooldowns() + with _core._lock: - _clear_connect_cooldowns() - _core._stop_mcp_loop(only_if_idle=scope is not None) - return + loop = _core._mcp_loop + if loop is not None and loop.is_running(): + from agent.async_utils import safe_schedule_threadsafe + future = safe_schedule_threadsafe(_shutdown(), loop, logger=logger, log_message="MCP shutdown: failed to schedule") + if future is not None: + try: + future.result(timeout=15) + except BaseException as exc: + logger.debug("Error during MCP shutdown: %s", exc) - async def _shutdown(): - results = await asyncio.gather(*(server.shutdown() for server in servers_snapshot), return_exceptions=True) - for server, result in zip(servers_snapshot, results): - if isinstance(result, Exception): - logger.debug("Error closing MCP server '%s': %s", server.name, result) - with _core._lock: - for name in selected: - _core._servers.pop(name, None) - _core._server_scope_keys.pop(name, None) - _clear_connect_cooldowns() - - with _core._lock: - loop = _core._mcp_loop - if loop is not None and loop.is_running(): - from agent.async_utils import safe_schedule_threadsafe - future = safe_schedule_threadsafe( - _shutdown(), loop, logger=logger, log_message="MCP shutdown: failed to schedule", - ) - if future is not None: - try: - future.result(timeout=15) - except BaseException as exc: - logger.debug("Error during MCP shutdown: %s", exc) - - # Unconditional final sweep: whether ``_shutdown`` ran, timed out, or was - # never scheduled, no stale connect-cooldown state may survive shutdown. + # Unconditional final sweep: whether ``_shutdown`` ran, timed out, or was never scheduled + # (a server that failed to connect is never in ``_servers`` — the most likely state for + # stale backoff entries), no connect-cooldown state may survive shutdown. with _core._lock: _clear_connect_cooldowns() _core._stop_mcp_loop(only_if_idle=scope is not None) def _take_reapable_pids(include_active: bool, server_name: Optional[str]) -> tuple[Dict[int, str], Dict[int, int]]: - """Pop the PIDs to reap (and their spawn-time pgids) out of the ledgers under - the lock, so a future spawn can't collide with stale state. - Returns ``(pid -> owner, pid -> pgid)``.""" + """Pop the PIDs to reap (and their spawn-time pgids) out of the ledgers under the lock, so + a future spawn can't collide with stale state. Returns ``(pid -> owner, pid -> pgid)``.""" def _owned(entries: Dict[int, str]) -> Dict[int, str]: return {pid: owner for pid, owner in entries.items() if server_name is None or owner == server_name} @@ -173,25 +161,21 @@ def _signal_mcp_process(pid: int, sig: int, server_name: str, pgid: Optional[int killpg = getattr(os, "killpg", None) if pgid is not None and killpg is not None: if my_pgid is not None and pgid == my_pgid: - # Child shares the gateway's pgroup: killpg would kill the gateway - # too, so use per-pid kill. Warn because per-pid kill can't reach - # grandchildren in this group (inherent trade-off). + # Child shares the gateway's pgroup: killpg would kill the gateway too, so use + # per-pid kill. Warn because per-pid kill can't reach grandchildren in this group. logger.warning( "MCP server '%s' pgid %d matches gateway pgid; skipping " "killpg to avoid self-kill and using per-pid kill — any " "grandchildren in this group may not be reaped", - server_name, pgid, - ) + server_name, pgid) else: try: killpg(pgid, sig) return except (ProcessLookupError, PermissionError, OSError) as exc: # Pgroup gone or refused — still try the direct child. - logger.debug( - "killpg(%d, %d) failed for MCP server '%s': %s; falling back to kill(pid)", - pgid, sig, server_name, exc, - ) + logger.debug("killpg(%d, %d) failed for MCP server '%s': %s; falling back to kill(pid)", + pgid, sig, server_name, exc) try: os.kill(pid, sig) except (ProcessLookupError, PermissionError, OSError): @@ -208,13 +192,10 @@ def _kill_orphaned_mcp_children(include_active: bool = False, server_name: Optio import signal as _signal pids, pgids = _take_reapable_pids(include_active, server_name) - # Fast path: nothing to reap — skip the 2s sleep every MCP-free shutdown - # would otherwise pay. - if not pids: + if not pids: # skip the 2s sleep every MCP-free shutdown would otherwise pay return - # Our own pgid, so we never killpg() the gateway itself. - try: + try: # our own pgid, so we never killpg() the gateway itself my_pgid = os.getpgrp() except (AttributeError, OSError): my_pgid = None # Windows or restricted environment @@ -236,18 +217,17 @@ def _kill_orphaned_mcp_children(include_active: bool = False, server_name: Optio def _stop_mcp_loop_if_idle() -> bool: - """Stop the MCP loop only when no registered server still owns it. Probe - paths create temporary MCPServerTasks not placed in ``_servers``; they may - clean up an idle loop but must not tear down the process-global loop under - live agent tools, or later calls fail with ``MCP event loop is not running``.""" + """Stop the MCP loop only when no registered server still owns it. Probe paths create + temporary MCPServerTasks not placed in ``_servers``; they may clean up an idle loop but + must not tear down the process-global loop under live agent tools.""" return _core._stop_mcp_loop(only_if_idle=True) async def _drain_mcp_loop_tasks(*, timeout: Optional[float] = None) -> None: - """Cancel every task still pending on the MCP loop and reap it. - ``Task.cancel()`` only schedules the throw, so tasks need a cancellation - cycle before the loop goes away; wait for them here, on their owning loop, - bounded so a task that suppresses cancellation cannot hang process exit.""" + """Cancel every task still pending on the MCP loop and reap it. ``Task.cancel()`` only + schedules the throw, so tasks need a cancellation cycle before the loop goes away; wait + for them here, on their owning loop, bounded so a task that suppresses cancellation + cannot hang process exit.""" if timeout is None: timeout = _core._MCP_LOOP_DRAIN_TIMEOUT current = asyncio.current_task() @@ -257,7 +237,6 @@ async def _drain_mcp_loop_tasks(*, timeout: Optional[float] = None) -> None: logger.debug("Draining %d pending task(s) from the MCP loop", len(pending)) for task in pending: task.cancel() - done, still_pending = await asyncio.wait(pending, timeout=timeout) for task in done: try: @@ -267,16 +246,14 @@ async def _drain_mcp_loop_tasks(*, timeout: Optional[float] = None) -> None: pass except Exception as exc: logger.debug("Pending MCP loop task ended during shutdown: %s", exc) - if still_pending: logger.warning("%d MCP loop task(s) still pending after %.1fs drain", len(still_pending), timeout) async def _drain_and_stop_mcp_loop() -> None: - """Drain pending tasks, then stop the loop from its owning thread. Both must - run as one loop-owned sequence: a ``loop.stop`` queued separately by a - timed-out caller can overtake the scheduled drain, leaving the drain - coroutine itself pending when the loop is closed.""" + """Drain pending tasks, then stop the loop from its owning thread. Both must run as one + loop-owned sequence: a ``loop.stop`` queued separately by a timed-out caller can overtake + the scheduled drain, leaving the drain coroutine itself pending when the loop is closed.""" loop = asyncio.get_running_loop() try: await _drain_mcp_loop_tasks(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT) diff --git a/tools/setup_mcp_tool.py b/tools/setup_mcp_tool.py index 624540016b..95539cc16c 100644 --- a/tools/setup_mcp_tool.py +++ b/tools/setup_mcp_tool.py @@ -1,17 +1,12 @@ #!/usr/bin/env python3 """Propose an MCP server to the user as an inline card in the desktop chat. -The card (install / enable / authorize + decline) lives in the desktop -renderer, so this tool round-trips through the gateway's blocking-prompt -bridge — the same one ``clarify`` uses: tui_gateway emits -``mcp.setup.request``, the renderer walks the user through the flow via the -existing REST endpoints (catalog install, enable, OAuth), and answers with -``mcp.setup.respond`` once the flow settles. This module is just schema + a -thin dispatcher over the platform-injected callback. - -Lives in the ``desktop_ui`` toolset, which the GUI gateway enables only for -desktop-sourced sessions — on every other surface the agent falls back to -``hermes mcp install `` in the terminal. +The card (install / enable / authorize + decline) lives in the desktop renderer, so this +tool round-trips through the gateway's blocking-prompt bridge (the one ``clarify`` uses): +tui_gateway emits ``mcp.setup.request``, the renderer walks the user through the existing +REST flows (catalog install, enable, OAuth) and answers with ``mcp.setup.respond``. Lives in +the ``desktop_ui`` toolset, which the GUI gateway enables only for desktop-sourced sessions; +elsewhere the agent falls back to ``hermes mcp install `` in the terminal. """ import json @@ -22,19 +17,13 @@ from tools.registry import registry, tool_error _ACTIONS = ("install", "enable", "authorize") -def setup_mcp_tool( - server: str = "", - action: str = "install", - reason: str = "", - callback: Optional[Callable] = None, -) -> str: +def setup_mcp_tool(server: str = "", action: str = "install", reason: str = "", callback: Optional[Callable] = None) -> str: """Ask the desktop GUI to run an MCP setup flow; return its JSON outcome.""" if callback is None: return tool_error( "setup_mcp is only available in the Hermes desktop app. Use the " "terminal instead: `hermes mcp install ` for catalog entries, " - "`hermes mcp login ` for OAuth." - ) + "`hermes mcp login ` for OAuth.") name = (server or "").strip() if not name: @@ -50,19 +39,14 @@ def setup_mcp_tool( return tool_error(f"MCP setup flow failed: {exc}") if not raw: - # The renderer never answered (timeout / closed window). Distinct from - # an explicit decline, which arrives as {"status": "declined"}. - return json.dumps( - { - "status": "unanswered", - "server": name, - "note": ( - "The user did not respond to the setup card. Do not retry " - "immediately; continue without the server or ask in chat." - ), - }, - ensure_ascii=False, - ) + # The renderer never answered (timeout / closed window). Distinct from an explicit + # decline, which arrives as {"status": "declined"}. + return json.dumps({ + "status": "unanswered", + "server": name, + "note": ("The user did not respond to the setup card. Do not retry " + "immediately; continue without the server or ask in chat."), + }, ensure_ascii=False) # Desktop answers with a JSON object; pass it through, else wrap the raw text. try: