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.
This commit is contained in:
Teknium
2026-09-02 23:28:07 -07:00
parent ee81b1abdd
commit 4efaf9ecd4
7 changed files with 175 additions and 296 deletions
+32 -50
View File
@@ -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
+11 -25
View File
@@ -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)
+20 -35
View File
@@ -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}
+19 -32
View File
@@ -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)
+12 -32
View File
@@ -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.<writer>(*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(
+35 -54
View File
@@ -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
+46 -68
View File
@@ -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)