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:
+104
-185
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user