refactor(tools): compact MCP oauth/config/health/lifecycle/agent/content modules

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