From 6e1089a23c64517c4095459b4ebcc7db371e49d6 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 17:05:22 -0700 Subject: [PATCH 1/4] refactor(mcp): compact mcp_oauth_manager docstrings/comments by hand; flatten single-use call sites --- tools/mcp_oauth_manager.py | 312 ++++++++++++++----------------------- 1 file changed, 116 insertions(+), 196 deletions(-) diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index a0d0d34b18..99891d21e3 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -1,18 +1,16 @@ #!/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 — Claude Code's ``invalidateOAuthCacheIfDiskChanged`` bug -class), 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). +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. +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. """ from __future__ import annotations @@ -31,28 +29,23 @@ 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.""" + """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("/") - ) + return pa.scheme == pb.scheme and pa.netloc.lower() == pb.netloc.lower() and pa.path.rstrip("/") == pb.path.rstrip("/") @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`` 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.""" server_url: str oauth_config: Optional[dict] @@ -64,29 +57,27 @@ class _ProviderEntry: # -- 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. - """ + """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.""" _hermes_logger = logger def __init__(self, *args: Any, server_name: str = "", preregistered: bool = False, **kwargs: Any): super().__init__(*args, **kwargs) - # mcp 2.0 uses a task-owned anyio.Lock held across the yielded resource - # request: a session-long GET blocks every concurrent POST, and HTTPX may - # close the auth-flow generator from another task. A binary semaphore - # keeps mutual exclusion without task ownership; async_auth_flow narrows - # its scope around resource I/O. + # mcp 2.0 uses a task-owned anyio.Lock held across the yielded resource request: a + # session-long GET blocks every concurrent POST, and HTTPX may close the auth-flow + # 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 (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. self._hermes_preregistered = preregistered def _hermes_storage(self): @@ -103,15 +94,14 @@ class _HermesRuntimeProviderMixin: """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. + ``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. """ await super()._initialize() tokens = self.context.current_tokens @@ -131,20 +121,16 @@ class _HermesRuntimeProviderMixin: 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 — defensive - # Non-fatal: the SDK's 401-branch discovery runs next request. + except Exception as exc: # pragma: no cover — non-fatal: 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. - Uses 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 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.""" # 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, which only its own + # AsyncClient can send (tools.mcp_tool.sdk_httpx). from tools.mcp_tool import sdk_httpx httpx = sdk_httpx() if httpx is None: # pragma: no cover — SDK import would have failed @@ -163,10 +149,7 @@ class _HermesRuntimeProviderMixin: try: return await client.send(create_oauth_metadata_request(url)) except httpx.HTTPError as exc: - logger.debug( - "MCP OAuth '%s': %s discovery to %s failed: %s", - self._hermes_server_name, label, url, exc, - ) + logger.debug("MCP OAuth '%s': %s discovery to %s failed: %s", self._hermes_server_name, label, url, exc) return None async with httpx.AsyncClient(timeout=10.0) as client: @@ -180,8 +163,7 @@ class _HermesRuntimeProviderMixin: 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). + # Step 2: 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: @@ -202,8 +184,8 @@ class _HermesRuntimeProviderMixin: 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) for future restarts; no-op when + absent, not our storage, or unchanged.""" meta = self.context.oauth_metadata storage = self._hermes_storage() if meta is None or storage is None: @@ -213,16 +195,11 @@ 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 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.""" 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 - ) + 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): @@ -233,17 +210,15 @@ 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``. + 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 + 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. """ try: @@ -253,18 +228,16 @@ class _HermesRuntimeProviderMixin: 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). + # 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.", + "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 @@ -282,17 +255,14 @@ class _HermesRuntimeProviderMixin: async def async_auth_flow(self, request): # type: ignore[override] # Pre-flow hook: reload from disk if it changed (non-fatal on error). try: - await get_manager().invalidate_if_disk_changed( - self._hermes_server_name, hermes_home=self._hermes_home - ) + await get_manager().invalidate_if_disk_changed(self._hermes_server_name, hermes_home=self._hermes_home) except Exception as exc: # pragma: no cover — defensive 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`` (tests/tools/test_mcp_oauth_bidirectional.py). + # 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``. inner = super().async_auth_flow(request) resource_lock_released = False sent_access_token = None @@ -300,10 +270,9 @@ class _HermesRuntimeProviderMixin: try: outgoing = await inner.__anext__() while True: - # The SDK holds context.lock for its whole generator, even while - # HTTPX waits on the MCP request. Release it for that request - # only; discovery/refresh/registration/exchange stay serialized - # exactly as the SDK implements them. + # The SDK holds context.lock for its whole generator, even while HTTPX waits on + # the MCP request. Release it for that request only; discovery/refresh/ + # registration/exchange stay serialized exactly as the SDK implements them. if outgoing is request: tokens = self.context.current_tokens sent_access_token = tokens.access_token if tokens is not None else None @@ -313,9 +282,9 @@ class _HermesRuntimeProviderMixin: if resource_lock_released: await self.context.lock.acquire() resource_lock_released = False - # Another request may have refreshed/authorized while this one - # was in flight: retry with that token instead of a duplicate - # OAuth transition from the stale 401/403. + # Another request may have refreshed/authorized while this one was in flight: + # 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) @@ -336,9 +305,8 @@ class _HermesRuntimeProviderMixin: 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. + # 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): @@ -351,21 +319,18 @@ class _HermesRuntimeProviderMixin: 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).""" + """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``. - """ + """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 @@ -376,50 +341,36 @@ _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: 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).""" def __init__(self) -> None: self._entries: dict[tuple[str, str], _ProviderEntry] = {} self._entries_lock = threading.Lock() - # Strong refs to in-flight 401 tasks so the loop's weak bookkeeping - # cannot GC them mid-run and leave `await pending` hanging forever. + # Strong refs to in-flight 401 tasks so the loop's weak bookkeeping cannot GC them + # mid-run and leave `await pending` hanging forever. self._inflight_tasks: set[asyncio.Task] = set() - # -- Provider construction / caching ------------------------------------- - + # -- 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. - - Returns None if the MCP SDK's OAuth support is unavailable. - """ + """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.""" key = self._key(server_name) with self._entries_lock: entry = self._entries.get(key) if entry is not None and entry.server_url != server_url: - logger.info( - "MCP OAuth '%s': URL changed from %s to %s, discarding cache", - server_name, entry.server_url, server_url, - ) + 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 - if entry.provider is None: entry.provider = self._build_provider(server_name, entry) if entry.provider is not None: entry.provider._hermes_home = key[0] - return entry.provider @staticmethod @@ -430,36 +381,24 @@ class MCPOAuthManager: 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 :class:`HermesMCPOAuthProvider` from the shared ``tools.mcp_oauth`` helpers; + 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 - # Local imports avoid circular deps at module import time. + 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) - - from tools.mcp_dashboard_oauth import get_dashboard_oauth_flow - - if ( - get_dashboard_oauth_flow() is None - and not _is_interactive() - and not storage.has_cached_tokens() - ): + if get_dashboard_oauth_flow() is None and not _is_interactive() and not storage.has_cached_tokens(): raise OAuthNonInteractiveError( - "MCP OAuth for " - f"'{server_name}': non-interactive environment and no " - "cached tokens found. Run `hermes mcp login " - f"{server_name}` interactively first to complete initial " - "authorization." + 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." ) - return _HERMES_PROVIDER_CLS( server_name=server_name, preregistered=bool(cfg.get("client_id")), @@ -468,11 +407,8 @@ class MCPOAuthManager: ) def remove(self, server_name: str, *, hermes_home: str | Path | None = None) -> _ProviderEntry | None: - """Evict the provider from cache AND delete tokens from disk. - - Called by ``hermes mcp remove `` and (indirectly) by - ``hermes mcp login `` during forced re-auth. - """ + """Evict the provider from cache AND delete tokens from disk (``hermes mcp remove`` and, + indirectly, ``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) @@ -493,43 +429,34 @@ class MCPOAuthManager: with self._entries_lock: return self._entries.pop(self._key(server_name, hermes_home), None) - # -- Disk watch ---------------------------------------------------------- - + # -- 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. - - Returns True if invalidated. This is the external-refresh fix: a cron - job writes fresh tokens and the next tool call picks them up. - """ + """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.""" 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 - async with entry.lock: tokens_path = _get_token_dir(hermes_home) / f"{_safe_filename(server_name)}.json" try: mtime_ns = tokens_path.stat().st_mtime_ns except (FileNotFoundError, OSError): return False - if mtime_ns == entry.last_mtime_ns: return False 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. + # `_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"): entry.provider._initialized = False # noqa: SLF001 - logger.info( - "MCP OAuth '%s': tokens file changed (mtime %d -> %d), forcing reload", - server_name, old, mtime_ns, - ) + logger.info("MCP OAuth '%s': tokens file changed (mtime %d -> %d), forcing reload", server_name, old, mtime_ns) return True - # -- 401 handler (dedup'd) ----------------------------------------------- - + # -- 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: @@ -538,9 +465,8 @@ class MCPOAuthManager: 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). + # 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 @@ -556,21 +482,16 @@ class MCPOAuthManager: entry.pending_401.pop(key, None) 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. - """ + """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.""" entry = self._entries.get(self._key(server_name)) if entry is None or entry.provider is None: return False - key = failed_access_token or "" loop = asyncio.get_running_loop() - async with entry.lock: pending = entry.pending_401.get(key) if pending is None: @@ -579,7 +500,6 @@ class MCPOAuthManager: 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) - try: return await pending except Exception as exc: # pragma: no cover — defensive From 607a5b711edc149d02655e5cd83037546a05c256 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 17:05:28 -0700 Subject: [PATCH 2/4] =?UTF-8?q?refactor(mcp):=20errors=20=E2=80=94=20one?= =?UTF-8?q?=20=5Fjsonrpc=5Fmatches=20for=20both=20code/marker=20classifier?= =?UTF-8?q?s,=20one=20=5Fexc=5Fchildren=20walk=20for=20connect-error=20ren?= =?UTF-8?q?dering,=20docstrings=20compacted=20by=20hand=20(455=20->=20384;?= =?UTF-8?q?=20differential=20fuzz=2020k=20cases=20identical)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/mcp_tool_errors.py | 281 +++++++++++++++------------------------ 1 file changed, 105 insertions(+), 176 deletions(-) diff --git a/tools/mcp_tool_errors.py b/tools/mcp_tool_errors.py index e650f8b234..0d769422b1 100644 --- a/tools/mcp_tool_errors.py +++ b/tools/mcp_tool_errors.py @@ -23,68 +23,53 @@ def _jsonrpc_code(exc: BaseException): return getattr(getattr(exc, "error", None), "code", None) -def _handshake_rejected_as_modern(exc: BaseException) -> bool: - """True when a failed ``initialize`` signals a stateless-only (2026-07-28) server. - - Structural code check first, then substring fallback — never ``isinstance`` on - SDK exception types (they arrive wrapped in ExceptionGroups and drift across generations). - """ - code = _jsonrpc_code(exc) or getattr(exc, "code", None) - if code in (_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, _core._JSONRPC_METHOD_NOT_FOUND): +def _jsonrpc_matches(exc: BaseException, code, codes: tuple, markers: tuple) -> bool: + """Structural *code* in *codes*, else any lowercased *marker* in ``str(exc)``. Never + ``isinstance`` on SDK exception types: they arrive wrapped in ExceptionGroups and drift + across generations.""" + if code in codes: return True msg = str(exc).lower() - return bool(msg) and ( - "unsupported protocol version" in msg - or str(_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION) in msg - or _is_method_not_found_error(exc) - ) + return bool(msg) and any(marker in msg for marker in markers) + + +def _handshake_rejected_as_modern(exc: BaseException) -> bool: + """True when a failed ``initialize`` signals a stateless-only (2026-07-28) server.""" + return _jsonrpc_matches( + exc, _jsonrpc_code(exc) or getattr(exc, "code", None), + (_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, _core._JSONRPC_METHOD_NOT_FOUND), + ("unsupported protocol version", str(_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION)), + ) or _is_method_not_found_error(exc) def _is_method_not_found_error(exc: BaseException) -> bool: - """True if *exc* is a JSON-RPC ``method not found`` (-32601). - - ``ping`` is optional in MCP; servers lacking it answer -32601. Structural - code check first, then substring fallback — including "Unknown method: ", - which some servers use; without it the ping→list_tools keepalive fallback - never latches and reconnect-loops. - """ - if _jsonrpc_code(exc) == _core._JSONRPC_METHOD_NOT_FOUND: - return True - msg = str(exc).lower() - return bool(msg) and ( - str(_core._JSONRPC_METHOD_NOT_FOUND) in msg - or "method not found" in msg - or "unknown method" in msg - or "not found: ping" in msg + """True if *exc* is a JSON-RPC ``method not found`` (-32601). ``ping`` is optional in MCP; + servers lacking it answer -32601. The substring fallback includes "Unknown method: ", + which some servers use; without it the ping→list_tools keepalive fallback never latches + and reconnect-loops.""" + return _jsonrpc_matches( + exc, _jsonrpc_code(exc), (_core._JSONRPC_METHOD_NOT_FOUND,), + (str(_core._JSONRPC_METHOD_NOT_FOUND), "method not found", "unknown method", "not found: ping"), ) class InvalidMcpUrlError(ValueError): - """A remote MCP server's ``url`` is not parseable http(s)://. - - Validated once at startup so we fail fast instead of burning the - reconnect-backoff loop on every attempt. - """ + """A remote MCP server's ``url`` is not parseable http(s)://. Validated once at startup so + we fail fast instead of burning the reconnect-backoff loop on every attempt.""" class NonMcpEndpointError(ConnectionError): - """An HTTP MCP URL served a non-MCP 2xx response (e.g. ``text/html``). - - Real Streamable-HTTP endpoints answer ``application/json`` or - ``text/event-stream``. Non-retryable: every attempt gets the same page, so - the backoff loop is skipped and the server is failed immediately. - Subclasses ConnectionError so broad catches still see a connection problem. - """ + """An HTTP MCP URL served a non-MCP 2xx response (e.g. ``text/html``); real Streamable-HTTP + endpoints answer ``application/json`` or ``text/event-stream``. Non-retryable: every attempt + gets the same page, so the backoff loop is skipped and the server is failed immediately. + Subclasses ConnectionError so broad catches still see a connection problem.""" def _unwrap_exception_group(exc: BaseException) -> BaseException: - """Extract the root-cause leaf from anyio ``(Base)ExceptionGroup`` wrappers. - - Group ``str()`` is opaque ("unhandled errors in a TaskGroup"), so log sites - must unwrap. Two rules: a ``KeyboardInterrupt``/``SystemExit`` leaf anywhere - is re-raised (never flattened into a loggable error); a non-cancellation - leaf is preferred over the ``CancelledError`` noise anyio sprays on siblings. - """ + """Extract the root-cause leaf from anyio ``(Base)ExceptionGroup`` wrappers (group ``str()`` + is opaque, so log sites must unwrap). Two rules: a ``KeyboardInterrupt``/``SystemExit`` leaf + anywhere is re-raised (never flattened into a loggable error); a non-cancellation leaf is + preferred over the ``CancelledError`` noise anyio sprays on siblings.""" while isinstance(exc, BaseExceptionGroup) and exc.exceptions: fatal, _rest = exc.split((KeyboardInterrupt, SystemExit)) if fatal is not None: @@ -104,19 +89,15 @@ def _contains_only_cancellation(exc: BaseException) -> bool: def _classify_mcp_failure(exc: BaseException) -> str: - """Classify a connection failure as ``'permanent'`` or ``'transient'``. - - Permanent (deterministic — ``run()`` parks immediately instead of burning the - retry ladder): auth 401/403, NonMcpEndpointError, InvalidMcpUrlError, missing - stdio command (FileNotFoundError / ENOENT). Everything else keeps backoff retry. - """ + """``'permanent'`` (deterministic — ``run()`` parks immediately instead of burning the retry + ladder: auth 401/403, NonMcpEndpointError, InvalidMcpUrlError, missing stdio command + FileNotFoundError / ENOENT) or ``'transient'`` (keeps backoff retry).""" root = _unwrap_exception_group(exc) permanent = ( _core._is_auth_error(root) or isinstance(root, (NonMcpEndpointError, InvalidMcpUrlError, FileNotFoundError)) or (isinstance(root, OSError) and getattr(root, "errno", None) == errno.ENOENT) - # 401/403 HTTPStatusError that _is_auth_error's type-gate missed - # (auth types not importable in this environment). + # 401/403 HTTPStatusError that _is_auth_error's type-gate missed (auth types not importable here). or _response_status(root) in (401, 403) ) return "permanent" if permanent else "transient" @@ -128,11 +109,9 @@ def _response_status(exc: BaseException): def _validate_remote_mcp_url(server_name: str, url: Any) -> str: - """Return the stripped URL if it is a valid http(s) remote MCP URL. - - Raises InvalidMcpUrlError naming the server for non-strings, missing/other - schemes (stdio servers use ``command``, not ``url``), and empty hosts. - """ + """The stripped URL if it is a valid http(s) remote MCP URL. Raises InvalidMcpUrlError naming + the server for non-strings, missing/other schemes (stdio servers use ``command``, not ``url``), + and empty hosts.""" def _bad(detail: str) -> InvalidMcpUrlError: return InvalidMcpUrlError(f"Invalid MCP URL for '{server_name}': {detail}") @@ -149,20 +128,16 @@ def _validate_remote_mcp_url(server_name: str, url: Any) -> str: raise _bad(f"scheme must be http or https, got {parsed.scheme!r} ({stripped!r})") if not parsed.netloc: raise _bad(f"missing host ({stripped!r})") - # ``urlparse`` accepts ``http://:8080`` (empty host, explicit port) — reject it. - if not parsed.hostname: + if not parsed.hostname: # ``urlparse`` accepts ``http://:8080`` (empty host, explicit port) raise _bad(f"missing hostname ({stripped!r})") return stripped def _resolve_client_cert(server_name: str, config: dict): - """Resolve ``client_cert`` / ``client_key`` into httpx's ``cert=`` shape. - - None when neither is set; a single path for a combined PEM; ``(cert, key)`` - or ``(cert, key, password)`` for the pair/list forms. ``~`` is expanded and - missing files raise a server-scoped FileNotFoundError instead of an opaque - TLS handshake error. - """ + """``client_cert`` / ``client_key`` in httpx's ``cert=`` shape: None when neither is set; a + single path for a combined PEM; ``(cert, key)`` or ``(cert, key, password)`` for the + pair/list forms. ``~`` is expanded and missing files raise a server-scoped + FileNotFoundError instead of an opaque TLS handshake error.""" raw_cert = config.get("client_cert") raw_key = config.get("client_key") if raw_cert is None and raw_key is None: @@ -179,20 +154,16 @@ def _resolve_client_cert(server_name: str, config: dict): if isinstance(raw_cert, (list, tuple)): if raw_key is not None: - raise ValueError( - f"{prefix}specify either client_cert as a list [cert, key] OR " - f"client_cert + client_key, not both" - ) + raise ValueError(f"{prefix}specify either client_cert as a list [cert, key] OR " + f"client_cert + client_key, not both") if len(raw_cert) not in (2, 3): raise ValueError(f"{prefix}client_cert list form must have 2 or 3 elements (got {len(raw_cert)})") pair = (_expand(raw_cert[0], "client_cert[0]"), _expand(raw_cert[1], "client_cert[1]")) if len(raw_cert) == 2: return pair - password = raw_cert[2] - if not isinstance(password, str): + if not isinstance(raw_cert[2], str): raise ValueError(f"{prefix}client_cert[2] (key passphrase) must be a string") - return (*pair, password) - + return (*pair, raw_cert[2]) cert_path = _expand(raw_cert, "client_cert") if raw_key is not None: return (cert_path, _expand(raw_key, "client_key")) @@ -200,13 +171,10 @@ def _resolve_client_cert(server_name: str, config: dict): def _resolve_identity_header(server_name: str, config: dict): - """Resolve the optional per-server ``identity_header`` config. - - Shape: ``{name: "X-User-Id", value_from: "static"|"profile", value: "..."}`` - (``value`` required for static). Returns ``(name, value)`` or None. Invalid - configs warn and are ignored — an identity header must never break the - connection. ``profile`` resolves once at connect time; no per-call mutation. - """ + """Optional per-server ``identity_header`` ``{name, value_from: "static"|"profile", value}`` + (``value`` required for static) → ``(name, value)`` or None. Invalid configs warn and are + ignored — an identity header must never break the connection. ``profile`` resolves once at + connect time; no per-call mutation.""" raw = config.get("identity_header") if raw is None: return None @@ -233,38 +201,28 @@ def _resolve_identity_header(server_name: str, config: dict): def _apply_identity_header(server_name: str, config: dict, headers: dict) -> dict: - """Merge the resolved identity header into ``headers`` in place. - - An explicit per-server ``headers`` entry with the same name (any casing) - wins — the identity header never silently overrides user config. - """ + """Merge the resolved identity header into ``headers`` in place. An explicit per-server + ``headers`` entry with the same name (any casing) wins — the identity header never silently + overrides user config.""" resolved = _resolve_identity_header(server_name, config) if resolved is None: return headers name, value = resolved if any(key.lower() == name.lower() for key in headers): - logger.debug( - "MCP server '%s': identity_header '%s' already set via explicit " - "headers config — keeping the explicit value", server_name, name, - ) + logger.debug("MCP server '%s': identity_header '%s' already set via explicit " + "headers config — keeping the explicit value", server_name, name) return headers headers[name] = value return headers -def _make_redirect_header_stripper( - original_url, - *, - strict: bool = False, - configured_header_names: "set[str] | frozenset[str]" = frozenset(), -): - """Build an httpx response hook that guards cross-origin redirects. - - Always strips ``Authorization`` when a redirect leaves the original origin. - With *strict* (Agent Plugins v1 ``strict_redirect_headers``) every configured - header (lowercase names in *configured_header_names*) is stripped too — the - v1 spec forbids forwarding package-configured headers cross-origin. - """ +def _make_redirect_header_stripper(original_url, *, strict: bool = False, + configured_header_names: "set[str] | frozenset[str]" = frozenset()): + """httpx response hook guarding cross-origin redirects: always strips ``Authorization`` when + a redirect leaves the original origin; with *strict* (Agent Plugins v1 + ``strict_redirect_headers``) every configured header (lowercase names in + *configured_header_names*) is stripped too — the v1 spec forbids forwarding + package-configured headers cross-origin.""" origin = (original_url.scheme, original_url.host, original_url.port) async def _strip_on_cross_origin_redirect(response): @@ -284,45 +242,41 @@ def _make_redirect_header_stripper( return _strip_on_cross_origin_redirect -def _exc_causes(exc: BaseException) -> List[BaseException]: - """``__cause__`` then ``__context__`` of *exc*, when they are exceptions.""" - return [nested for nested in (exc.__cause__, exc.__context__) if isinstance(nested, BaseException)] +def _exc_children(exc: BaseException) -> List[BaseException]: + """Sub-exceptions of a group, else ``__cause__``/``__context__`` when they are exceptions.""" + nested = getattr(exc, "exceptions", None) + if nested: + return list(nested) + return [c for c in (exc.__cause__, exc.__context__) if isinstance(c, BaseException)] def _format_connect_error(exc: BaseException) -> str: """Render nested MCP connection errors into an actionable short message.""" def _find_missing(current: BaseException) -> Optional[str]: - nested = getattr(current, "exceptions", None) - if nested: - return next(filter(None, map(_find_missing, nested)), None) if isinstance(current, FileNotFoundError): if getattr(current, "filename", None): return str(current.filename) match = re.search(r"No such file or directory: '([^']+)'", str(current)) if match: return match.group(1) - return next(filter(None, map(_find_missing, _exc_causes(current))), None) + return next(filter(None, map(_find_missing, _exc_children(current))), None) def _flatten_messages(current: BaseException) -> List[str]: - nested = getattr(current, "exceptions", None) - if nested: - return [m for child in nested for m in _flatten_messages(child)] - text = str(current).strip() + # A group's own str() is opaque — only its children speak. + text = "" if getattr(current, "exceptions", None) else str(current).strip() messages = [text] if text else [] - for nested_exc in _exc_causes(current): - messages.extend(_flatten_messages(nested_exc)) + for child in _exc_children(current): + messages.extend(_flatten_messages(child)) return messages or [current.__class__.__name__] missing = _find_missing(exc) if missing: message = f"missing executable '{missing}'" if os.path.basename(missing) in {"npx", "npm", "node"}: - message += ( - " (ensure Node.js is installed and PATH includes its bin directory, " - "or set mcp_servers..command to an absolute path and include " - "that directory in mcp_servers..env.PATH)" - ) + message += (" (ensure Node.js is installed and PATH includes its bin directory, " + "or set mcp_servers..command to an absolute path and include " + "that directory in mcp_servers..env.PATH)") return _sanitize_error(message) deduped = list(dict.fromkeys(_flatten_messages(exc))) return _sanitize_error("; ".join(deduped[:3])) @@ -343,31 +297,21 @@ def _optional_types(module: str, *names: str) -> list: def _http_status_error_types() -> tuple: - """``HTTPStatusError`` classes from both httpx flavours. - - A 401 may come from the SDK's own stack (``httpx2`` on mcp >= 2.0) or from - Hermes' pinned ``httpx``; the classes are unrelated, so both go in the tuple. - """ + """``HTTPStatusError`` classes from both httpx flavours: a 401 may come from the SDK's own + stack (``httpx2`` on mcp >= 2.0) or from Hermes' pinned ``httpx``; the classes are unrelated.""" global _HTTP_STATUS_ERROR_TYPES if _HTTP_STATUS_ERROR_TYPES is None: - found: list = [] sdk_mod = _core.sdk_httpx() - if sdk_mod is not None: - found.append(sdk_mod.HTTPStatusError) - for cls in _optional_types("httpx", "HTTPStatusError"): - if cls not in found: - found.append(cls) + found: list = [sdk_mod.HTTPStatusError] if sdk_mod is not None else [] + found += [cls for cls in _optional_types("httpx", "HTTPStatusError") if cls not in found] _HTTP_STATUS_ERROR_TYPES = tuple(found) return _HTTP_STATUS_ERROR_TYPES def _get_auth_error_types() -> tuple: - """Cached tuple of exception types indicating MCP OAuth failure. - - SDK ``OAuthFlowError``/``OAuthTokenError`` (+ legacy ``UnauthorizedError``), - our ``OAuthNonInteractiveError``, and ``HTTPStatusError`` from both httpx - flavours — the latter needs the 401 status check in :func:`_is_auth_error`. - """ + """Cached exception types indicating MCP OAuth failure: SDK ``OAuthFlowError``/``OAuthTokenError`` + (+ legacy ``UnauthorizedError``), our ``OAuthNonInteractiveError``, and ``HTTPStatusError`` from + both httpx flavours — the latter needs the 401 status check in :func:`_is_auth_error`.""" global _AUTH_ERROR_TYPES if not _AUTH_ERROR_TYPES: _AUTH_ERROR_TYPES = tuple( @@ -380,11 +324,8 @@ def _get_auth_error_types() -> tuple: def _is_auth_error(exc: BaseException) -> bool: - """True if ``exc`` indicates an MCP OAuth failure. - - ``HTTPStatusError`` counts only with status 401; other HTTP errors fall - through to the generic error path. - """ + """True if ``exc`` indicates an MCP OAuth failure. ``HTTPStatusError`` counts only with + status 401; other HTTP errors fall through to the generic error path.""" types = _get_auth_error_types() if not types or not isinstance(exc, types): return False @@ -394,39 +335,33 @@ def _is_auth_error(exc: BaseException) -> bool: return True -# Lower-cased substrings meaning the server-side transport session expired / -# was GC'd. The OAuth token is still valid — only the transport needs rebuilding. +# Lower-cased substrings meaning the server-side transport session expired / was GC'd. +# The OAuth token is still valid — only the transport needs rebuilding. _SESSION_EXPIRED_MARKERS: tuple = ( "invalid or expired session", "expired session", "session expired", "session not found", "unknown session", "session terminated", "closedresourceerror", "closed resource", "transport is closed", "connection closed", "broken pipe", "end of file", ) - -# Node budget for ``_is_session_expired_error``. The visited set breaks cycles; -# the budget bounds pathological acyclic graphs. Kept well above -# ``sys.getrecursionlimit()`` so deep task-group nesting is still fully scanned. +# Node budget for ``_is_session_expired_error``. The visited set breaks cycles; the budget +# bounds pathological acyclic graphs. Kept well above ``sys.getrecursionlimit()`` so deep +# task-group nesting is still fully scanned. _EXC_TRAVERSAL_MAX_NODES = 10_000 def _is_session_expired_error(exc: BaseException) -> bool: - """True if ``exc`` looks like an MCP transport session expiry. - - Streamable-HTTP servers GC session state (idle TTL, restart, pod rotation) - while the OAuth token stays valid, so unlike :func:`_is_auth_error` the fix - is a transport reconnect (``_reconnect_event``), not an OAuth refresh. - """ + """True if ``exc`` looks like an MCP transport session expiry. Streamable-HTTP servers GC + session state (idle TTL, restart, pod rotation) while the OAuth token stays valid, so unlike + :func:`_is_auth_error` the fix is a transport reconnect (``_reconnect_event``), not an OAuth refresh.""" # AnyIO stream exceptions are often message-less (``str(ClosedResourceError()) == ""``), # so type checks are needed in addition to marker matching. - transport_error_types = tuple( - _optional_types("anyio", "BrokenResourceError", "ClosedResourceError", "EndOfStream") - ) + transport_error_types = tuple(_optional_types("anyio", "BrokenResourceError", "ClosedResourceError", "EndOfStream")) - # Iterative traversal over ``exceptions`` / ``__cause__`` / ``__context__`` - # with an identity-visited set AND a node budget (graphs can be deep or - # cyclic). Every reachable node is inspected so an InterruptedError anywhere - # overrides transport markers; the chain walk matters because SDK wrappers - # often raise a generic RuntimeError *from* the message-less ClosedResourceError. + # Iterative traversal over ``exceptions`` / ``__cause__`` / ``__context__`` with an + # identity-visited set AND a node budget (graphs can be deep or cyclic). Every reachable + # node is inspected so an InterruptedError anywhere overrides transport markers; the chain + # walk matters because SDK wrappers often raise a generic RuntimeError *from* the + # message-less ClosedResourceError. stack: "list[BaseException | None]" = [exc] seen: set[int] = set() transport_error_found = False @@ -437,19 +372,13 @@ def _is_session_expired_error(exc: BaseException) -> bool: continue seen.add(id(current)) budget -= 1 - if isinstance(current, InterruptedError): return False - # Messages vary across SDK versions and servers: match a narrow - # allow-list of stable substrings, not exception type, to avoid false positives. + # Messages vary across SDK versions and servers: match a narrow allow-list of stable + # substrings, not exception type, to avoid false positives. msg = str(current).lower() - if isinstance(current, transport_error_types) or ( - msg and any(marker in msg for marker in _SESSION_EXPIRED_MARKERS) - ): + if isinstance(current, transport_error_types) or (msg and any(marker in msg for marker in _SESSION_EXPIRED_MARKERS)): transport_error_found = True - stack.extend(getattr(current, "exceptions", ())) - stack.append(getattr(current, "__cause__", None)) - stack.append(getattr(current, "__context__", None)) - + stack.extend((getattr(current, "__cause__", None), getattr(current, "__context__", None))) return transport_error_found From 5a8ffa517930b7b3546e6706798a3a262cec2458 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 17:07:08 -0700 Subject: [PATCH 3/4] =?UTF-8?q?refactor(mcp):=20health=20=E2=80=94=20recyc?= =?UTF-8?q?le=20reason=20via=20next(),=20change-list=20comprehension=20in?= =?UTF-8?q?=20=5Frefresh=5Ftools,=20docstrings=20compacted=20by=20hand=20(?= =?UTF-8?q?380=20->=20316)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/mcp_tool_health.py | 230 ++++++++++++++------------------------- 1 file changed, 83 insertions(+), 147 deletions(-) diff --git a/tools/mcp_tool_health.py b/tools/mcp_tool_health.py index 4d5c611492..a2dbd242ae 100644 --- a/tools/mcp_tool_health.py +++ b/tools/mcp_tool_health.py @@ -16,8 +16,8 @@ _KEEPALIVE_RPC_TIMEOUT = 30.0 def _stdio_children_dead_impl(pids, is_http: bool) -> bool: - """True when every pid has exited. Best-effort: False (unknown → don't fail - fast) for HTTP, no captured PIDs, missing psutil, or a failed probe.""" + """True when every pid has exited. Best-effort: False (unknown → don't fail fast) for HTTP, + no captured PIDs, missing psutil, or a failed probe.""" if not pids or is_http: return False try: @@ -25,9 +25,8 @@ def _stdio_children_dead_impl(pids, is_http: bool) -> bool: except ImportError: return False for pid in pids: - # pid_exists handles Windows without signal-permission noise. try: - if psutil.pid_exists(pid): + if psutil.pid_exists(pid): # handles Windows without signal-permission noise return False except Exception: return False @@ -40,11 +39,10 @@ class MCPServerHealthMixin: __slots__ = () def _is_http(self) -> bool: - """Check if this server uses HTTP transport.""" return "url" in self._config def _is_recycled_stdio(self) -> bool: - """Return True when a stdio server was intentionally recycled.""" + """True when a stdio server was intentionally recycled.""" return not self._is_http() and self._recycled_reason is not None def mark_tool_call(self) -> None: @@ -52,16 +50,14 @@ class MCPServerHealthMixin: self._last_tool_call_at = time.monotonic() def _mark_lifecycle_started(self) -> None: - now = time.monotonic() - self._lifecycle_started_at = now - self._last_tool_call_at = now + self._lifecycle_started_at = self._last_tool_call_at = time.monotonic() self._recycled_reason = None # ------------------------------------------------------- stdio recycling def _stdio_recycle_deadlines(self): - """``[(deadline, reason), ...]`` for the configured lifetime/idle limits; - empty for HTTP servers or while an RPC holds the lock.""" + """``[(deadline, reason), ...]`` for the configured lifetime/idle limits; empty for HTTP + servers or while an RPC holds the lock.""" if self._is_http() or self._rpc_lock.locked(): return [] deadlines = [] @@ -72,15 +68,12 @@ class MCPServerHealthMixin: return deadlines def _stdio_recycle_reason(self, now: Optional[float] = None) -> Optional[str]: - """Return the stdio recycle reason if idle/age limits have elapsed (lifetime wins).""" + """The stdio recycle reason if idle/age limits have elapsed (lifetime wins), else None.""" now = time.monotonic() if now is None else now - for deadline, reason in self._stdio_recycle_deadlines(): - if now >= deadline: - return reason - return None + return next((reason for deadline, reason in self._stdio_recycle_deadlines() if now >= deadline), None) def _next_stdio_recycle_deadline(self) -> Optional[float]: - """Return the next monotonic recycle deadline for stdio, if any.""" + """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 @@ -106,13 +99,11 @@ class MCPServerHealthMixin: return task def _make_logging_callback(self): - """Build a ``logging_callback`` that forwards server ``notifications/message`` - into Hermes logging tagged with the server name (the SDK default drops them).""" + """``logging_callback`` forwarding server ``notifications/message`` into Hermes logging + tagged with the server name (the SDK default drops them).""" async def _on_log(params): try: - level = _core._MCP_LOG_LEVEL_MAP.get( - str(getattr(params, "level", "info")).lower(), logging.INFO, - ) + level = _core._MCP_LOG_LEVEL_MAP.get(str(getattr(params, "level", "info")).lower(), logging.INFO) data = getattr(params, "data", None) if not isinstance(data, str): try: @@ -130,31 +121,27 @@ class MCPServerHealthMixin: return _on_log def _make_message_handler(self): - """Build a ``message_handler`` for ``ClientSession``: only - ``ToolListChangedNotification`` triggers a refresh; prompt/resource changes are logged.""" + """``message_handler`` for ``ClientSession``: only ``ToolListChangedNotification`` triggers + a refresh; prompt/resource changes are logged.""" async def _handler(message): try: 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. + # 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. + 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. + # 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) @@ -167,8 +154,8 @@ class MCPServerHealthMixin: 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* that 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}" @@ -179,69 +166,44 @@ class MCPServerHealthMixin: _forget_mcp_tool_server(tool_name) async def _refresh_tools(self): - """Re-fetch tools on ``tools/list_changed`` and update the registry. - - The lock serializes rapid-fire notifications. After the list_tools - ``await``, all mutations are synchronous — atomic on the event loop. - """ + """Re-fetch tools on ``tools/list_changed`` and update the registry. The lock serializes + rapid-fire notifications; after the list_tools ``await`` all mutations are synchronous — + atomic on the event loop.""" if not self._advertises_tools(): - # Shouldn't happen, but tools/list would raise MCPError(-32601). - return - + return # tools/list would raise MCPError(-32601) async with self._refresh_lock: old_tool_names = set(self._registered_tool_names) - - # 1. Fetch the current tool list (follow nextCursor). async with self._rpc_lock: new_mcp_tools = await _core._paginate_full_list(self.session.list_tools, "tools", self.name) - - # 2. Remove only stale names first — no nuke-and-repave: live agent - # turns may hold tool-call IDs pointing at existing handlers, and - # in-place replacement avoids transient "tool not connected" races. - self._deregister_owned(old_tool_names - { - mcp_prefixed_tool_name(self.name, tool.name) for tool in new_mcp_tools - }) - - # 3. Re-register; the helper may skip names ambiguous after normalization. + # Remove only stale names first — no nuke-and-repave: live agent turns may hold + # tool-call IDs pointing at existing handlers, and in-place replacement avoids + # transient "tool not connected" races. + self._deregister_owned(old_tool_names - {mcp_prefixed_tool_name(self.name, tool.name) for tool in new_mcp_tools}) + # Re-register; the helper may skip names ambiguous after normalization. A raw name + # can become ambiguous without changing its normalized name, so the pre-pass misses + # it: drop any old entry the final collision-checked registration no longer owns. self._tools = new_mcp_tools registered_names = _core._register_server_tools(self.name, self, self._config) - # A raw name can become ambiguous without changing its normalized - # name, so the pre-pass misses it: drop any old entry the final - # collision-checked registration no longer owns. self._deregister_owned(old_tool_names - set(registered_names)) self._registered_tool_names = registered_names - - # 4. Log what changed (user-visible). - new_tool_names = set(self._registered_tool_names) - added = new_tool_names - old_tool_names - removed = old_tool_names - new_tool_names - changes = [] - if added: - changes.append(f"added: {', '.join(sorted(added))}") - if removed: - changes.append(f"removed: {', '.join(sorted(removed))}") + # 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] if changes: - logger.warning( - "MCP server '%s': tools changed dynamically — %s. " - "Verify these changes are expected.", - self.name, "; ".join(changes), - ) + logger.warning("MCP server '%s': tools changed dynamically — %s. " + "Verify these changes are expected.", self.name, "; ".join(changes)) else: - logger.info( - "MCP server '%s': dynamically refreshed %d tool(s) (no changes)", - self.name, len(self._registered_tool_names), - ) + logger.info("MCP server '%s': dynamically refreshed %d tool(s) (no changes)", + self.name, len(self._registered_tool_names)) # ------------------------------------------------------ keepalive / health async def _keepalive_probe(self) -> None: - """Exercise the session; raise on a genuine connection failure. - - ``ping`` first (cheap, OPTIONAL utility). On -32601 latch - ``_ping_unsupported`` and fall back to ``list_tools`` when the server - advertises tools; otherwise the -32601 propagates (no liveness primitive - left). The latch resets on each fresh transport connection. - """ + """Exercise the session; raise on a genuine connection failure. ``ping`` first (cheap, + OPTIONAL utility). On -32601 latch ``_ping_unsupported`` and fall back to ``list_tools`` + when the server advertises tools; otherwise the -32601 propagates (no liveness primitive + left). The latch resets on each fresh transport connection.""" if not self._ping_unsupported: try: await asyncio.wait_for(self.session.send_ping(), timeout=_KEEPALIVE_RPC_TIMEOUT) @@ -252,74 +214,55 @@ class MCPServerHealthMixin: if not self._advertises_tools(): raise self._ping_unsupported = True - logger.info( - "MCP server '%s': does not implement the optional 'ping' utility (-32601); " - "using 'list_tools' for keepalive on this connection.", - self.name, - ) + logger.info("MCP server '%s': does not implement the optional 'ping' utility (-32601); " + "using 'list_tools' for keepalive on this connection.", self.name) elif isinstance(exc, (TimeoutError, asyncio.TimeoutError)) and self._advertises_tools(): - # A server that silently drops ping looks like a dead transport. - # Confirm with list_tools before declaring it dead; if that - # also fails, propagate the original failure. + # A server that silently drops ping looks like a dead transport. Confirm with + # list_tools before declaring it dead; if that also fails, propagate the + # original failure. try: 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 - 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, - ) + 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: - # Closed transport, expired session, etc. — real failure. - raise - + 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 + instead of respawning forever.""" if self._session_proven: return self._session_proven = True self._reconnect_retries = 0 if self._was_parked: self._was_parked = False - logger.warning( - "MCP server '%s': revived — session healthy again after " - "parking (state: parked → connected)", - self.name, - ) - # A proven fresh transport clears the one-time permanent-failure - # grace and any race bookkeeping. + logger.warning("MCP server '%s': revived — session healthy again after " + "parking (state: parked → connected)", self.name) + # A proven fresh transport clears the one-time permanent-failure grace and any race bookkeeping. self._permanent_grace_used = False self._teardown_race = False def mark_suspect(self, reason: str) -> None: - """Latch a suspicion (no I/O). The NEXT call verifies via - :meth:`ensure_healthy` and recycles the transport if the probe fails.""" + """Latch a suspicion (no I/O). The NEXT call verifies via :meth:`ensure_healthy` and + recycles the transport if the probe fails.""" if self._suspect_reason is None and reason: - logger.warning( - "MCP server '%s': connection marked suspect (%s); next call will health-check it", - self.name, reason, - ) + logger.warning("MCP server '%s': connection marked suspect (%s); next call will health-check it", + self.name, reason) self._suspect_reason = reason or None async def ensure_healthy(self, timeout: float = 5.0) -> bool: - """Verify a suspect connection before reuse; recycle if dead. - - True when healthy (suspicion cleared). On failure requests a reconnect, - drops the stale session so the caller's no-session path takes over, and - returns False. Never raises. - """ + """Verify a suspect connection before reuse; recycle if dead. True when healthy (suspicion + cleared). On failure requests a reconnect, drops the stale session so the caller's + no-session path takes over, and returns False. Never raises.""" reason = self._suspect_reason if not reason: return True @@ -332,34 +275,27 @@ class MCPServerHealthMixin: await asyncio.wait_for(self._keepalive_probe(), timeout=timeout) except Exception as exc: root = _unwrap_exception_group(exc) - logger.warning( - "MCP server '%s': suspect connection (%s) failed health check (%s: %s) — " - "requesting reconnect (state: suspect → degraded)", - self.name, reason, type(root).__name__, root, - ) + logger.warning("MCP server '%s': suspect connection (%s) failed health check (%s: %s) — " + "requesting reconnect (state: suspect → degraded)", + self.name, reason, type(root).__name__, root) self._suspect_reason = None self.mark_suspect(f"health check failed after {reason}") self.session = None self._ready.clear() self._reconnect_event.set() return False - logger.info( - "MCP server '%s': suspect connection passed health check (%s) — clearing suspicion", - self.name, reason, - ) + logger.info("MCP server '%s': suspect connection passed health check (%s) — clearing suspicion", + self.name, reason) self._suspect_reason = None self._mark_session_proven() return True 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 - ``_teardown_race`` so run() treats the next reconnect as recovery - rather than charging the rapid-drop budget. - """ + """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 + ``_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()] if not victims: return @@ -374,7 +310,7 @@ class MCPServerHealthMixin: return _stdio_children_dead_impl(getattr(self, "_stdio_child_pids", None), self._is_http()) async def _watch_stdio_children(self) -> None: - """Poll child liveness while a stdio RPC is in flight; resolves when a - tracked child dies so the caller cancels the RPC instead of waiting out the timeout.""" + """Poll child liveness while a stdio RPC is in flight; resolves when a tracked child dies + so the caller cancels the RPC instead of waiting out the timeout.""" while not self._stdio_children_dead(): await asyncio.sleep(0.25) From 493c3841ede92402c64734e27e12241b79310859 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 17:07:42 -0700 Subject: [PATCH 4/4] =?UTF-8?q?refactor(mcp):=20handlers=20=E2=80=94=20fol?= =?UTF-8?q?d=20utility=20on=5Ffailure=20into=20an=20inline=20lambda,=20sin?= =?UTF-8?q?gle=20guard=20in=20reconnectable=20lookup=20(602=20->=20598)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tools/mcp_tool_handlers.py | 10 +++------- 1 file changed, 3 insertions(+), 7 deletions(-) diff --git a/tools/mcp_tool_handlers.py b/tools/mcp_tool_handlers.py index 9b121ae3ab..3beebc7ddb 100644 --- a/tools/mcp_tool_handlers.py +++ b/tools/mcp_tool_handlers.py @@ -115,9 +115,7 @@ def _lookup_reconnectable_server(server_name: str, require_loop: bool = False): With *require_loop*, also None unless the MCP loop is running (nothing to wait on).""" with _core._lock: srv = _core._servers.get(server_name) - if srv is None or not hasattr(srv, "_reconnect_event"): - return None - if require_loop and not _mcp_loop_running(): + if srv is None or not hasattr(srv, "_reconnect_event") or (require_loop and not _mcp_loop_running()): return None return srv @@ -481,12 +479,10 @@ def _make_utility_handler(server_name: str, tool_timeout: float, op: str, log_la result = await rpc(server.session, args) return json.dumps(render(result, server_name), ensure_ascii=False) - def _on_failure(exc): - logger.error("MCP %s/%s failed: %s", server_name, log_label, exc) - return _invoke_with_recovery( server_name, lambda: _core._run_on_mcp_loop(_call, timeout=tool_timeout), op, - (_handle_auth_error_and_retry, _handle_session_expired_and_retry), _on_failure, + (_handle_auth_error_and_retry, _handle_session_expired_and_retry), + lambda exc: logger.error("MCP %s/%s failed: %s", server_name, log_label, exc), ) return _handler