From 4efaf9ecd43a5ce4542fff2bae3aff38284672e1 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 23:28:07 -0700 Subject: [PATCH] refactor(tools): fold single-use MCP helpers, collapse defensive layers, compact docstrings _stdio_children_dead_impl/_refresh_tools_task folded into their methods, _recover_401/_is_invalid_client_at_token_endpoint defensive getattr chains collapsed, lifecycle pid ledgers and drain loop tightened, WHY-preserving docstring compaction across the group. Schemas byte-identical. --- tools/mcp_oauth_manager.py | 82 ++++++++++---------------- tools/mcp_oauth_provider.py | 36 ++++-------- tools/mcp_tool_agent.py | 55 +++++++---------- tools/mcp_tool_config.py | 51 ++++++---------- tools/mcp_tool_content.py | 44 ++++---------- tools/mcp_tool_health.py | 89 +++++++++++----------------- tools/mcp_tool_lifecycle.py | 114 +++++++++++++++--------------------- 7 files changed, 175 insertions(+), 296 deletions(-) diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index ae3945ca71..c0141159a8 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -53,10 +53,9 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): 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 POST; HTTPX may close the generator from another task). + # A binary semaphore keeps mutual exclusion without task ownership. import anyio self.context.lock = anyio.Semaphore(1, max_value=1) self._hermes_server_name = server_name @@ -75,14 +74,12 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): logger.debug("MCP OAuth '%s': %s failed (non-fatal): %s", self._hermes_server_name, what, exc) async def _initialize(self) -> None: - """Load stored state, seed ``token_expiry_time``, restore/prefetch metadata. - - 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`` + """Load stored state, seed ``token_expiry_time``, restore/prefetch metadata. 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; + seeding the expiry (``HermesTokenStorage`` persists absolute ``expires_at``) makes the SDK + refresh first. 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 @@ -113,8 +110,7 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): 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, - ) + handle_auth_metadata_response, handle_protected_resource_response) server_url = self.context.server_url async def _send(client, url: str, label: str): @@ -167,14 +163,12 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): ``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 + token_endpoint = getattr(getattr(self.context, "oauth_metadata", None), "token_endpoint", 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: + if not token_endpoint or req is None: return False try: - pa, pb = urlsplit(req_url), urlsplit(token_endpoint) + pa, pb = urlsplit(str(req.url)), urlsplit(str(token_endpoint)) except ValueError: # pragma: no cover — malformed URL return False if not (pa.scheme == pb.scheme and pa.netloc.lower() == pb.netloc.lower() @@ -184,15 +178,12 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): return re.search(rb"\binvalid_client\b", body.lower()) is not None 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 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 + """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 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``.""" + in the body; pre-registered clients are never poisoned; any failure is swallowed. The + browser-side "Redirect URI Mismatch" case has no HTTP signal (``hermes mcp reauth``).""" try: if self._hermes_preregistered or getattr(response, "status_code", None) not in (400, 401): return @@ -226,9 +217,8 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): 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. + # Bridge the bidirectional generator by hand: a naive ``async for item in inner: yield + # item`` DISCARDS the responses httpx sends back via ``asend``, and the SDK crashes on None. inner = super().async_auth_flow(request) resource_lock_released = False sent_access_token = None @@ -237,8 +227,7 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): 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 MCP request. Release it for that request only; OAuth transitions stay serialized. if outgoing is request: tokens = self.context.current_tokens sent_access_token = tokens.access_token if tokens is not None else None @@ -249,8 +238,7 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): 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. + # retry with that token instead of a duplicate OAuth transition from a 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): @@ -262,7 +250,6 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): outgoing = await inner.asend(incoming) except StopAsyncIteration: 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 @@ -274,7 +261,6 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): if retry_after_concurrent_auth: yield request self._persist_oauth_metadata_if_changed() - return # Cached at import time; None when the SDK's OAuth module is unavailable. @@ -282,9 +268,8 @@ _HERMES_PROVIDER_CLS: Optional[type] = HermesMCPOAuthProvider if _SDK_BASES else class MCPOAuthManager: - """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).""" + """Single source of truth for per-server MCP OAuth state. ``_entries`` is guarded by + ``_entries_lock`` (get-or-create); per-entry state by the entry's ``asyncio.Lock``.""" def __init__(self) -> None: self._entries: dict[tuple[str, str], _ProviderEntry] = {} @@ -294,9 +279,8 @@ class MCPOAuthManager: self._inflight_tasks: set[asyncio.Task] = set() def get_or_build_provider(self, server_name: str, server_url: str, oauth_config: Optional[dict]) -> Optional[Any]: - """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.""" + """Cached OAuth provider for ``server_name``, built on first use (rebuilt when + ``server_url`` changes). None if the MCP SDK's OAuth support is unavailable.""" key = self._key(server_name) with self._entries_lock: entry = self._entries.get(key) @@ -388,10 +372,9 @@ class MCPOAuthManager: # 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 = bool(entry.provider.context.can_refresh_token()) + except Exception: # no context / not callable / probe failed can_refresh = False if not pending.done(): pending.set_result(can_refresh) @@ -403,11 +386,10 @@ 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 — 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.""" + """Handle a 401 from a tool call. True: a (possibly new) token is available — reconnect + and retry. False: no recovery path — surface ``needs_reauth`` so the model stops + hallucinating manual refreshes. 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 diff --git a/tools/mcp_oauth_provider.py b/tools/mcp_oauth_provider.py index b8d3cea01e..cfe6247f3d 100644 --- a/tools/mcp_oauth_provider.py +++ b/tools/mcp_oauth_provider.py @@ -13,25 +13,19 @@ from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from tools.mcp_oauth import HermesTokenStorage - logger = logging.getLogger(__name__) class HermesProviderMixin: - """Token-endpoint fixes layered over the SDK's ``OAuthClientProvider``. + """Token-endpoint fixes layered over the SDK's ``OAuthClientProvider`` (must precede it in + the MRO; subclasses set ``_hermes_logger`` to keep their own logger name). - 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`` 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. - """ + ``token_endpoint_auth_method``; the SDK then treats the client as public and the token + endpoint rejects the exchange (looping the browser page) — coerce ``client_secret_post``. + - ``token_user_agent`` (``oauth.user_agent``) is stamped onto token-endpoint requests only + (some authorization servers/WAFs reject httpx's default). + - Any 2xx token/refresh response is accepted; token bodies never leak into errors/logs.""" _hermes_logger: logging.Logger = logger @@ -54,7 +48,6 @@ class HermesProviderMixin: return from mcp.shared.auth import OAuthClientInformationFull from tools.mcp_oauth import HermesTokenStorage - data = info.model_dump(mode="json", exclude_none=True) if HermesTokenStorage._coerce_secret_auth_method(data): self.context.client_info = OAuthClientInformationFull.model_validate(data) @@ -75,12 +68,10 @@ class HermesProviderMixin: async def _handle_token_response(self, response): """Accept any 2xx token response; never echo the body into errors.""" from mcp.client.auth.oauth2 import OAuthTokenError - if not (200 <= response.status_code < 300): raise OAuthTokenError(f"Token exchange failed ({response.status_code})") from httpx import HTTPError from mcp.client.auth.utils import handle_token_response_scopes - try: token_response = await handle_token_response_scopes(response) except (HTTPError, OAuthTokenError): @@ -96,7 +87,6 @@ class HermesProviderMixin: from httpx import HTTPError from mcp.shared.auth import OAuthToken from pydantic import ValidationError - try: token_response = OAuthToken.model_validate_json(await response.aread()) except (HTTPError, ValidationError): @@ -112,21 +102,17 @@ def prepare_oauth_config(server_name: str, server_url: str, oauth_config: dict | 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 {}) mo.apply_oauth_provider_defaults(cfg, server_name=server_name, server_url=server_url) return cfg, mo.HermesTokenStorage(server_name) 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 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.""" + """Resolve the callback port and return the shared provider constructor kwargs. 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`` so tests can patch them.""" from tools import mcp_oauth as mo - port = mo._configure_callback_port(cfg, storage) client_metadata = mo._build_client_metadata(cfg) mo._maybe_preregister_client(storage, cfg, client_metadata) diff --git a/tools/mcp_tool_agent.py b/tools/mcp_tool_agent.py index ac3407e4e9..6ff6e27ee6 100644 --- a/tools/mcp_tool_agent.py +++ b/tools/mcp_tool_agent.py @@ -33,8 +33,7 @@ def _resolve_refresh_toolsets(agent, enabled_override, disabled_override): if enabled_override is not None or disabled_override is not None: enabled = enabled_override if enabled_override is not None else enabled disabled = disabled_override if disabled_override is not None else disabled - agent.enabled_toolsets = enabled - agent.disabled_toolsets = disabled + agent.enabled_toolsets, agent.disabled_toolsets = enabled, disabled return enabled, disabled @@ -50,8 +49,7 @@ def _tool_defs_content_changed(agent, new_defs: list) -> bool: 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]: + 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 won).""" @@ -85,35 +83,27 @@ 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). + newly-added tool names (empty when unchanged). The agent snapshots ``agent.tools`` at build + time, so servers that connect later (slow OAuth, ``/reload-mcp``) are invisible until + rebuilt. Shared by the TUI RPC, gateway reload, late-binding thread and between-turns + refresh: respects the toolset filter, diffs by tool NAME (a count compare misses an + equal-size swap), re-injects the memory-provider / context-engine tools ``agent_init`` + appends after ``get_tool_definitions``, publishes ``(tools, valid_tool_names)`` together. - 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 (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).""" + ``preserve_prefix``: for rebuilds inside a live conversation 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 (``check_fn`` gates exposure, never invocation), a deregistered tool is + dropped, new tools append at the tail. The caller owns the prompt-cache contract.""" 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 already published. + # Generation captured BEFORE the slow get_tool_definitions call (a slower caller holding an + # OLDER set must not clobber a newer one); definitions computed OUTSIDE the lock. 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 []) 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. + # Post-build families re-appended on LOCALS only; live attributes untouched until 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. @@ -128,9 +118,7 @@ def refresh_agent_mcp_tools( 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. - persist_agent_tool_names(agent) + persist_agent_tool_names(agent) # re-pin so a rebuild after agent-cache eviction restores this order return added @@ -140,7 +128,6 @@ def reprobe_tool_availability() -> None: otherwise replay the stale verdicts).""" from model_tools import _clear_tool_defs_cache from tools.registry import invalidate_check_fn_cache - invalidate_check_fn_cache() _clear_tool_defs_cache() @@ -159,14 +146,12 @@ 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 (same merge rule as - ``_merge_preserving_prefix``; a saved tool still registered but failing its probe is - carried forward from the registry schema).""" + After agent-cache eviction the gateway rebuilds a NEW AIAgent with no predecessor to + preserve, so the saved name list stands in (``_merge_preserving_prefix`` rule; 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 - fresh_defs = _agent_tool_defs(agent) fresh = {_def_name(t): t for t in fresh_defs} diff --git a/tools/mcp_tool_config.py b/tools/mcp_tool_config.py index 72b1e98d9d..0e39bf1a50 100644 --- a/tools/mcp_tool_config.py +++ b/tools/mcp_tool_config.py @@ -66,7 +66,8 @@ _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"\$\{([^}]+)\}") @@ -77,7 +78,6 @@ def _workspace_folder() -> str: (terminal cwd / task override / $TERMINAL_CWD), else cwd.""" try: from tools.file_tools import _authoritative_workspace_root - root = _authoritative_workspace_root() if root: return root @@ -97,7 +97,8 @@ _CONTEXT_VAR_RESOLVERS = { "workspaceFolder": lambda: _core._workspace_folder(), "workspaceFolderBasename": _workspace_basename, "pathSeparator": lambda: os.sep, - "/": lambda: os.sep} + "/": lambda: os.sep, +} def _build_safe_env(user_env: Optional[dict]) -> dict: @@ -142,9 +143,8 @@ def _node_fallback(command: str) -> str: 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 hand-authored env.PATH omits it: npx's shebang re-execs - # /usr/bin/env node, so a symlink workaround fails one layer deeper. + # Canonical Node location (from-source Linux, Hermes Docker image, Intel Homebrew). Needed + # when a hand-authored env.PATH omits it: npx's shebang re-execs /usr/bin/env node. 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): @@ -189,18 +189,14 @@ 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. 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 and context vars. Env + refs resolve from the active profile's secret scope when multiplexing (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): 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 resolver() if resolver is not None else (_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): return {k: _interpolate_env_vars(v) for k, v in value.items()} @@ -215,16 +211,14 @@ _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 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 failures, invisible in config.yaml). Advisory only: values + are never mutated (could be intentional) nor logged (often secrets). Returns flagged paths.""" flagged: List[str] = [] def _walk(value: Any, path: str) -> None: - if isinstance(value, str): - if value != value.strip(): - flagged.append(path) + if isinstance(value, str) and value != value.strip(): + flagged.append(path) elif isinstance(value, dict): for k, v in value.items(): _walk(v, f"{path}.{k}" if path else str(k)) @@ -234,10 +228,9 @@ def _warn_hidden_whitespace(server_name: str, config: dict) -> List[str]: _walk(config, "") for key_path in flagged: - dedupe_key = (server_name, key_path) - if dedupe_key in _whitespace_warned: + if (server_name, key_path) in _whitespace_warned: continue - _whitespace_warned.add(dedupe_key) + _whitespace_warned.add((server_name, key_path)) logger.warning( "MCP server '%s': config value '%s' has hidden leading or " "trailing whitespace — this often causes authentication or " @@ -269,7 +262,6 @@ def _portable_mcp_servers(safe_servers: Dict[str, dict]) -> None: on a name clash. Never raises.""" try: from hermes_cli.plugins import discover_plugins, get_plugin_manager - discover_plugins() portable = get_plugin_manager().get_portable_mcp_servers() for name, cfg in _core._filter_suspicious_mcp_servers(portable).items(): @@ -283,25 +275,20 @@ 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.""" + mode); ``${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 - if _env_enabled("HERMES_SAFE_MODE"): return {} servers = load_config().get("mcp_servers") - if not isinstance(servers, dict): - servers = {} try: # ensure .env vars are available for interpolation from hermes_cli.env_loader import load_hermes_dotenv load_hermes_dotenv() except Exception: pass safe_servers: Dict[str, dict] = {} - for name, cfg in _core._filter_suspicious_mcp_servers(servers).items(): + for name, cfg in _core._filter_suspicious_mcp_servers(servers if isinstance(servers, dict) else {}).items(): interpolated = _interpolate_env_vars(cfg) if isinstance(interpolated, dict): _warn_hidden_whitespace(name, interpolated) diff --git a/tools/mcp_tool_content.py b/tools/mcp_tool_content.py index 3eeed70a7a..826a144f06 100644 --- a/tools/mcp_tool_content.py +++ b/tools/mcp_tool_content.py @@ -13,14 +13,11 @@ 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. +# Hard ceiling for one MCP text payload (chars), deliberately far ABOVE the budget layer's 50K +# spillover threshold so ordinary large results reach spillover intact; only floods are lossy. _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. Base64 expands ~4/3; oversized payloads are rejected BEFORE decoding so a -# multi-GB blob string is never transiently doubled in memory. +# Cap on decoded resource bytes per block (a misbehaving server can't fill the cache disk). +# Base64 expands ~4/3; oversized payloads are rejected BEFORE decoding (never doubled in memory). _MCP_RESOURCE_MAX_BYTES = 50 * 1024 * 1024 _MCP_RESOURCE_MAX_B64_CHARS = _MCP_RESOURCE_MAX_BYTES * 4 // 3 + 4 @@ -33,11 +30,8 @@ def _truncate_mcp_text_result(text: str, max_chars: int = _MCP_HARD_RESULT_CAP_C head_chars = int(max_chars * 0.4) tail_chars = max_chars - head_chars omitted = len(text) - head_chars - tail_chars - return ( - text[:head_chars] - + f"\n\n... [MCP RESULT TRUNCATED - {omitted:,} chars omitted " - f"out of {len(text):,} total] ...\n\n" - + text[-tail_chars:]) + return (text[:head_chars] + f"\n\n... [MCP RESULT TRUNCATED - {omitted:,} chars omitted " + f"out of {len(text):,} total] ...\n\n" + text[-tail_chars:]) def _is_reserved_mcp_meta_key(key: str) -> bool: @@ -93,12 +87,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.""" + marker)``. Fail-open so one bad block never kills the tool result: gateway deps missing + (cron without gateway) → ``unavailable``; any other cache error → warning + ``failed``.""" try: import gateway.platforms.base as _base - return getattr(_base, writer)(*args, **kwargs), "" except ImportError: logger.debug("MCP %s caching skipped — gateway.platforms.base unavailable", skip_label) @@ -146,22 +138,18 @@ def _mcp_resource_filename(uri: str, mime_type: str) -> str: import re as _re from pathlib import Path from urllib.parse import urlparse, unquote - name = "" if uri: try: name = Path(unquote(urlparse(str(uri)).path or "")).name except (ValueError, TypeError): - name = "" + pass # 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(".") - if dot and 0 < len(ext) <= 12: - name = stem[: 150 - len(ext) - 1] + "." + ext - else: - name = name[:150] + name = stem[: 150 - len(ext) - 1] + "." + ext if dot and 0 < len(ext) <= 12 else name[:150] if not name or name in {".", ".."}: ext = mimetypes.guess_extension(_base_mime(mime_type)) or ".bin" name = f"resource{ext}" @@ -175,22 +163,15 @@ def _render_mcp_resource_block(block, server_name: str = "") -> str: 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"): + if block_type == "resource_link" or (hasattr(block, "uri") and not hasattr(block, "resource") and block_type != "text"): uri = getattr(block, "uri", None) if not uri: return "" name = getattr(block, "name", "") or "" mime = mcp_field(block, "mime_type", "mimeType", "") or "" - details = f"uri={uri}" - if name: - details += f", name={name}" - if mime: - details += f", mimeType={mime}" + details = f"uri={uri}" + (f", name={name}" if name else "") + (f", mimeType={mime}" if mime else "") 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) if resource is None: return "" @@ -200,7 +181,6 @@ def _render_mcp_resource_block(block, server_name: str = "") -> str: blob = getattr(resource, "blob", None) if blob is None: return "" - uri = str(getattr(resource, "uri", "") or "") mime = str(mcp_field(resource, "mime_type", "mimeType", "") or "") raw_bytes, err = _decode_block_b64( diff --git a/tools/mcp_tool_health.py b/tools/mcp_tool_health.py index 75cf1ce2fe..4d01fc0867 100644 --- a/tools/mcp_tool_health.py +++ b/tools/mcp_tool_health.py @@ -17,24 +17,6 @@ logger = logging.getLogger("tools.mcp_tool") _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.""" - if not pids or is_http: - return False - try: - import psutil - except ImportError: - return False - for pid in pids: - try: - if psutil.pid_exists(pid): # handles Windows without signal-permission noise - return False - except Exception: - return False - return True - - class MCPServerHealthMixin: """Methods of :class:`tools.mcp_tool.MCPServerTask` (mixed in; relies on its attributes).""" @@ -85,15 +67,15 @@ class MCPServerHealthMixin: # -- notifications / logs -- - async def _refresh_tools_task(self): - try: - await self._refresh_tools() - except Exception: - logger.exception("MCP server '%s': dynamic tool refresh failed", self.name) - def _schedule_tools_refresh(self) -> asyncio.Task: - """Schedule a background tool refresh and keep it strongly referenced.""" - task = asyncio.create_task(self._refresh_tools_task()) + """Schedule a background tool refresh (failures logged) and keep it strongly referenced.""" + async def _run(): + try: + await self._refresh_tools() + except Exception: + logger.exception("MCP server '%s': dynamic tool refresh failed", self.name) + + task = asyncio.create_task(_run()) self._pending_refresh_tasks.add(task) task.add_done_callback(self._pending_refresh_tasks.discard) return task @@ -129,19 +111,15 @@ class MCPServerHealthMixin: return 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. + # mcp 2.0 made ServerNotification a plain union (payload IS the message) instead of + # a RootModel (payload under ``.root``); without this unwrap refreshes silently 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. + # Separate task: refreshing synchronously inside the handler can wedge the stdio + # JSON-RPC stream when list_changed arrives while another request is in flight. self._schedule_tools_refresh() - # Yield one tick so short-lived notification contexts (and tests) can observe - # the scheduled refresh. - await asyncio.sleep(0) + await asyncio.sleep(0) # one tick so short-lived contexts (and tests) observe it elif isinstance(payload, _core.PromptListChangedNotification): logger.debug("MCP server '%s': prompts/list_changed (ignored)", self.name) elif isinstance(payload, _core.ResourceListChangedNotification): @@ -154,7 +132,6 @@ class MCPServerHealthMixin: """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}" for tool_name in tool_names: if registry.get_toolset_for_tool(tool_name) != toolset_name: @@ -172,13 +149,11 @@ class MCPServerHealthMixin: old_tool_names = set(self._registered_tool_names) async with self._rpc_lock: new_mcp_tools = await _core._paginate_full_list(self.session.list_tools, "tools", self.name) - # 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. + # Remove only stale names first — no nuke-and-repave: live turns may hold tool-call + # IDs pointing at existing handlers; in-place replacement avoids "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. + # Re-register; a raw name can become ambiguous after normalization without changing + # its normalized name, so also drop old entries the final registration no longer owns. self._tools = new_mcp_tools registered_names = _core._register_server_tools(self.name, self, self._config) self._deregister_owned(old_tool_names - set(registered_names)) @@ -197,9 +172,8 @@ class MCPServerHealthMixin: 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.""" + OPTIONAL); on -32601 latch ``_ping_unsupported`` (reset per transport connection) and fall + back to ``list_tools`` when the server advertises tools, else the -32601 propagates.""" if not self._ping_unsupported: try: await asyncio.wait_for(self.session.send_ping(), timeout=_KEEPALIVE_RPC_TIMEOUT) @@ -212,9 +186,8 @@ class MCPServerHealthMixin: 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, else propagate the original failure. try: await asyncio.wait_for(self.session.list_tools(), timeout=_KEEPALIVE_RPC_TIMEOUT) except Exception: @@ -282,11 +255,10 @@ class MCPServerHealthMixin: 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 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 @@ -297,7 +269,16 @@ class MCPServerHealthMixin: task.cancel() def _stdio_children_dead(self) -> bool: - return _stdio_children_dead_impl(getattr(self, "_stdio_child_pids", None), self._is_http()) + """True when every stdio child we spawned has exited. Best-effort: False (unknown → don't + fail fast) for HTTP, no captured PIDs, missing psutil, or a failed probe.""" + pids = getattr(self, "_stdio_child_pids", None) + if not pids or self._is_http(): + return False + try: + import psutil + return not any(psutil.pid_exists(pid) for pid in pids) # Windows-safe, no signal noise + except Exception: # missing psutil or failed probe → unknown → don't fail fast + return False async def _watch_stdio_children(self) -> None: """Poll child liveness while a stdio RPC is in flight; resolves when a tracked child dies diff --git a/tools/mcp_tool_lifecycle.py b/tools/mcp_tool_lifecycle.py index 7543035ad2..a8b409b196 100644 --- a/tools/mcp_tool_lifecycle.py +++ b/tools/mcp_tool_lifecycle.py @@ -10,21 +10,16 @@ from tools.mcp_tool_common import _core logger = logging.getLogger("tools.mcp_tool") -# Live stdio MCP children (pid -> server_name), added after connection and -# removed on normal shutdown, so they can be force-killed if SDK teardown fails. +# Live stdio MCP children (pid -> server_name), added after connection and removed on normal +# shutdown, so they can be force-killed if SDK teardown fails. _stdio_pids: Dict[int, str] = {} - -# PIDs that survived their session context exit (SDK teardown failed to kill -# them); detected in _run_stdio's finally, reaped by _kill_orphaned_mcp_children(). -# Kept separate from _stdio_pids so cleanup sweeps never race active sessions. +# PIDs that survived their session context exit (detected in _run_stdio's finally, reaped by +# _kill_orphaned_mcp_children). Separate from _stdio_pids so sweeps never race active sessions. _orphan_stdio_pids: set = set() _orphan_stdio_pid_servers: Dict[int, str] = {} - -# pid -> pgid captured at spawn. The SDK spawns children with -# start_new_session=True (PGID == PID); grandchildren inherit that PGID and -# keep it after the direct child exits, so killpg still reaches them. Tracked -# separately from _stdio_pids so the PGID survives the child's removal. -# Empty on Windows (os.getpgid is POSIX-only). +# pid -> pgid captured at spawn. The SDK spawns with start_new_session=True (PGID == PID); +# grandchildren keep that PGID after the direct child exits, so killpg still reaches them. +# Separate from _stdio_pids so the PGID survives the child's removal. Empty on Windows. _stdio_pgids: Dict[int, int] = {} @@ -53,53 +48,49 @@ def _snapshot_child_pids() -> set: return set() -# argv markers of non-MCP gateway children that can race into the snapshot -# delta during an MCP spawn (defense-in-depth; LSP/slash_worker already use -# start_new_session). Matched against argv[1:] because Python/Java children -# start with the interpreter path. +# argv markers of non-MCP gateway children that can race into the snapshot delta during an +# MCP spawn (defense-in-depth; LSP/slash_worker already use start_new_session). Matched against +# argv[1:] because Python/Java children start with the interpreter path. _NON_MCP_CHILD_CMDLINE_MARKERS: tuple[str, ...] = ( - "tui_gateway.slash_worker", - "tui_gateway.entry", - "-dorg.eclipse.equinox.launcher", # jdtls (legacy arg style) - "eclipse.jdt.ls", - "org.eclipse.equinox.launcher_") + "tui_gateway.slash_worker", "tui_gateway.entry", + "-dorg.eclipse.equinox.launcher", "eclipse.jdt.ls", "org.eclipse.equinox.launcher_", # jdtls +) def _filter_mcp_children(pids: set) -> set: - """Drop non-MCP children from a PID snapshot delta. Tracking a stray child in - _stdio_pgids is catastrophic if it lacks start_new_session: its pgid can be - the TUI parent's, so the shutdown killpg() would kill the TUI itself.""" + """Drop non-MCP children from a PID snapshot delta. Tracking a stray child in _stdio_pgids + is catastrophic if it lacks start_new_session: its pgid can be the TUI parent's, so the + shutdown killpg() would kill the TUI itself.""" if not pids: return pids try: import psutil except ImportError: return pids # keep all PIDs (prior behavior) - - def _is_mcp(pid: int) -> bool: + kept = set() + for pid in pids: try: argv = psutil.Process(pid).cmdline() except (psutil.NoSuchProcess, psutil.AccessDenied, OSError): - return False # raced away or zombie — cannot be our fresh server, unsafe to track - return not any(marker in arg for arg in argv[1:] for marker in _NON_MCP_CHILD_CMDLINE_MARKERS) - - return {pid for pid in pids if _is_mcp(pid)} + continue # raced away or zombie — cannot be our fresh server, unsafe to track + if not any(marker in arg for arg in argv[1:] for marker in _NON_MCP_CHILD_CMDLINE_MARKERS): + kept.add(pid) + return kept def _clear_connect_cooldowns() -> None: - """Drop connect-retry cooldowns: a restart must re-attempt every server - immediately, not honour a stale per-server backoff. Caller holds ``_core._lock``.""" + """Drop connect-retry cooldowns: a restart must re-attempt every server immediately, not + honour a stale per-server backoff. Caller holds ``_core._lock``.""" _core._server_connect_retry_after.clear() _core._server_connect_failures.clear() def shutdown_mcp_servers(*, scope: Optional[str] = None): - """Close MCP server connections (in parallel) and stop the background loop. - Each server Task is signalled to exit its own ``async with`` so the anyio - cancel-scope cleanup runs in the Task that opened it. ``scope`` restricts - teardown to one multiplexed profile's servers (its ``/reload-mcp`` must not - kill other profiles') and leaves the shared loop running if anything else is - still connected.""" + """Close MCP server connections (in parallel) and stop the background loop. Each server + Task is signalled to exit its own ``async with`` so the anyio cancel-scope cleanup runs in + the Task that opened it. ``scope`` restricts teardown to one multiplexed profile's servers + (its ``/reload-mcp`` must not kill other profiles') and leaves the shared loop running if + anything else is still connected.""" with _core._lock: 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] @@ -143,8 +134,8 @@ def _take_reapable_pids(include_active: bool, server_name: Optional[str]) -> tup with _core._lock: pids = _owned({opid: _orphan_stdio_pid_servers.get(opid, "orphan") for opid in _orphan_stdio_pids}) + _orphan_stdio_pids.difference_update(pids) for opid in pids: - _orphan_stdio_pids.discard(opid) _orphan_stdio_pid_servers.pop(opid, None) if include_active: active = _owned(_stdio_pids) @@ -156,18 +147,16 @@ def _take_reapable_pids(include_active: bool, server_name: Optional[str]) -> tup def _signal_mcp_process(pid: int, sig: int, server_name: str, pgid: Optional[int], my_pgid: Optional[int]) -> None: - """SIGTERM/SIGKILL via the spawn-time pgroup on POSIX (reaches reparented - grandchildren), falling back to a per-pid signal.""" + """SIGTERM/SIGKILL via the spawn-time pgroup on POSIX (reaches reparented grandchildren), + falling back to a per-pid signal.""" 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. - 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) + 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) else: try: killpg(pgid, sig) @@ -183,14 +172,12 @@ def _signal_mcp_process(pid: int, sig: int, server_name: str, pgid: Optional[int def _kill_orphaned_mcp_children(include_active: bool = False, server_name: Optional[str] = None) -> None: - """Best-effort reap of stdio MCP subprocesses: SIGTERM, wait 2s, SIGKILL survivors. - By default only ``_orphan_stdio_pids`` (PIDs that outlived their session - context) are reaped so concurrent cron jobs / live sessions are untouched; - ``include_active=True`` also kills every ``_stdio_pids`` entry and is only - for final shutdown after the MCP loop has stopped. ``server_name`` limits the - sweep to one server (stdio reconnects cleaning up their old transport).""" + """Best-effort reap of stdio MCP subprocesses: SIGTERM, wait 2s, SIGKILL survivors. By + default only ``_orphan_stdio_pids`` are reaped so concurrent cron jobs / live sessions are + untouched; ``include_active=True`` also kills every ``_stdio_pids`` entry and is only for + final shutdown after the MCP loop has stopped. ``server_name`` limits the sweep to one + server (stdio reconnects cleaning up their old transport).""" import signal as _signal - pids, pgids = _take_reapable_pids(include_active, server_name) if not pids: # skip the 2s sleep every MCP-free shutdown would otherwise pay return @@ -203,17 +190,13 @@ def _kill_orphaned_mcp_children(include_active: bool = False, server_name: Optio for pid, owner in pids.items(): _signal_mcp_process(pid, _signal.SIGTERM, owner, pgids.get(pid), my_pgid) logger.debug("Sent SIGTERM to orphaned MCP process %d (%s)", pid, owner) - time.sleep(2) - sigkill = getattr(_signal, "SIGKILL", _signal.SIGTERM) - # ``os.kill(pid, 0)`` is NOT a no-op on Windows; use the portable check. - from gateway.status import _pid_exists + from gateway.status import _pid_exists # ``os.kill(pid, 0)`` is NOT a no-op on Windows for pid, owner in pids.items(): - if not _pid_exists(pid): - continue # exited after SIGTERM - _signal_mcp_process(pid, sigkill, owner, pgids.get(pid), my_pgid) - logger.warning("Force-killed MCP process %d (%s) after SIGTERM timeout", pid, owner) + if _pid_exists(pid): # survived SIGTERM + _signal_mcp_process(pid, sigkill, owner, pgids.get(pid), my_pgid) + logger.warning("Force-killed MCP process %d (%s) after SIGTERM timeout", pid, owner) def _stop_mcp_loop_if_idle() -> bool: @@ -239,13 +222,8 @@ async def _drain_mcp_loop_tasks(*, timeout: Optional[float] = None) -> None: task.cancel() done, still_pending = await asyncio.wait(pending, timeout=timeout) for task in done: - try: - if not task.cancelled(): - task.exception() - except asyncio.CancelledError: - pass - except Exception as exc: - logger.debug("Pending MCP loop task ended during shutdown: %s", exc) + if not task.cancelled(): + task.exception() # mark retrieved so asyncio doesn't warn "exception was never retrieved" if still_pending: logger.warning("%d MCP loop task(s) still pending after %.1fs drain", len(still_pending), timeout)