Merge branch 'simp/r3-33-F' into simp/r3-33

This commit is contained in:
Teknium
2026-09-03 00:25:07 -07:00
4 changed files with 444 additions and 774 deletions

View File

@@ -1,4 +1,6 @@
"""MCP connection/transport error classification: URL validation, TLS client certs, identity headers, redirect header stripping, exception-group unwrapping, auth/session-expired/method-not-found detection and connect-error formatting. Split from tools/mcp_tool.py."""
"""MCP connection/transport error classification: URL validation, TLS client certs, identity
headers, redirect header stripping, exception-group unwrapping, auth/session-expired/
method-not-found detection and connect-error formatting. Split from tools/mcp_tool.py."""
import asyncio
import errno
@@ -12,7 +14,6 @@ from tools.mcp_tool_common import _sanitize_error, _core
logger = logging.getLogger("tools.mcp_tool")
# Stateless (2026-07-28) servers reject a legacy ``initialize`` with
# UnsupportedProtocolVersion (-32022) or plain method-not-found.
_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION = -32022
@@ -24,13 +25,9 @@ def _jsonrpc_code(exc: BaseException):
def _jsonrpc_matches(exc: BaseException, code, codes: tuple, markers: tuple) -> bool:
"""Structural *code* in *codes*, else any lowercased *marker* in ``str(exc)``. Never
``isinstance`` on SDK exception types: they arrive wrapped in ExceptionGroups and drift
across generations."""
if code in codes:
return True
msg = str(exc).lower()
return bool(msg) and any(marker in msg for marker in markers)
"""Structural *code* in *codes*, else any *marker* in ``str(exc).lower()``. Never ``isinstance``
on SDK exception types: they arrive wrapped in ExceptionGroups and drift across generations."""
return code in codes or any(marker in str(exc).lower() for marker in markers)
def _handshake_rejected_as_modern(exc: BaseException) -> bool:
@@ -43,37 +40,32 @@ def _handshake_rejected_as_modern(exc: BaseException) -> bool:
def _is_method_not_found_error(exc: BaseException) -> bool:
"""True if *exc* is a JSON-RPC ``method not found`` (-32601). ``ping`` is optional in MCP;
servers lacking it answer -32601. The substring fallback includes "Unknown method: <name>",
which some servers use; without it the ping→list_tools keepalive fallback never latches
and reconnect-loops."""
"""True if *exc* is a JSON-RPC ``method not found`` (-32601; ``ping`` is optional in MCP). The
substring fallback includes "Unknown method: <name>" — without it the ping→list_tools keepalive
fallback never latches and reconnect-loops."""
return _jsonrpc_matches(
exc, _jsonrpc_code(exc), (_core._JSONRPC_METHOD_NOT_FOUND,),
(str(_core._JSONRPC_METHOD_NOT_FOUND), "method not found", "unknown method", "not found: ping"),
)
(str(_core._JSONRPC_METHOD_NOT_FOUND), "method not found", "unknown method", "not found: ping"))
class InvalidMcpUrlError(ValueError):
"""A remote MCP server's ``url`` is not parseable http(s)://. Validated once at startup so
we fail fast instead of burning the reconnect-backoff loop on every attempt."""
"""A remote MCP server's ``url`` is not parseable http(s):// — validated once at startup so we
fail fast instead of burning the reconnect-backoff loop."""
class NonMcpEndpointError(ConnectionError):
"""An HTTP MCP URL served a non-MCP 2xx response (e.g. ``text/html``); real Streamable-HTTP
endpoints answer ``application/json`` or ``text/event-stream``. Non-retryable: every attempt
gets the same page, so the backoff loop is skipped and the server is failed immediately.
Subclasses ConnectionError so broad catches still see a connection problem."""
"""An HTTP MCP URL served a non-MCP 2xx (e.g. ``text/html``). Non-retryable: every attempt gets
the same page, so backoff is skipped and the server fails immediately. Subclasses ConnectionError
so broad catches still see a connection problem."""
def _unwrap_exception_group(exc: BaseException) -> BaseException:
"""Extract the root-cause leaf from anyio ``(Base)ExceptionGroup`` wrappers (group ``str()``
is opaque, so log sites must unwrap). Two rules: a ``KeyboardInterrupt``/``SystemExit`` leaf
anywhere is re-raised (never flattened into a loggable error); a non-cancellation leaf is
preferred over the ``CancelledError`` noise anyio sprays on siblings."""
"""Root-cause leaf of anyio ``(Base)ExceptionGroup`` wrappers (group ``str()`` is opaque). A
``KeyboardInterrupt``/``SystemExit`` leaf anywhere is re-raised, never flattened into a loggable
error; a non-cancellation leaf is preferred over the ``CancelledError`` noise anyio sprays on siblings."""
while isinstance(exc, BaseExceptionGroup) and exc.exceptions:
fatal, _rest = exc.split((KeyboardInterrupt, SystemExit))
if fatal is not None:
leaf: BaseException = fatal
leaf: BaseException = exc.split((KeyboardInterrupt, SystemExit))[0]
if leaf is not None:
while isinstance(leaf, BaseExceptionGroup) and leaf.exceptions:
leaf = leaf.exceptions[0]
raise leaf
@@ -89,29 +81,20 @@ def _contains_only_cancellation(exc: BaseException) -> bool:
def _classify_mcp_failure(exc: BaseException) -> str:
"""``'permanent'`` (deterministic — ``run()`` parks immediately instead of burning the retry
ladder: auth 401/403, NonMcpEndpointError, InvalidMcpUrlError, missing stdio command
FileNotFoundError / ENOENT) or ``'transient'`` (keeps backoff retry)."""
"""``'permanent'`` (``run()`` parks instead of burning the retry ladder: auth 401/403,
NonMcpEndpointError, InvalidMcpUrlError, missing stdio command) or ``'transient'`` (backoff retry)."""
root = _unwrap_exception_group(exc)
permanent = (
_core._is_auth_error(root)
or isinstance(root, (NonMcpEndpointError, InvalidMcpUrlError, FileNotFoundError))
or (isinstance(root, OSError) and getattr(root, "errno", None) == errno.ENOENT)
# 401/403 HTTPStatusError that _is_auth_error's type-gate missed (auth types not importable here).
or _response_status(root) in (401, 403)
)
permanent = (_core._is_auth_error(root)
or isinstance(root, (NonMcpEndpointError, InvalidMcpUrlError, FileNotFoundError))
or (isinstance(root, OSError) and getattr(root, "errno", None) == errno.ENOENT)
# 401/403 HTTPStatusError that _is_auth_error's type-gate missed (auth types not importable here)
or getattr(getattr(root, "response", None), "status_code", None) in (401, 403))
return "permanent" if permanent else "transient"
def _response_status(exc: BaseException):
"""``exc.response.status_code`` for httpx-shaped errors, else None."""
return getattr(getattr(exc, "response", None), "status_code", None)
def _validate_remote_mcp_url(server_name: str, url: Any) -> str:
"""The stripped URL if it is a valid http(s) remote MCP URL. Raises InvalidMcpUrlError naming
the server for non-strings, missing/other schemes (stdio servers use ``command``, not ``url``),
and empty hosts."""
"""The stripped URL if it is a valid http(s) URL; else InvalidMcpUrlError naming the server
(non-string, other scheme — stdio servers use ``command`` — or empty host)."""
def _bad(detail: str) -> InvalidMcpUrlError:
return InvalidMcpUrlError(f"Invalid MCP URL for '{server_name}': {detail}")
@@ -134,10 +117,9 @@ def _validate_remote_mcp_url(server_name: str, url: Any) -> str:
def _resolve_client_cert(server_name: str, config: dict):
"""``client_cert`` / ``client_key`` in httpx's ``cert=`` shape: None when neither is set; a
single path for a combined PEM; ``(cert, key)`` or ``(cert, key, password)`` for the
pair/list forms. ``~`` is expanded and missing files raise a server-scoped
FileNotFoundError instead of an opaque TLS handshake error."""
"""``client_cert`` / ``client_key`` in httpx's ``cert=`` shape: None, a combined-PEM path,
``(cert, key)`` or ``(cert, key, password)``. ``~`` is expanded; missing files raise a
server-scoped FileNotFoundError instead of an opaque TLS handshake error."""
raw_cert = config.get("client_cert")
raw_key = config.get("client_key")
if raw_cert is None and raw_key is None:
@@ -152,29 +134,26 @@ def _resolve_client_cert(server_name: str, config: dict):
raise FileNotFoundError(f"{prefix}{label} not found at {expanded!r}")
return expanded
if isinstance(raw_cert, (list, tuple)):
if raw_key is not None:
raise ValueError(f"{prefix}specify either client_cert as a list [cert, key] OR "
f"client_cert + client_key, not both")
if len(raw_cert) not in (2, 3):
raise ValueError(f"{prefix}client_cert list form must have 2 or 3 elements (got {len(raw_cert)})")
pair = (_expand(raw_cert[0], "client_cert[0]"), _expand(raw_cert[1], "client_cert[1]"))
if len(raw_cert) == 2:
return pair
if not isinstance(raw_cert[2], str):
raise ValueError(f"{prefix}client_cert[2] (key passphrase) must be a string")
return (*pair, raw_cert[2])
cert_path = _expand(raw_cert, "client_cert")
if not isinstance(raw_cert, (list, tuple)):
cert_path = _expand(raw_cert, "client_cert")
return (cert_path, _expand(raw_key, "client_key")) if raw_key is not None else cert_path # combined PEM
if raw_key is not None:
return (cert_path, _expand(raw_key, "client_key"))
return cert_path # single combined PEM (cert + key)
raise ValueError(f"{prefix}specify either client_cert as a list [cert, key] OR "
f"client_cert + client_key, not both")
if len(raw_cert) not in (2, 3):
raise ValueError(f"{prefix}client_cert list form must have 2 or 3 elements (got {len(raw_cert)})")
pair = (_expand(raw_cert[0], "client_cert[0]"), _expand(raw_cert[1], "client_cert[1]"))
if len(raw_cert) == 2:
return pair
if not isinstance(raw_cert[2], str):
raise ValueError(f"{prefix}client_cert[2] (key passphrase) must be a string")
return (*pair, raw_cert[2])
def _resolve_identity_header(server_name: str, config: dict):
"""Optional per-server ``identity_header`` ``{name, value_from: "static"|"profile", value}``
(``value`` required for static) → ``(name, value)`` or None. Invalid configs warn and are
ignored — an identity header must never break the connection. ``profile`` resolves once at
connect time; no per-call mutation."""
"""``identity_header`` ``{name, value_from: "static"|"profile", value}`` → ``(name, value)`` or
None. Invalid configs warn and are ignored — an identity header must never break the connection.
``profile`` resolves once at connect time."""
raw = config.get("identity_header")
if raw is None:
return None
@@ -189,55 +168,48 @@ def _resolve_identity_header(server_name: str, config: dict):
if not isinstance(name, str) or not name.strip():
return _ignore("requires a non-empty 'name'")
value_from = (raw.get("value_from") or "static").strip().lower()
if value_from == "static":
value = raw.get("value")
if not isinstance(value, str) or not value.strip():
return _ignore("with value_from: static requires a non-empty string 'value'")
return (name.strip(), value)
if value_from == "profile":
from hermes_cli.profiles import get_active_profile_name
return (name.strip(), get_active_profile_name())
return _ignore("value_from must be 'static' or 'profile' (got %r)", value_from)
if value_from != "static":
return _ignore("value_from must be 'static' or 'profile' (got %r)", value_from)
value = raw.get("value")
if not isinstance(value, str) or not value.strip():
return _ignore("with value_from: static requires a non-empty string 'value'")
return (name.strip(), value)
def _apply_identity_header(server_name: str, config: dict, headers: dict) -> dict:
"""Merge the resolved identity header into ``headers`` in place. An explicit per-server
``headers`` entry with the same name (any casing) wins — the identity header never silently
overrides user config."""
resolved = _resolve_identity_header(server_name, config)
if resolved is None:
"""Merge the identity header into ``headers`` in place; an explicit entry of the same name (any
casing) wins — never silently override user config."""
name, value = _resolve_identity_header(server_name, config) or (None, None)
if name is None:
return headers
name, value = resolved
if any(key.lower() == name.lower() for key in headers):
logger.debug("MCP server '%s': identity_header '%s' already set via explicit "
"headers config — keeping the explicit value", server_name, name)
return headers
headers[name] = value
else:
headers[name] = value
return headers
def _make_redirect_header_stripper(original_url, *, strict: bool = False,
configured_header_names: "set[str] | frozenset[str]" = frozenset()):
"""httpx response hook guarding cross-origin redirects: always strips ``Authorization`` when
a redirect leaves the original origin; with *strict* (Agent Plugins v1
``strict_redirect_headers``) every configured header (lowercase names in
*configured_header_names*) is stripped too — the v1 spec forbids forwarding
package-configured headers cross-origin."""
"""httpx response hook: strips ``Authorization`` when a redirect leaves the original origin;
with *strict* (Agent Plugins v1 ``strict_redirect_headers``) every configured header (lowercase
names in *configured_header_names*) is stripped too — v1 forbids forwarding them cross-origin."""
origin = (original_url.scheme, original_url.host, original_url.port)
async def _strip_on_cross_origin_redirect(response):
if not (response.is_redirect and response.next_request):
return
target = response.next_request.url
if (target.scheme, target.host, target.port) == origin:
target = response.next_request.url if response.is_redirect and response.next_request else None
if target is None or (target.scheme, target.host, target.port) == origin:
return
headers = response.next_request.headers
headers.pop("authorization", None)
headers.pop("Authorization", None)
if strict:
for _name in configured_header_names:
while _name in headers:
del headers[_name]
for _name in configured_header_names if strict else ():
while _name in headers:
del headers[_name]
return _strip_on_cross_origin_redirect
@@ -245,14 +217,11 @@ def _make_redirect_header_stripper(original_url, *, strict: bool = False,
def _exc_children(exc: BaseException) -> List[BaseException]:
"""Sub-exceptions of a group, else ``__cause__``/``__context__`` when they are exceptions."""
nested = getattr(exc, "exceptions", None)
if nested:
return list(nested)
return [c for c in (exc.__cause__, exc.__context__) if isinstance(c, BaseException)]
return list(nested) if nested else [c for c in (exc.__cause__, exc.__context__) if isinstance(c, BaseException)]
def _format_connect_error(exc: BaseException) -> str:
"""Render nested MCP connection errors into an actionable short message."""
def _find_missing(current: BaseException) -> Optional[str]:
if isinstance(current, FileNotFoundError):
if getattr(current, "filename", None):
@@ -265,24 +234,21 @@ def _format_connect_error(exc: BaseException) -> str:
def _flatten_messages(current: BaseException) -> List[str]:
# A group's own str() is opaque — only its children speak.
text = "" if getattr(current, "exceptions", None) else str(current).strip()
messages = [text] if text else []
for child in _exc_children(current):
messages.extend(_flatten_messages(child))
messages = ([text] if text else []) + [m for child in _exc_children(current) for m in _flatten_messages(child)]
return messages or [current.__class__.__name__]
missing = _find_missing(exc)
if missing:
message = f"missing executable '{missing}'"
if os.path.basename(missing) in {"npx", "npm", "node"}:
message += (" (ensure Node.js is installed and PATH includes its bin directory, "
"or set mcp_servers.<name>.command to an absolute path and include "
"that directory in mcp_servers.<name>.env.PATH)")
return _sanitize_error(message)
deduped = list(dict.fromkeys(_flatten_messages(exc)))
return _sanitize_error("; ".join(deduped[:3]))
if not missing:
return _sanitize_error("; ".join(list(dict.fromkeys(_flatten_messages(exc)))[:3]))
message = f"missing executable '{missing}'"
if os.path.basename(missing) in {"npx", "npm", "node"}:
message += (" (ensure Node.js is installed and PATH includes its bin directory, "
"or set mcp_servers.<name>.command to an absolute path and include "
"that directory in mcp_servers.<name>.env.PATH)")
return _sanitize_error(message)
# Lazily-built caches so this module imports even without the SDK OAuth module.
# Lazily-built caches so this module imports without the SDK OAuth module.
_AUTH_ERROR_TYPES: tuple = ()
_HTTP_STATUS_ERROR_TYPES: Optional[tuple] = None
@@ -297,74 +263,59 @@ def _optional_types(module: str, *names: str) -> list:
def _http_status_error_types() -> tuple:
"""``HTTPStatusError`` classes from both httpx flavours: a 401 may come from the SDK's own
stack (``httpx2`` on mcp >= 2.0) or from Hermes' pinned ``httpx``; the classes are unrelated."""
"""``HTTPStatusError`` from both httpx flavours: a 401 may come from the SDK's own stack
(``httpx2`` on mcp >= 2.0) or Hermes' pinned ``httpx``; the classes are unrelated."""
global _HTTP_STATUS_ERROR_TYPES
if _HTTP_STATUS_ERROR_TYPES is None:
sdk_mod = _core.sdk_httpx()
found: list = [sdk_mod.HTTPStatusError] if sdk_mod is not None else []
found += [cls for cls in _optional_types("httpx", "HTTPStatusError") if cls not in found]
_HTTP_STATUS_ERROR_TYPES = tuple(found)
_HTTP_STATUS_ERROR_TYPES = tuple(dict.fromkeys(
([sdk_mod.HTTPStatusError] if sdk_mod is not None else []) + _optional_types("httpx", "HTTPStatusError")))
return _HTTP_STATUS_ERROR_TYPES
def _get_auth_error_types() -> tuple:
"""Cached exception types indicating MCP OAuth failure: SDK ``OAuthFlowError``/``OAuthTokenError``
(+ legacy ``UnauthorizedError``), our ``OAuthNonInteractiveError``, and ``HTTPStatusError`` from
both httpx flavours — the latter needs the 401 status check in :func:`_is_auth_error`."""
"""Cached MCP OAuth failure types: SDK ``OAuthFlowError``/``OAuthTokenError`` (+ legacy
``UnauthorizedError``), our ``OAuthNonInteractiveError``, and both ``HTTPStatusError`` flavours
(which still need the 401 check in :func:`_is_auth_error`)."""
global _AUTH_ERROR_TYPES
if not _AUTH_ERROR_TYPES:
_AUTH_ERROR_TYPES = tuple(
_optional_types("mcp.client.auth", "OAuthFlowError", "OAuthTokenError")
+ _optional_types("mcp.client.auth", "UnauthorizedError") # older SDKs
+ _optional_types("tools.mcp_oauth", "OAuthNonInteractiveError")
+ list(_http_status_error_types())
)
_AUTH_ERROR_TYPES = (*_optional_types("mcp.client.auth", "OAuthFlowError", "OAuthTokenError"),
*_optional_types("mcp.client.auth", "UnauthorizedError"), # older SDKs
*_optional_types("tools.mcp_oauth", "OAuthNonInteractiveError"),
*_http_status_error_types())
return _AUTH_ERROR_TYPES
def _is_auth_error(exc: BaseException) -> bool:
"""True if ``exc`` indicates an MCP OAuth failure. ``HTTPStatusError`` counts only with
status 401; other HTTP errors fall through to the generic error path."""
types = _get_auth_error_types()
if not types or not isinstance(exc, types):
"""True if ``exc`` indicates an MCP OAuth failure; ``HTTPStatusError`` counts only with status 401."""
if not isinstance(exc, _get_auth_error_types()):
return False
status_error_types = _http_status_error_types()
if status_error_types and isinstance(exc, status_error_types):
return getattr(exc.response, "status_code", None) == 401
return True
return getattr(exc.response, "status_code", None) == 401 if isinstance(exc, _http_status_error_types()) else True
# Lower-cased substrings meaning the server-side transport session expired / was GC'd.
# The OAuth token is still valid — only the transport needs rebuilding.
# Lower-cased substrings meaning the transport session expired / was GC'd (OAuth token still valid).
_SESSION_EXPIRED_MARKERS: tuple = (
"invalid or expired session", "expired session", "session expired", "session not found",
"unknown session", "session terminated", "closedresourceerror", "closed resource",
"transport is closed", "connection closed", "broken pipe", "end of file",
)
"transport is closed", "connection closed", "broken pipe", "end of file")
# Node budget for ``_is_session_expired_error``. The visited set breaks cycles; the budget
# bounds pathological acyclic graphs. Kept well above ``sys.getrecursionlimit()`` so deep
# task-group nesting is still fully scanned.
# Node budget for ``_is_session_expired_error`` (the visited set breaks cycles; this bounds acyclic
# blow-ups). Well above ``sys.getrecursionlimit()`` so deep task-group nesting is fully scanned.
_EXC_TRAVERSAL_MAX_NODES = 10_000
def _is_session_expired_error(exc: BaseException) -> bool:
"""True if ``exc`` looks like an MCP transport session expiry. Streamable-HTTP servers GC
session state (idle TTL, restart, pod rotation) while the OAuth token stays valid, so unlike
:func:`_is_auth_error` the fix is a transport reconnect (``_reconnect_event``), not an OAuth refresh."""
# AnyIO stream exceptions are often message-less (``str(ClosedResourceError()) == ""``),
# so type checks are needed in addition to marker matching.
"""True if ``exc`` looks like a transport session expiry (Streamable-HTTP servers GC session
state on idle TTL / restart / pod rotation while the OAuth token stays valid) — the fix is a
transport reconnect, not an OAuth refresh. Iterative walk over ``exceptions`` / ``__cause__`` /
``__context__`` with a visited set AND a node budget; every reachable node is inspected so an
InterruptedError anywhere overrides transport markers, and the chain walk matters because SDK
wrappers raise a generic RuntimeError *from* a message-less ClosedResourceError."""
# AnyIO stream exceptions are often message-less, so type checks complement marker matching.
transport_error_types = tuple(_optional_types("anyio", "BrokenResourceError", "ClosedResourceError", "EndOfStream"))
# Iterative traversal over ``exceptions`` / ``__cause__`` / ``__context__`` with an
# identity-visited set AND a node budget (graphs can be deep or cyclic). Every reachable
# node is inspected so an InterruptedError anywhere overrides transport markers; the chain
# walk matters because SDK wrappers often raise a generic RuntimeError *from* the
# message-less ClosedResourceError.
stack: "list[BaseException | None]" = [exc]
seen: set[int] = set()
transport_error_found = False
found = False
budget = _EXC_TRAVERSAL_MAX_NODES
while stack and budget > 0:
current = stack.pop()
@@ -374,11 +325,10 @@ def _is_session_expired_error(exc: BaseException) -> bool:
budget -= 1
if isinstance(current, InterruptedError):
return False
# Messages vary across SDK versions and servers: match a narrow allow-list of stable
# substrings, not exception type, to avoid false positives.
# Messages vary across SDK versions/servers: a narrow allow-list of stable substrings avoids
# false positives.
msg = str(current).lower()
if isinstance(current, transport_error_types) or (msg and any(marker in msg for marker in _SESSION_EXPIRED_MARKERS)):
transport_error_found = True
found = found or isinstance(current, transport_error_types) or any(m in msg for m in _SESSION_EXPIRED_MARKERS)
stack.extend(getattr(current, "exceptions", ()))
stack.extend((getattr(current, "__cause__", None), getattr(current, "__context__", None)))
return transport_error_found
return found

View File

@@ -1,13 +1,13 @@
"""Background-loop plumbing for tools.mcp_tool: the cross-process discovery file lock,
scheduling coroutines onto the MCP loop from caller threads (with profile HOME override
and dashboard OAuth flow propagation) and the loop's exception handler. Split from
tools/mcp_tool.py; origin state (``_lock``, ``_mcp_loop``) is read through ``_core`` so
``mock.patch("tools.mcp_tool.X")`` keeps working."""
"""Background-loop plumbing for tools.mcp_tool: cross-process discovery file lock, scheduling
coroutines onto the MCP loop from caller threads (with profile HOME override and dashboard OAuth
flow propagation) and the loop's exception handler. Origin state (``_lock``, ``_mcp_loop``) is
read through ``_core`` so ``mock.patch("tools.mcp_tool.X")`` keeps working."""
from __future__ import annotations
import asyncio
import concurrent.futures
import contextlib
import errno
import logging
import os
@@ -20,11 +20,8 @@ logger = logging.getLogger("tools.mcp_tool")
class _LockCookie:
"""Holds a cross-process file lock; ``release()`` drops it.
The file object MUST stay open while the lock is held: both the fcntl and
the portalocker lock are tied to the descriptor's lifetime.
"""
"""Holds a cross-process file lock; ``release()`` drops it. The file object MUST stay open while
held: both fcntl and portalocker locks are tied to the descriptor."""
def __init__(self, fh: Any) -> None:
self._fh = fh
@@ -32,76 +29,56 @@ class _LockCookie:
def release(self) -> None:
if self._fh is None:
return
# Best effort on every step: an unlock/close failure must never
# propagate out of discovery.
try:
# Best effort: an unlock/close failure must never propagate out of discovery.
with contextlib.suppress(Exception):
if os.name == "posix":
import fcntl
fcntl.flock(self._fh.fileno(), fcntl.LOCK_UN)
else:
import portalocker
portalocker.unlock(self._fh)
except Exception:
pass
try:
with contextlib.suppress(Exception):
self._fh.close()
except Exception:
pass
self._fh = None
def _acquire_lock_on_fh(fh: Any) -> bool:
"""Non-blocking exclusive lock (fcntl on POSIX, portalocker elsewhere).
False when another process holds it; unexpected errors propagate so the
caller can treat locking as unavailable.
"""
fd = fh.fileno()
"""Non-blocking exclusive lock (fcntl on POSIX, portalocker elsewhere). False when another process
holds it; unexpected errors propagate so the caller can treat locking as unavailable."""
if os.name == "posix":
import fcntl
try:
fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
return True
fcntl.flock(fh.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)
except OSError as e:
if e.errno in (errno.EACCES, errno.EAGAIN, errno.EWOULDBLOCK):
return False
raise
else:
import portalocker
try:
portalocker.lock(fh, portalocker.LOCK_EX | portalocker.LOCK_NB)
return True
except portalocker.LockException:
return False
return True
import portalocker
try:
portalocker.lock(fh, portalocker.LOCK_EX | portalocker.LOCK_NB)
except portalocker.LockException:
return False
return True
def _try_acquire_mcp_discovery_lock() -> Any:
"""Return a ``_LockCookie`` (acquired), ``None`` (held by another process)
or ``_LOCK_UNAVAILABLE`` (locking broken: run discovery unguarded)."""
# The cached path lives on the ORIGIN module (tests reset
# ``tools.mcp_tool._MCP_DISCOVERY_LOCK_PATH = None``), so write it there.
"""``_LockCookie`` (acquired), ``None`` (held by another process) or ``_LOCK_UNAVAILABLE``
(locking broken: run discovery unguarded)."""
# The cached path lives on the ORIGIN module (tests reset ``tools.mcp_tool._MCP_DISCOVERY_LOCK_PATH``).
from tools import mcp_tool as _origin
try:
from hermes_constants import get_hermes_home
if _origin._MCP_DISCOVERY_LOCK_PATH is None:
_origin._MCP_DISCOVERY_LOCK_PATH = str(
get_hermes_home() / ".mcp-discovery.lock"
)
lock_path = _origin._MCP_DISCOVERY_LOCK_PATH
_origin._MCP_DISCOVERY_LOCK_PATH = str(get_hermes_home() / ".mcp-discovery.lock")
fh = open(_origin._MCP_DISCOVERY_LOCK_PATH, "w", encoding="utf-8")
except Exception:
return _core._LOCK_UNAVAILABLE
try:
fh = open(lock_path, "w", encoding="utf-8")
except Exception:
return _core._LOCK_UNAVAILABLE
try:
acquired = _core._acquire_lock_on_fh(fh)
except Exception:
fh.close()
return _core._LOCK_UNAVAILABLE
if acquired:
return _core._LockCookie(fh)
fh.close()
@@ -109,27 +86,21 @@ def _try_acquire_mcp_discovery_lock() -> Any:
def _mcp_loop_exception_handler(loop, context):
"""Suppress the benign 'Event loop is closed' RuntimeError that httpx
finalizers raise against the dead loop during shutdown; forward the rest."""
"""Suppress the benign 'Event loop is closed' RuntimeError httpx finalizers raise against the
dead loop during shutdown; forward the rest."""
exc = context.get("exception")
if isinstance(exc, RuntimeError) and "Event loop is closed" in str(exc):
return
loop.default_exception_handler(context)
if not (isinstance(exc, RuntimeError) and "Event loop is closed" in str(exc)):
loop.default_exception_handler(context)
def _wrap_with_home_override(coro: "Coroutine") -> "Coroutine":
"""Carry the caller's context-local HERMES_HOME override into ``coro``
(task-local on the MCP loop, so concurrent scopes don't interfere)."""
"""Carry the caller's context-local HERMES_HOME override into ``coro`` (task-local on the MCP
loop, so concurrent scopes don't interfere)."""
try:
from hermes_constants import (
get_hermes_home_override,
reset_hermes_home_override,
set_hermes_home_override,
)
from hermes_constants import get_hermes_home_override, reset_hermes_home_override, set_hermes_home_override
home_override = get_hermes_home_override()
except Exception:
return coro
home_override = None
if not home_override:
return coro
@@ -146,14 +117,10 @@ def _wrap_with_home_override(coro: "Coroutine") -> "Coroutine":
def _wrap_with_dashboard_oauth_flow(coro):
"""Propagate a dashboard OAuth flow onto the dedicated MCP loop task."""
try:
from tools.mcp_dashboard_oauth import (
dashboard_oauth_flow,
get_dashboard_oauth_flow,
)
from tools.mcp_dashboard_oauth import dashboard_oauth_flow, get_dashboard_oauth_flow
flow = get_dashboard_oauth_flow()
except Exception:
return coro
flow = None
if flow is None:
return coro
@@ -172,12 +139,8 @@ def _running_loop() -> Optional[asyncio.AbstractEventLoop]:
def _run_on_mcp_loop(coro_or_factory, timeout: float = 30):
"""Schedule a coroutine on the MCP loop and block until done.
Accepts a coroutine or a zero-arg factory (a factory avoids leaking a
never-awaited coroutine when the loop is down). Polls in short intervals
so the calling thread can honor user interrupts.
"""
"""Schedule a coroutine (or zero-arg factory — avoids leaking a never-awaited coroutine when the
loop is down) on the MCP loop and block until done, polling so user interrupts are honored."""
from tools.interrupt import is_interrupted
from agent.async_utils import safe_schedule_threadsafe
@@ -186,58 +149,37 @@ def _run_on_mcp_loop(coro_or_factory, timeout: float = 30):
if asyncio.iscoroutine(coro_or_factory):
coro_or_factory.close()
raise RuntimeError("MCP event loop is not running")
coro = coro_or_factory() if callable(coro_or_factory) else coro_or_factory
# Tasks created via run_coroutine_threadsafe copy the LOOP thread's
# context, so a per-request profile scope would vanish here; re-establish
# it inside the task's own context.
coro = _core._wrap_with_home_override(coro)
coro = _core._wrap_with_dashboard_oauth_flow(coro)
future = safe_schedule_threadsafe(
coro, loop,
logger=logger,
log_message="MCP scheduling failed",
)
# run_coroutine_threadsafe copies the LOOP thread's context, so a per-request profile scope
# would vanish here; re-establish it inside the task's own context.
coro = _core._wrap_with_dashboard_oauth_flow(_core._wrap_with_home_override(
coro_or_factory() if callable(coro_or_factory) else coro_or_factory))
future = safe_schedule_threadsafe(coro, loop, logger=logger, log_message="MCP scheduling failed")
if future is None:
raise RuntimeError("MCP event loop unavailable (failed to schedule)")
start_time = time.monotonic()
deadline = None if timeout is None else start_time + timeout
while True:
if is_interrupted():
future.cancel()
raise InterruptedError("User sent a new message")
wait_timeout = 0.1
if deadline is not None:
remaining = deadline - time.monotonic()
if remaining <= 0:
future.cancel()
elapsed = time.monotonic() - start_time
raise TimeoutError(
f"MCP call timed out after {elapsed:.1f}s "
f"(configured timeout: {float(timeout):.1f}s)"
)
wait_timeout = min(wait_timeout, remaining)
remaining = 0.1 if deadline is None else deadline - time.monotonic()
if remaining <= 0:
future.cancel()
raise TimeoutError(f"MCP call timed out after {time.monotonic() - start_time:.1f}s "
f"(configured timeout: {float(timeout):.1f}s)")
try:
return future.result(timeout=wait_timeout)
return future.result(timeout=min(0.1, remaining))
except concurrent.futures.TimeoutError:
# Aliases builtin TimeoutError, so this also fires for the
# coroutine's own timeout: a done future must yield its outcome.
# Aliases builtin TimeoutError, so it also fires for the coroutine's own timeout: a done
# future must yield its outcome.
if future.done():
return future.result()
continue
def _signal_reconnect(server: Any) -> bool:
"""Ask a server task to rebuild its transport, thread-safely.
Handlers run on caller threads while the event lives on the MCP loop, so
it is set via ``call_soon_threadsafe`` when the loop runs (direct
``.set()`` otherwise). False when the server has no reconnect machinery.
"""
"""Ask a server task to rebuild its transport, thread-safely: the event lives on the MCP loop,
so set via ``call_soon_threadsafe`` when it runs (direct ``.set()`` otherwise). False when the
server has no reconnect machinery."""
event = getattr(server, "_reconnect_event", None)
if event is None:
return False
@@ -253,97 +195,58 @@ def reconnect_mcp_server(server_name: str) -> bool:
"""Ask a currently-live MCP server to rebuild after external re-auth."""
with _core._lock:
server = _core._servers.get(server_name)
if server is None:
return False
return _core._signal_reconnect(server)
return server is not None and _core._signal_reconnect(server)
def _wait_for_server_session_ready(
srv: Any,
*,
old_session: Any = None,
timeout: float = 15.0,
) -> bool:
"""Poll until the server exposes a usable, ready session.
During a reconnect ``srv.session`` is briefly None or still the stale
object; retrying blindly there burns breaker strikes. With
``old_session`` the observed session must differ from it. Iteration-
bounded, not deadline-bounded: tests freeze ``time.monotonic``.
"""
poll_interval = 0.25
iterations = max(1, int(max(float(timeout), 0.0) / poll_interval))
def _wait_for_server_session_ready(srv: Any, *, old_session: Any = None, timeout: float = 15.0) -> bool:
"""Poll until the server exposes a usable, ready session (during a reconnect ``srv.session`` is
briefly None or stale; retrying blindly burns breaker strikes). With ``old_session`` the observed
session must differ. Iteration-bounded, not deadline-bounded: tests freeze ``time.monotonic``."""
iterations = max(1, int(max(float(timeout), 0.0) / 0.25))
for i in range(iterations):
session = getattr(srv, "session", None)
ready = getattr(srv, "_ready", None)
is_ready = True
if ready is not None and hasattr(ready, "is_set"):
try:
is_ready = bool(ready.is_set())
except Exception:
is_ready = True
try:
is_ready = bool(ready.is_set()) if hasattr(ready, "is_set") else True
except Exception:
is_ready = True
if session is not None and session is not old_session and is_ready:
return True
if i < iterations - 1:
time.sleep(poll_interval)
time.sleep(0.25)
return False
def _signal_reconnect_and_wait(
server_name: str,
srv: Any,
*,
op_description: str,
timeout: float = 15.0,
) -> bool:
"""Request a transport rebuild and wait for the fresh session.
``_ready`` is cleared on the loop BEFORE ``_reconnect_event`` is set;
otherwise the readiness poll returns immediately and retries against the
same dead session.
"""
def _signal_reconnect_and_wait(server_name: str, srv: Any, *, op_description: str, timeout: float = 15.0) -> bool:
"""Request a transport rebuild and wait for the fresh session. ``_ready`` is cleared on the loop
BEFORE ``_reconnect_event`` is set, else the readiness poll returns at once on the dead session."""
loop = _core._mcp_loop
if loop is None or not loop.is_running():
return False
old_session = getattr(srv, "session", None)
def _request_reconnect() -> None:
ready = getattr(srv, "_ready", None)
if ready is not None and hasattr(ready, "clear"):
ready, reconnect_event = getattr(srv, "_ready", None), getattr(srv, "_reconnect_event", None)
if hasattr(ready, "clear"):
ready.clear()
reconnect_event = getattr(srv, "_reconnect_event", None)
if reconnect_event is not None and hasattr(reconnect_event, "set"):
if hasattr(reconnect_event, "set"):
reconnect_event.set()
logger.info(
"MCP server '%s': %s requesting transport reconnect",
server_name, op_description,
)
old_session = getattr(srv, "session", None)
logger.info("MCP server '%s': %s requesting transport reconnect", server_name, op_description)
loop.call_soon_threadsafe(_request_reconnect)
return _core._wait_for_server_session_ready(
srv,
old_session=old_session,
timeout=timeout,
)
return _core._wait_for_server_session_ready(srv, old_session=old_session, timeout=timeout)
def _ensure_mcp_loop():
"""Start the background event loop thread if not already running.
The loop/thread handles live on the ORIGIN module (tests read and reset
``tools.mcp_tool._mcp_loop``), so they are written there, never here.
"""
"""Start the background loop thread if not running. The loop/thread handles live on the ORIGIN
module (tests read and reset ``tools.mcp_tool._mcp_loop``), so they are written there."""
from tools import mcp_tool as _origin
with _core._lock:
if _origin._mcp_loop is not None and _origin._mcp_loop.is_running():
return
_origin._mcp_loop = asyncio.new_event_loop()
_origin._mcp_loop.set_exception_handler(_core._mcp_loop_exception_handler)
_origin._mcp_thread = threading.Thread(
target=_origin._mcp_loop.run_forever,
name="mcp-event-loop",
daemon=True,
)
loop = _origin._mcp_loop = asyncio.new_event_loop()
loop.set_exception_handler(_core._mcp_loop_exception_handler)
_origin._mcp_thread = threading.Thread(target=loop.run_forever, name="mcp-event-loop", daemon=True)
_origin._mcp_thread.start()
@@ -354,53 +257,41 @@ def _stop_mcp_loop(*, only_if_idle: bool = False) -> bool:
if only_if_idle and (_core._servers or _core._server_connecting):
logger.debug("Leaving MCP event loop running; active servers are registered or connecting")
return False
loop = _origin._mcp_loop
thread = _origin._mcp_thread
_origin._mcp_loop = None
_origin._mcp_thread = None
if loop is not None:
# Drain before stopping: tasks still suspended when the loop closes
# get resumed by the GC against a closed loop. shutdown_mcp_servers
# only reaps servers held in _servers; everything else ends up here.
stop_owned_by_loop = False
if loop.is_running():
from agent.async_utils import safe_schedule_threadsafe
loop, thread = _origin._mcp_loop, _origin._mcp_thread
_origin._mcp_loop = _origin._mcp_thread = None
if loop is None:
return True
# Drain before stopping: tasks still suspended when the loop closes get resumed by the GC
# against a closed loop. shutdown_mcp_servers only reaps _servers; everything else ends here.
future = None
if loop.is_running():
from agent.async_utils import safe_schedule_threadsafe
future = safe_schedule_threadsafe(
_core._drain_and_stop_mcp_loop(), loop,
logger=logger,
log_message="MCP loop drain: failed to schedule",
log_level=logging.WARNING,
)
if future is not None:
stop_owned_by_loop = True
try:
future.result(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT + 1)
except TimeoutError:
logger.warning(
"Timed out waiting for MCP loop drain after %.1fs",
_core._MCP_LOOP_DRAIN_TIMEOUT + 1,
)
except BaseException as exc:
logger.warning("Error draining MCP loop tasks: %s", exc)
elif not loop.is_closed():
future = safe_schedule_threadsafe(
_core._drain_and_stop_mcp_loop(), loop, logger=logger,
log_message="MCP loop drain: failed to schedule", log_level=logging.WARNING)
if future is not None:
try:
loop.run_until_complete(
_core._drain_mcp_loop_tasks(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT)
)
future.result(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT + 1)
except TimeoutError:
logger.warning("Timed out waiting for MCP loop drain after %.1fs", _core._MCP_LOOP_DRAIN_TIMEOUT + 1)
except BaseException as exc:
logger.warning("Error draining stopped MCP loop tasks: %s", exc)
if not stop_owned_by_loop and loop.is_running():
loop.call_soon_threadsafe(loop.stop)
if thread is not None:
thread.join(timeout=5)
if thread.is_alive():
logger.warning("MCP event loop thread did not stop within 5.0s")
logger.warning("Error draining MCP loop tasks: %s", exc)
elif not loop.is_closed():
try:
loop.close()
except Exception as exc:
logger.warning("Unable to close MCP event loop cleanly: %s", exc)
# The loop is gone, so no session can be in flight: reap active too.
_core._kill_orphaned_mcp_children(include_active=True)
loop.run_until_complete(_core._drain_mcp_loop_tasks(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT))
except BaseException as exc:
logger.warning("Error draining stopped MCP loop tasks: %s", exc)
if future is None and loop.is_running(): # drain-and-stop wasn't scheduled: stop it ourselves
loop.call_soon_threadsafe(loop.stop)
if thread is not None:
thread.join(timeout=5)
if thread.is_alive():
logger.warning("MCP event loop thread did not stop within 5.0s")
try:
loop.close()
except Exception as exc:
logger.warning("Unable to close MCP event loop cleanly: %s", exc)
# The loop is gone, so no session can be in flight: reap active too.
_core._kill_orphaned_mcp_children(include_active=True)
return True

View File

@@ -13,16 +13,11 @@ logger = logging.getLogger("tools.mcp_tool")
def _tool_use_id(block):
"""Tool-use id (the discriminator for a tool *result* block), read under both
SDK spellings — on mcp 2.x a bare ``hasattr(b, "toolUseId")`` is False and
would silently drop tool results."""
"""Tool-use id (marks a tool *result* block) under both SDK spellings — on mcp 2.x a bare
``hasattr(b, "toolUseId")`` is False and would silently drop tool results."""
return mcp_field(block, "tool_use_id", "toolUseId", _MISSING)
def _is_tool_use(block) -> bool:
return hasattr(block, "name") and hasattr(block, "input")
def _tool_result_text(block) -> str:
"""Text of a ToolResultContent block ("" when it carries no content)."""
content = getattr(block, "content", None)
@@ -45,26 +40,19 @@ def _content_part(block) -> Optional[dict]:
def _tool_call_dict(tu, index: int) -> dict:
args = tu.input
return {
"id": getattr(tu, "id", f"call_{index}"),
"type": "function",
"function": {
"name": tu.name,
"arguments": json.dumps(args, ensure_ascii=False) if isinstance(args, dict) else str(args),
},
}
return {"id": getattr(tu, "id", f"call_{index}"), "type": "function", "function": {
"name": tu.name, "arguments": json.dumps(args, ensure_ascii=False) if isinstance(args, dict) else str(args)}}
def _convert_sampling_message(msg) -> List[dict]:
"""One MCP SamplingMessage -> OpenAI-format messages (tool results first,
then either an assistant tool_calls message or plain content)."""
"""One MCP SamplingMessage -> OpenAI messages: tool results first, then either an assistant
tool_calls message or plain content."""
blocks = msg.content_as_list if hasattr(msg, "content_as_list") else (
msg.content if isinstance(msg.content, list) else [msg.content]
)
msg.content if isinstance(msg.content, list) else [msg.content])
tool_results = [b for b in blocks if _tool_use_id(b) is not _MISSING]
tool_uses = [b for b in blocks if _is_tool_use(b) and _tool_use_id(b) is _MISSING]
content_blocks = [b for b in blocks if _tool_use_id(b) is _MISSING and not _is_tool_use(b)]
others = [b for b in blocks if _tool_use_id(b) is _MISSING]
tool_uses = [b for b in others if hasattr(b, "name") and hasattr(b, "input")]
content_blocks = [b for b in others if not (hasattr(b, "name") and hasattr(b, "input"))]
out = [{"role": "tool", "tool_call_id": _tool_use_id(tr), "content": _tool_result_text(tr)} for tr in tool_results]
if tool_uses:
msg_dict: dict = {"role": msg.role, "tool_calls": [_tool_call_dict(tu, i) for i, tu in enumerate(tool_uses)]}
@@ -72,45 +60,32 @@ def _convert_sampling_message(msg) -> List[dict]:
if text_parts:
msg_dict["content"] = "\n".join(text_parts)
out.append(msg_dict)
elif len(content_blocks) == 1 and hasattr(content_blocks[0], "text"):
out.append({"role": msg.role, "content": content_blocks[0].text})
elif content_blocks:
if len(content_blocks) == 1 and hasattr(content_blocks[0], "text"):
out.append({"role": msg.role, "content": content_blocks[0].text})
else:
parts = [p for p in map(_content_part, content_blocks) if p is not None]
if parts:
out.append({"role": msg.role, "content": parts})
parts = [p for p in map(_content_part, content_blocks) if p is not None]
if parts:
out.append({"role": msg.role, "content": parts})
return out
def _parse_tool_call_arguments(server_name: str, args) -> dict:
"""LLM tool_calls arguments -> dict; malformed JSON / non-dict values are
wrapped as ``{"_raw": ...}`` rather than dropped."""
"""LLM tool_calls arguments -> dict; malformed JSON / non-dicts become ``{"_raw": ...}``, not dropped."""
if isinstance(args, str):
try:
return json.loads(args)
except (json.JSONDecodeError, ValueError):
logger.warning(
"MCP server '%s': malformed tool_calls arguments from LLM (wrapping as raw): %.100s",
server_name, args,
)
logger.warning("MCP server '%s': malformed tool_calls arguments from LLM (wrapping as raw): %.100s",
server_name, args)
return {"_raw": args}
return args if isinstance(args, dict) else {"_raw": str(args)}
def _response_total_tokens(response, default):
return getattr(getattr(response, "usage", None), "total_tokens", default)
class SamplingHandler:
"""Handles sampling/createMessage requests for one MCP server; passed to
``ClientSession`` as ``sampling_callback``. All state (rate-limit
timestamps, metrics, tool-loop counter) is per instance. Runs on the MCP
background loop; the sync LLM call is offloaded via ``asyncio.to_thread``.
Deprecated upstream (MCP 2026-07-28, SEP-2577, 12-month window): stays fully
functional because handshake-era servers still issue it, but do NOT grow new
capability here — modern servers use MRTR, handled by the SDK session layer.
"""
"""``sampling_callback`` for one MCP server (per-instance rate-limit, metrics, tool-loop state).
Runs on the MCP loop; the sync LLM call is offloaded via ``asyncio.to_thread``. Deprecated
upstream (MCP 2026-07-28, SEP-2577): stays functional for handshake-era servers, but do NOT grow
new capability here — modern servers use MRTR, handled by the SDK session layer."""
_STOP_REASON_MAP = {"stop": "endTurn", "length": "maxTokens", "tool_calls": "toolUse"}
_LOG_LEVELS = {"debug": logging.DEBUG, "info": logging.INFO, "warning": logging.WARNING}
@@ -141,22 +116,19 @@ class SamplingHandler:
"""Config override > server hint > None (use default)."""
if self.model_override:
return self.model_override
for hint in (getattr(preferences, "hints", None) or []):
if getattr(hint, "name", None):
return hint.name
return None
hints = getattr(preferences, "hints", None) or []
return next((hint.name for hint in hints if getattr(hint, "name", None)), None)
def _convert_messages(self, params) -> List[dict]:
"""Convert MCP SamplingMessages to OpenAI format (``content_as_list``
when the SDK provides it; per-block duck-typed dispatch)."""
"""MCP SamplingMessages -> OpenAI format (per-block duck-typed dispatch)."""
return [m for msg in params.messages for m in _convert_sampling_message(msg)]
@staticmethod
def _error(message: str, code: int = -1):
"""Return ErrorData (MCP spec) or raise as fallback."""
if _core._MCP_SAMPLING_TYPES:
return _core.ErrorData(code=code, message=message)
raise Exception(message)
if not _core._MCP_SAMPLING_TYPES:
raise Exception(message)
return _core.ErrorData(code=code, message=message)
def _fail(self, message: str):
"""Count an error and return the ErrorData for it."""
@@ -164,77 +136,55 @@ class SamplingHandler:
return self._error(message)
def _log_response(self, response, suffix: str = "", *args) -> None:
logger.log(
self.audit_level, "MCP server '%s' sampling response: model=%s, tokens=%s" + suffix,
self.server_name, response.model, _response_total_tokens(response, "?"), *args,
)
logger.log(self.audit_level, "MCP server '%s' sampling response: model=%s, tokens=%s" + suffix,
self.server_name, response.model, getattr(getattr(response, "usage", None), "total_tokens", "?"), *args)
def _build_tool_use_result(self, choice, response):
"""Build a CreateMessageResultWithTools from an LLM tool_calls response,
subject to tool-loop governance (``max_tool_rounds``; 0 disables)."""
"""CreateMessageResultWithTools from a tool_calls response, under ``max_tool_rounds`` (0 disables)."""
self.metrics["tool_use_count"] += 1
if self.max_tool_rounds == 0:
self._tool_loop_count = 0
return self._error(f"Tool loops disabled for server '{self.server_name}' (max_tool_rounds=0)")
self._tool_loop_count += 1
if self._tool_loop_count > self.max_tool_rounds:
if self.max_tool_rounds == 0 or self._tool_loop_count > self.max_tool_rounds:
self._tool_loop_count = 0
return self._error(
f"Tool loop limit exceeded for server '{self.server_name}' (max {self.max_tool_rounds} rounds)"
)
content_blocks = [
_core.ToolUseContent(
type="tool_use", id=tc.id, name=tc.function.name,
input=_parse_tool_call_arguments(self.server_name, tc.function.arguments),
)
for tc in choice.message.tool_calls
]
f"Tool loops disabled for server '{self.server_name}' (max_tool_rounds=0)" if self.max_tool_rounds == 0
else f"Tool loop limit exceeded for server '{self.server_name}' (max {self.max_tool_rounds} rounds)")
content_blocks = [_core.ToolUseContent(type="tool_use", id=tc.id, name=tc.function.name,
input=_parse_tool_call_arguments(self.server_name, tc.function.arguments))
for tc in choice.message.tool_calls]
self._log_response(response, ", tool_calls=%d", len(content_blocks))
return _core.CreateMessageResultWithTools(
role="assistant", content=content_blocks, model=response.model, stopReason="toolUse",
)
role="assistant", content=content_blocks, model=response.model, stopReason="toolUse")
def _build_text_result(self, choice, response):
"""Build a CreateMessageResult from a normal text response (resets the tool loop)."""
"""CreateMessageResult from a normal text response (resets the tool loop)."""
self._tool_loop_count = 0
self._log_response(response)
return _core.CreateMessageResult(
role="assistant",
role="assistant", model=response.model,
content=_core.TextContent(type="text", text=_sanitize_error(choice.message.content or "")),
model=response.model,
stopReason=self._STOP_REASON_MAP.get(choice.finish_reason, "endTurn"),
)
stopReason=self._STOP_REASON_MAP.get(choice.finish_reason, "endTurn"))
def session_kwargs(self) -> dict:
"""Kwargs to pass to ClientSession for sampling support."""
return {
"sampling_callback": self,
"sampling_capabilities": _core.SamplingCapability(tools=_core.SamplingToolsCapability()),
}
return {"sampling_callback": self,
"sampling_capabilities": _core.SamplingCapability(tools=_core.SamplingToolsCapability())}
def _admit(self, params):
"""Rate-limit + allowed_models gate. Returns ``(resolved_model, None)``
or ``(None, ErrorData)``."""
"""Rate-limit + allowed_models gate. Returns ``(resolved_model, None)`` or ``(None, ErrorData)``."""
if not self._check_rate_limit():
logger.warning("MCP server '%s' sampling rate limit exceeded (%d/min)", self.server_name, self.max_rpm)
return None, self._fail(
f"Sampling rate limit exceeded for server '{self.server_name}' ({self.max_rpm} requests/minute)"
)
model = self._resolve_model(mcp_field(params, "model_preferences", "modelPreferences"))
resolved_model = model or self.model_override or ""
f"Sampling rate limit exceeded for server '{self.server_name}' ({self.max_rpm} requests/minute)")
resolved_model = self._resolve_model(mcp_field(params, "model_preferences", "modelPreferences")) or ""
if self.allowed_models and resolved_model and resolved_model not in self.allowed_models:
logger.warning(
"MCP server '%s' requested model '%s' not in allowed_models", self.server_name, resolved_model,
)
return None, self._fail(
f"Model '{resolved_model}' not allowed for server "
f"'{self.server_name}'. Allowed: {', '.join(self.allowed_models)}"
)
logger.warning("MCP server '%s' requested model '%s' not in allowed_models",
self.server_name, resolved_model)
return None, self._fail(f"Model '{resolved_model}' not allowed for server "
f"'{self.server_name}'. Allowed: {', '.join(self.allowed_models)}")
return resolved_model, None
def _build_llm_call(self, params, resolved_model: str) -> Callable[[], object]:
"""Translate the sampling params into a zero-arg sync ``call_llm`` thunk
(run off-loop so the MCP loop is not blocked)."""
"""Sampling params -> zero-arg sync ``call_llm`` thunk (run off-loop); server tools are forwarded."""
from agent.auxiliary_client import call_llm
messages = self._convert_messages(params)
@@ -242,96 +192,69 @@ class SamplingHandler:
if system_prompt:
messages.insert(0, {"role": "system", "content": system_prompt})
max_tokens = min(mcp_field(params, "max_tokens", "maxTokens", self.max_tokens_cap), self.max_tokens_cap)
temperature = getattr(params, "temperature", None)
# Forward server-provided tools.
server_tools = getattr(params, "tools", None)
tools = [
{
"type": "function",
"function": {
"name": getattr(t, "name", ""),
"description": getattr(t, "description", "") or "",
"parameters": _normalize_mcp_input_schema(mcp_field(t, "input_schema", "inputSchema")),
},
}
for t in server_tools
] if server_tools else None
logger.log(
self.audit_level,
"MCP server '%s' sampling request: model=%s, max_tokens=%d, messages=%d",
self.server_name, resolved_model, max_tokens, len(messages),
)
return lambda: call_llm(
task="mcp", model=resolved_model or None, messages=messages, temperature=temperature,
max_tokens=max_tokens, tools=tools, timeout=self.timeout,
)
tools = [{"type": "function", "function": {
"name": getattr(t, "name", ""), "description": getattr(t, "description", "") or "",
"parameters": _normalize_mcp_input_schema(mcp_field(t, "input_schema", "inputSchema"))}}
for t in server_tools] if server_tools else None
logger.log(self.audit_level, "MCP server '%s' sampling request: model=%s, max_tokens=%d, messages=%d",
self.server_name, resolved_model, max_tokens, len(messages))
return lambda: call_llm(task="mcp", model=resolved_model or None, messages=messages, max_tokens=max_tokens,
temperature=getattr(params, "temperature", None), tools=tools, timeout=self.timeout)
async def __call__(self, context, params):
"""SDK sampling callback (``SamplingFnT``). Returns CreateMessageResult,
CreateMessageResultWithTools, or ErrorData."""
"""SDK ``SamplingFnT``: CreateMessageResult, CreateMessageResultWithTools, or ErrorData."""
resolved_model, err = self._admit(params)
if err is not None:
return err
sync_call = self._build_llm_call(params, resolved_model)
sync_call = self._build_llm_call(params, resolved_model) # outside the try: its errors propagate, not _fail
try:
response = await asyncio.wait_for(asyncio.to_thread(sync_call), timeout=self.timeout)
except asyncio.TimeoutError:
return self._fail(f"Sampling LLM call timed out after {self.timeout}s for server '{self.server_name}'")
except Exception as exc:
return self._fail(f"Sampling LLM call failed: {_sanitize_error(_exc_str(exc))}")
# Empty choices happen on content filtering / provider errors.
if not getattr(response, "choices", None):
return self._fail(f"LLM returned empty response (no choices) for server '{self.server_name}'")
choice = response.choices[0]
self.metrics["requests"] += 1
total_tokens = _response_total_tokens(response, 0)
if isinstance(total_tokens, int):
self.metrics["tokens_used"] += total_tokens
total_tokens = getattr(getattr(response, "usage", None), "total_tokens", 0)
self.metrics["tokens_used"] += total_tokens if isinstance(total_tokens, int) else 0
if choice.finish_reason == "tool_calls" and getattr(choice.message, "tool_calls", None):
return self._build_tool_use_result(choice, response)
return self._build_text_result(choice, response)
def _format_elicitation_schema_summary(schema: dict, server_name: str) -> str:
"""Render a flat-object requested_schema as a human-readable field list
(names, types, descriptions) so the user knows what they're approving."""
"""Flat-object requested_schema -> readable field list so the user knows what they're approving."""
props = schema.get("properties") if isinstance(schema, dict) else None
if not isinstance(props, dict) or not props:
return f"Approval requested by MCP server '{server_name}'."
lines = [f"Fields requested by MCP server '{server_name}':"]
for field_name, field_spec in props.items():
spec = field_spec if isinstance(field_spec, dict) else {}
field_type = str(spec.get("type", "") or "")
field_desc = str(spec.get("description", "") or "")
suffix = f" ({field_type})" if field_type else ""
lines.append(f" - {field_name}{suffix}: {field_desc}" if field_desc else f" - {field_name}{suffix}")
field_type, field_desc = str(spec.get("type", "") or ""), str(spec.get("description", "") or "")
lines.append(f" - {field_name}" + (f" ({field_type})" if field_type else "") + (f": {field_desc}" if field_desc else ""))
return "\n".join(lines)
class ElicitationHandler:
"""Handles ``elicitation/create`` requests for one MCP server; passed to
``ClientSession`` as ``elicitation_callback``. Form-mode requests route
through Hermes' approval system (CLI, TUI, Telegram, ...); URL-mode is
declined as unsupported. Fail-closed: any timeout, exception or unexpected
"""``elicitation_callback`` for one MCP server. Form-mode routes through Hermes' approval system
(CLI, TUI, Telegram, ...); URL-mode is declined. Fail-closed: any timeout, exception or unexpected
state returns decline/cancel, never a silent accept."""
# asyncio-side safety net over the approval's own input() timeout so the
# MCP loop never blocks indefinitely if the inner timeout is bypassed.
# asyncio-side safety net over the approval's own input() timeout so the MCP loop never blocks
# indefinitely if the inner timeout is bypassed.
_OUTER_TIMEOUT_GRACE_SECONDS = 5
# consent answer -> (ElicitResult action, metric); anything else declines.
_ANSWER_RESULTS = {"accept": ("accept", "accepted"), "cancel": ("cancel", "errors")}
def __init__(self, server_name: str, config: dict, owner: Optional["MCPServerTask"] = None):
self.server_name = server_name
# Default 5 min mirrors the gateway approval default so async surfaces
# (Telegram, Slack) have time to respond.
# 5 min mirrors the gateway approval default so async surfaces (Telegram, Slack) can respond.
self.timeout = _safe_numeric(config.get("timeout", 300), 300, float)
# Back-reference for the agent's contextvars snapshot; optional so the
# handler stays unit-testable in isolation.
# Back-reference for the agent's contextvars snapshot; optional for isolated unit tests.
self.owner = owner
self.metrics = {"requests": 0, "accepted": 0, "declined": 0, "errors": 0}
@@ -342,15 +265,12 @@ class ElicitationHandler:
def _result(self, action: str, metric: str):
"""Count *metric* and return ``ElicitResult(action)`` (accept carries empty content)."""
self.metrics[metric] += 1
if action == "accept":
return _core.ElicitResult(action="accept", content={})
return _core.ElicitResult(action=action)
return _core.ElicitResult(action=action, **({"content": {}} if action == "accept" else {}))
def _consent_thunk(self, message: str, description: str) -> Callable[[], str]:
"""Sync consent call, replaying the agent's contextvars snapshot when the
owner captured one: the recv-loop task does NOT inherit them, and
gateway-platform detection needs them. ``Context.run`` executes a
context once, so it is copied per elicitation."""
"""Sync consent call replaying the agent's contextvars snapshot when the owner captured one
(the recv-loop task does NOT inherit them; gateway-platform detection needs them).
``Context.run`` runs a context once, so it is copied per elicitation."""
from tools.approval import request_elicitation_consent
kwargs = {"timeout_seconds": int(self.timeout), "surface": f"mcp-elicitation/{self.server_name}"}
@@ -362,42 +282,28 @@ class ElicitationHandler:
async def __call__(self, context, params):
"""SDK elicitation callback (``ElicitationFnT``). Returns ElicitResult or ErrorData."""
self.metrics["requests"] += 1
# URL-mode (OAuth, payment) would need a browser + waiting for
# notifications/elicitation/complete — not implemented; decline cleanly.
if getattr(params, "mode", "form") == "url":
logger.info(
"MCP server '%s' requested URL-mode elicitation; declining (URL-mode elicitation not implemented)",
self.server_name,
)
if getattr(params, "mode", "form") == "url": # OAuth/payment: needs a browser + elicitation/complete; unsupported
logger.info("MCP server '%s' requested URL-mode elicitation; declining "
"(URL-mode elicitation not implemented)", self.server_name)
return self._result("decline", "declined")
message = getattr(params, "message", "") or f"MCP server '{self.server_name}' is requesting your approval"
# ``requestedSchema`` on mcp 1.x, ``requested_schema`` on 2.0 (pydantic
# aliases don't apply to attribute access) — read both or the user is
# asked to approve without seeing which fields the server wants.
# ``requestedSchema`` on mcp 1.x, ``requested_schema`` on 2.0 (aliases don't apply to attribute
# access) — read both or the user approves without seeing the fields.
schema = getattr(params, "requestedSchema", None) or getattr(params, "requested_schema", None) or {}
description = _format_elicitation_schema_summary(schema, self.server_name)
logger.info("MCP server '%s' elicitation request: %s", self.server_name, _sanitize_error(message)[:200])
# Lazy import avoids import-order coupling with early-bootstrap tools.approval.
try:
invoke_consent = self._consent_thunk(message, description)
try: # lazy import inside avoids import-order coupling with early-bootstrap tools.approval
invoke_consent = self._consent_thunk(message, _format_elicitation_schema_summary(schema, self.server_name))
except Exception as exc: # pragma: no cover -- defensive
logger.error("MCP server '%s' elicitation: approval system unavailable: %s", self.server_name, exc)
return self._result("decline", "errors")
# Offload the sync consent flow to a thread — inline it would freeze the
# MCP loop and every other RPC on this session.
try:
try: # off-thread: inline, the sync consent flow would freeze the MCP loop and every RPC on it
answer = await asyncio.wait_for(
asyncio.to_thread(invoke_consent), timeout=self.timeout + self._OUTER_TIMEOUT_GRACE_SECONDS,
)
asyncio.to_thread(invoke_consent), timeout=self.timeout + self._OUTER_TIMEOUT_GRACE_SECONDS)
except asyncio.TimeoutError:
logger.warning("MCP server '%s' elicitation timed out after %ds", self.server_name, int(self.timeout))
return self._result("cancel", "errors")
except Exception as exc:
logger.error("MCP server '%s' elicitation failed: %s", self.server_name, exc, exc_info=True)
return self._result("decline", "errors")
return self._result(*self._ANSWER_RESULTS.get(answer, ("decline", "declined")))

View File

@@ -1,4 +1,6 @@
"""Transport bring-up for MCPServerTask: stdio spawn (OSV preflight, watchdog wrap, child PID ledger), Streamable HTTP / SSE connect (preflight, identity header, client certs, OAuth), protocol negotiation and initial tool discovery. Split from tools/mcp_tool.py."""
"""Transport bring-up for MCPServerTask: stdio spawn (OSV preflight, watchdog wrap, child PID
ledger), Streamable HTTP / SSE connect (preflight, identity header, client certs, OAuth),
protocol negotiation and initial tool discovery. Split from tools/mcp_tool.py."""
import logging
import asyncio
@@ -14,13 +16,8 @@ logger = logging.getLogger("tools.mcp_tool")
# JSON-RPC ``initialize`` body used by the content-type preflight POST.
_PROBE_INITIALIZE_BODY = (
'{"jsonrpc":"2.0","id":"_probe",'
'"method":"initialize",'
'"params":{"protocolVersion":"2025-03-26",'
'"capabilities":{},'
'"clientInfo":{"name":"hermes-probe",'
'"version":"0.1"}}}'
)
'{"jsonrpc":"2.0","id":"_probe","method":"initialize","params":{"protocolVersion":"2025-03-26",'
'"capabilities":{},"clientInfo":{"name":"hermes-probe","version":"0.1"}}}')
def _content_type_base(resp) -> str:
@@ -32,18 +29,6 @@ def _is_2xx(resp) -> bool:
return 200 <= resp.status_code < 300
def _capture_pgids(pids: Set[int]) -> Dict[int, int]:
"""pgid per live pid. Captured while the child is alive — getpgid fails once it exits,
and the sweep needs it to reach reparented descendants."""
pgids: Dict[int, int] = {}
for pid in pids:
try:
pgids[pid] = os.getpgid(pid)
except (AttributeError, ProcessLookupError, OSError): # Windows / already exited
pass
return pgids
def _pgroup_alive(pgid: Optional[int]) -> bool:
"""Signal 0 to the group succeeds iff any member is alive (POSIX only)."""
_killpg = getattr(os, "killpg", None)
@@ -57,9 +42,8 @@ def _pgroup_alive(pgid: Optional[int]) -> bool:
async def _osv_malware_preflight(server_name: str, command: str, args: list) -> None:
"""OSV malware preflight: off-loop (blocking HTTPS) with a wall-clock bound so a stalled
handshake can't freeze discovery; fail-open on timeout. Must run against the REAL
command/args — the watchdog wrap rewrites argv to the supervisor, turning the check into a no-op."""
"""OSV malware preflight, off-loop with a wall-clock bound (fail-open on timeout). Must run on
the REAL command/args — the watchdog wrap rewrites argv to the supervisor (check becomes a no-op)."""
from tools.osv_check import check_package_for_malware
try:
malware_error = await asyncio.wait_for(
@@ -79,9 +63,8 @@ class MCPServerTransportMixin:
__slots__ = ()
def _advertises_tools(self) -> bool:
"""Whether the server advertises the ``tools`` capability. Prompt-/resource-only servers
omit it, and ``tools/list`` against them raises ``MCPError(-32601)``. True when no
capability info was captured (legacy fallback: always call list_tools)."""
"""Whether the server advertises ``tools`` (prompt-/resource-only servers omit it and
``tools/list`` raises -32601). True when no capability info was captured (legacy fallback)."""
caps = getattr(self.initialize_result, "capabilities", None)
return caps is None or getattr(caps, "tools", None) is not None
@@ -97,53 +80,45 @@ class MCPServerTransportMixin:
return kwargs
async def _negotiate_session(self, session, connect_timeout: float):
"""Negotiate the protocol era (``initialize`` vs ``server/discover``) and return its result.
Per-server ``protocol`` key: ``auto`` (default) tries the legacy handshake FIRST and falls
back to ``server/discover`` only when the server signals modern-only (-32022 / initialize
-32601) — the reverse of the SDK's discover-first mode, on purpose: zero extra round-trips
for the handshake-era servers that dominate today. ``stateless`` probes discover first (one
legacy retry on any error); ``legacy`` is handshake only, no fallback. Both result types
expose ``.capabilities``. A handshake TIMEOUT never triggers a fallback — it propagates."""
def initialize():
return asyncio.wait_for(session.initialize(), timeout=connect_timeout)
def discover():
return asyncio.wait_for(session.discover(), timeout=connect_timeout)
"""Negotiate the protocol era (``initialize`` vs ``server/discover``); both results expose
``.capabilities``. ``protocol: auto`` (default) tries the legacy handshake FIRST and falls back
to discover only on a modern-only signal (-32022 / initialize -32601) — deliberately the
reverse of the SDK's discover-first mode: zero extra round-trips for the handshake-era
servers that dominate today. ``stateless`` probes discover first (one legacy retry on any
error); ``legacy`` is handshake only. A handshake TIMEOUT never falls back — it propagates."""
def call(method: str):
return asyncio.wait_for(getattr(session, method)(), timeout=connect_timeout)
async def attempt(primary, fallback, should_fallback, log_fmt, *log_extra):
try:
return await primary()
except asyncio.TimeoutError:
raise
return await call(primary)
except Exception as exc:
if not should_fallback(exc):
if isinstance(exc, asyncio.TimeoutError) or not should_fallback(exc):
raise
logger.info(log_fmt, self.name, exc, *log_extra)
return await fallback()
return await call(fallback)
mode = str((self._config or {}).get("protocol", "auto")).lower().strip()
if mode in ("stateless", "modern", "2026-07-28"):
return await attempt(discover, initialize, lambda exc: True,
return await attempt("discover", "initialize", lambda exc: True,
"MCP server '%s': server/discover rejected (%s) despite "
"protocol=%s — falling back to the legacy handshake", mode)
if mode in ("legacy", "handshake"):
return await initialize()
return await call("initialize")
if mode != "auto":
logger.warning("MCP server '%s': unknown protocol=%r — treating as 'auto' "
"(valid: auto, stateless, legacy)", self.name, mode)
# mcp 1.x has no server/discover client — nothing to fall back to.
return await attempt(
initialize, discover,
lambda exc: _handshake_rejected_as_modern(exc) and hasattr(session, "discover"),
"initialize", "discover", lambda exc: _handshake_rejected_as_modern(exc) and hasattr(session, "discover"),
"MCP server '%s': legacy handshake rejected (%s) — "
"retrying via server/discover (2026-07-28 stateless server)")
async def _serve_session(self, session, connect_timeout: float,
label: str = "", mark_lifecycle: bool = False) -> str:
"""Handshake, discover, publish readiness, then serve until a lifecycle event.
Clears stale breaker state from a prior outage but leaves the session UNPROVEN: a
completed handshake is not proof of health (flapping transports handshake fine and drop
moments later); only keepalive or tool-call success clears the reconnect budget."""
"""Handshake, discover, publish readiness, then serve until a lifecycle event. Clears stale
breaker state but leaves the session UNPROVEN: flapping transports handshake fine and drop
moments later, so only keepalive or tool-call success clears the reconnect budget."""
self.initialize_result = await self._negotiate_session(session, connect_timeout)
self.session = session
if mark_lifecycle:
@@ -159,8 +134,8 @@ class MCPServerTransportMixin:
return reason
async def _serve_transport(self, transport_cm, label: str, connect_timeout: float) -> str:
"""Open *transport_cm*, wrap its streams in a ClientSession and serve it. Streams are
unpacked positionally: mcp 1.x yields ``(read, write, get_session_id)``, 2.x ``(read, write)``.
"""Open *transport_cm*, wrap its streams in a ClientSession and serve it. Streams are indexed,
not unpacked: mcp 1.x yields ``(read, write, get_session_id)``, 2.x ``(read, write)``.
A transport TaskGroup drop maps to ``"reconnect"`` instead of backoff/park."""
try:
async with transport_cm as _streams:
@@ -171,47 +146,39 @@ class MCPServerTransportMixin:
# ------------------------------------------------------------------ stdio
def _resolve_stdio_config(self, config: dict):
"""``(command, args, safe_env)`` from config, with the command resolved against the safe env."""
command = config.get("command")
if not command:
raise ValueError(f"MCP server '{self.name}' has no 'command' in config")
safe_env = _core._build_safe_env(config.get("env"))
command, safe_env = _core._resolve_stdio_command(command, safe_env)
return command, config.get("args", []), safe_env
def _track_spawned_children(self, new_pids: Set[int]) -> None:
"""Ledger the freshly spawned stdio children (pids, pgids, machine spawn ledger)."""
new_pgids = _capture_pgids(new_pids)
"""Ledger the freshly spawned stdio children (pids, pgids, machine spawn ledger). pgids are
captured while alive (getpgid fails once it exits; the sweep needs it for reparented descendants)."""
new_pgids: Dict[int, int] = {}
for pid in new_pids:
try:
new_pgids[pid] = os.getpgid(pid)
except (AttributeError, ProcessLookupError, OSError): # Windows / already exited
pass
with _core._lock:
for _pid in new_pids:
_stdio_pids[_pid] = self.name
_stdio_pids.update(dict.fromkeys(new_pids, self.name))
_stdio_pgids.update(new_pgids)
# Machine spawn ledger so startup sweeps can reap orphans after an unclean parent
# exit. Best-effort — never break startup.
# Machine spawn ledger (startup sweeps reap orphans after an unclean exit); best-effort.
for _pid in new_pids:
try:
from hermes_cli.process_identity import register_child
register_child(_pid, "mcp-helper")
except Exception:
logger.debug("spawn-ledger register_child failed for MCP helper pid %s", _pid, exc_info=True)
def _release_spawned_children(self, new_pids: Set[int]) -> None:
"""Drop the ledger entries; any child (or its pgroup) still alive means SDK teardown
failed (common on cancel mid-way on Linux, where setsid() children escape the cgroup)
— mark it orphaned for the next cleanup sweep."""
"""Drop the ledger entries; a child (or its pgroup) still alive means SDK teardown failed
(common on mid-way cancel on Linux: setsid() children escape) — mark it orphaned for the sweep."""
from gateway.status import _pid_exists
with _core._lock:
for pid in new_pids:
_stdio_pids.pop(pid, None)
# ``os.kill(pid, 0)`` is NOT a no-op on Windows; use the cross-platform check.
# The child may have exited while descendants remain in its pgroup.
# ``os.kill(pid, 0)`` is NOT a no-op on Windows; the child may be gone while
# descendants remain in its pgroup.
if _pid_exists(pid) or _pgroup_alive(_stdio_pgids.get(pid)):
_orphan_stdio_pids.add(pid)
_orphan_stdio_pid_servers[pid] = self.name
else:
# Nothing to reap — drop the pgid so PID reuse can't surface stale pgroup state.
else: # nothing to reap — drop the pgid so PID reuse can't surface stale pgroup state
_stdio_pgids.pop(pid, None)
async def _run_stdio(self, config: dict):
@@ -223,46 +190,40 @@ class MCPServerTransportMixin:
if not _core._ensure_mcp_sdk():
raise ImportError(f"MCP server '{self.name}' requires the 'mcp' Python SDK, but "
"it is not installed. Run `hermes setup` to install MCP support, then retry.")
command, args, safe_env = self._resolve_stdio_config(config)
await _osv_malware_preflight(self.name, command, args)
# Parent-death watchdog: an ungraceful Hermes exit (kill -9, crash) can't leave the
# child and its descendants running. POSIX-only (process groups); no-op elsewhere.
command = config.get("command")
if not command:
raise ValueError(f"MCP server '{self.name}' has no 'command' in config")
command, safe_env = _core._resolve_stdio_command(command, _core._build_safe_env(config.get("env")))
await _osv_malware_preflight(self.name, command, config.get("args", []))
# Parent-death watchdog so kill -9 / crash can't leave the child tree running (POSIX-only).
# AFTER the OSV preflight so the check inspects the real package.
command, args = _wrap_command_with_watchdog(command, args)
command, args = _wrap_command_with_watchdog(command, config.get("args", []))
server_params = _core.StdioServerParameters(
command=command, args=args, env=safe_env if safe_env else None, cwd=config.get("cwd"),
# Windows pipes can split non-UTF-8 bytes at chunk boundaries; substitute U+FFFD
# instead of raising UnicodeDecodeError.
encoding_error_handler="replace",
)
command=command, args=args, env=safe_env or None, cwd=config.get("cwd"),
# Windows pipes can split non-UTF-8 bytes at chunk boundaries; substitute, don't raise.
encoding_error_handler="replace")
session_kwargs = self._session_kwargs()
# Reap orphans from prior failed attempts before spawning, else each reconnect retry
# piles up zombie pairs. Unscoped on purpose (also reaps orphans of servers that never
# reconnect). Worker thread: the reaper blocks up to 2s (SIGTERM → wait → SIGKILL).
# Reap orphans of prior attempts first, else each retry piles up zombie pairs. Unscoped on
# purpose (also reaps servers that never reconnect). Off-loop: the reaper blocks up to 2s.
await asyncio.to_thread(_core._kill_orphaned_mcp_children)
# Snapshot child PIDs before spawning so the new one can be identified.
pids_before = _core._snapshot_child_pids()
pids_before = _core._snapshot_child_pids() # so the new child can be identified after spawn
new_pids: set = set()
# Route subprocess stderr to ~/.hermes/logs/mcp-stderr.log so server banners don't
# land on the user's TTY and corrupt the TUI.
# Subprocess stderr goes to ~/.hermes/logs/mcp-stderr.log so banners can't corrupt the TUI.
_core._write_stderr_log_header(self.name)
_errlog = _core._get_mcp_stderr_log()
try:
async with _core.stdio_client(server_params, errlog=_errlog) as (read_stream, write_stream):
# Capture the new PID for force-kill cleanup, filtering non-MCP children
# (slash_worker, LSP servers) that race into the snapshot window: they share
# the TUI parent's pgid, so leaking them into _stdio_pgids makes the shutdown
# killpg() kill the TUI itself.
errlog = _core._get_mcp_stderr_log()
async with _core.stdio_client(server_params, errlog=errlog) as (read_stream, write_stream):
# New PIDs for force-kill cleanup, minus non-MCP children (slash_worker, LSP) that
# race into the window: they share the TUI's pgid, so leaking them into _stdio_pgids
# would make the shutdown killpg() kill the TUI itself.
new_pids = _filter_mcp_children(_core._snapshot_child_pids() - pids_before)
if new_pids:
self._track_spawned_children(new_pids)
# Tracked on the connection so in-flight calls fail fast when the subprocess dies.
self._stdio_child_pids = set(new_pids)
self._stdio_child_pids = set(new_pids) # so in-flight calls fail fast when the child dies
async with _core.ClientSession(read_stream, write_stream, **session_kwargs) as session:
# Bound the handshake: ``connect_timeout`` only bounds the caller's
# ``.result()`` wait, not this coroutine. A server that never answers
# ``initialize`` would otherwise hang here forever, the ``finally`` below
# would never run, and the child + pipes would leak on every retry until EMFILE.
# Bound the handshake here: ``connect_timeout`` only bounds the caller's ``.result()``.
# A server that never answers ``initialize`` would otherwise hang forever, skip the
# ``finally`` and leak child + pipes on every retry until EMFILE.
connect_timeout = float(config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT))
return await self._serve_session(session, connect_timeout, mark_lifecycle=True)
finally:
@@ -274,67 +235,56 @@ class MCPServerTransportMixin:
async def _preflight_content_type(self, url: str, *, headers: Optional[dict] = None,
ssl_verify: bool = True, client_cert=None, timeout: float = 5.0) -> None:
"""Probe *url* for an MCP-shaped response before the SDK connects. A URL pointing at a
plain web page makes the SDK sit out the full ``connect_timeout`` before an opaque
``CancelledError``; this raises :class:`NonMcpEndpointError` within ``timeout`` instead.
Allow-list based: only a 2xx with a definite non-MCP content type is rejected, and only
after a JSON-RPC ``initialize`` POST also fails to look like MCP (some servers serve a UI
on GET but speak Streamable HTTP via POST). Missing content type, non-2xx, or transport
errors pass silently — the real handshake stays the source of truth. Uses its own httpx
client OUTSIDE the SDK's anyio task group so the error isn't wrapped in an ExceptionGroup."""
"""Probe *url* before the SDK connects: a plain web page makes the SDK sit out the full
``connect_timeout`` before an opaque ``CancelledError``; this raises NonMcpEndpointError within
``timeout`` instead. Allow-list based: only a 2xx with a definite non-MCP content type is
rejected, and only after a JSON-RPC ``initialize`` POST also fails to look like MCP (some
servers serve a UI on GET but speak MCP via POST). Missing content type, non-2xx or transport
errors pass silently — the real handshake stays the source of truth. Own httpx client, OUTSIDE
the SDK's anyio task group, so the error isn't wrapped in an ExceptionGroup."""
try:
import httpx as _httpx
except ImportError:
return # No httpx → skip probe; SDK import would have failed first.
client_kwargs: dict = {"verify": ssl_verify, "follow_redirects": True, "timeout": _httpx.Timeout(timeout)}
if client_cert is not None:
client_kwargs["cert"] = client_cert
probe_headers = dict(headers) if headers else {}
try:
async with _httpx.AsyncClient(**client_kwargs) as client:
async with _httpx.AsyncClient(verify=ssl_verify, follow_redirects=True, timeout=_httpx.Timeout(timeout),
**({"cert": client_cert} if client_cert is not None else {})) as client:
# HEAD is cheapest; fall back to GET on 405/501.
resp = await client.head(url, headers=probe_headers)
if resp.status_code in (405, 501):
resp = await client.get(url, headers=probe_headers)
# Non-MCP content type on HEAD/GET: try a JSON-RPC POST before rejecting, so
# POST-only servers aren't false positives.
# Non-MCP content type on HEAD/GET: try a JSON-RPC POST so POST-only servers pass.
ct = _content_type_base(resp)
if ct and ct not in self._MCP_CONTENT_TYPES and _is_2xx(resp):
post_resp = await client.post(
url,
url, content=_PROBE_INITIALIZE_BODY,
headers={**probe_headers, "Content-Type": "application/json",
"Accept": "application/json, text/event-stream"},
content=_PROBE_INITIALIZE_BODY,
)
"Accept": "application/json, text/event-stream"})
if _is_2xx(post_resp) and _content_type_base(post_resp) in self._MCP_CONTENT_TYPES:
resp = post_resp
except _httpx.HTTPError:
return # DNS/connect/timeout/transport error — let the SDK try.
# Only judge 2xx: a 4xx/5xx may be an auth challenge or transient error the real
# handshake handles correctly. No content type advertised → don't second-guess the SDK.
if not _is_2xx(resp):
return
# Only judge 2xx (4xx/5xx may be an auth challenge the handshake handles); no content type
# advertised → don't second-guess the SDK.
ct_base = _content_type_base(resp)
if not ct_base or ct_base in self._MCP_CONTENT_TYPES:
if not _is_2xx(resp) or not ct_base or ct_base in self._MCP_CONTENT_TYPES:
return
raise NonMcpEndpointError(
f"MCP server '{self.name}' at {url} returned Content-Type '{ct_base}', not an MCP "
f"response (expected one of: {', '.join(self._MCP_CONTENT_TYPES)}). The URL most likely "
"points at a web page rather than an MCP endpoint — check it resolves to a Streamable "
"HTTP / SSE endpoint (e.g. https://host/mcp, not https://host/)."
)
"HTTP / SSE endpoint (e.g. https://host/mcp, not https://host/).")
def _reconnect_or_reraise_group(self, eg: BaseExceptionGroup) -> str:
"""Map an SDK transport TaskGroup failure to a clean ``"reconnect"``. HTTP/SSE stream
pumps run in an anyio TaskGroup, so a transient stream drop escapes as a
``BaseExceptionGroup``; unmapped, ``run()`` would back off and eventually park the server
for 300s (deregistering its tools) over a sub-second glitch. Re-raise when it is not a
transient drop: shutdown in progress (``_shutdown_event`` is set before the task is
cancelled), the group carries KeyboardInterrupt/SystemExit or a real CancelledError (must
propagate), or no live session was reached this attempt (``_ready`` unset — connect
failures must go through backoff, not hot-loop)."""
"""Map an SDK transport TaskGroup failure to a clean ``"reconnect"``: HTTP/SSE stream pumps
run in an anyio TaskGroup, so a transient drop escapes as a ``BaseExceptionGroup`` that would
otherwise back off and park the server for 300s over a sub-second glitch. Re-raise when it is
not a transient drop: shutdown in progress (``_shutdown_event`` is set before cancel), the
group carries KeyboardInterrupt/SystemExit or a real CancelledError, or no live session was
reached this attempt (``_ready`` unset — connect failures must back off, not hot-loop)."""
if (self._shutdown_event.is_set()
or eg.split((KeyboardInterrupt, SystemExit))[0] is not None
or eg.split(asyncio.CancelledError)[0] is not None
@@ -345,9 +295,9 @@ class MCPServerTransportMixin:
return "reconnect"
def _build_oauth_auth(self, url: str, config: dict):
"""OAuth 2.1 PKCE via the central MCPOAuthManager so one provider is reused across
reconnects and shared with config-time CLI paths. On setup failure (e.g.
non-interactive without cached tokens) re-raise so only this server is reported failed."""
"""OAuth 2.1 PKCE via the central MCPOAuthManager (one provider reused across reconnects and
CLI paths). Setup failures (e.g. non-interactive without cached tokens) re-raise so only this
server is reported failed."""
if self._auth_type != "oauth":
return None
try:
@@ -368,46 +318,44 @@ class MCPServerTransportMixin:
raise ImportError(f"MCP server '{self.name}' requires SSE transport but "
"mcp.client.sse.sse_client is not available. "
"Upgrade the mcp package to get SSE support.")
# sse_read_timeout bounds the gap between SSE events. SSE servers commonly idle for
# minutes, so tool_timeout (60s) would drop the stream; 300s matches the Streamable
# HTTP read timeout.
sse_kwargs: dict = {"url": url, "headers": headers or None, "timeout": float(connect_timeout), "sse_read_timeout": 300.0}
if oauth_auth is not None:
# Forward OAuth to sse_client, else OAuth SSE servers 401 silently.
sse_kwargs["auth"] = oauth_auth
# sse_read_timeout bounds the gap between events: SSE servers idle for minutes, so 300s
# (matching the Streamable HTTP read timeout), not tool_timeout. ``auth`` must be forwarded
# or OAuth SSE servers 401 silently.
sse_kwargs: dict = {"url": url, "headers": headers or None, "timeout": float(connect_timeout),
"sse_read_timeout": 300.0, **({"auth": oauth_auth} if oauth_auth is not None else {})}
if client_cert is not None or ssl_verify is not True:
# sse_client has no verify/cert kwargs: wrap the SDK defaults (follow_redirects=True)
# in an httpx_client_factory, forwarding the SDK's (headers, auth, timeout) and
# layering TLS on top. The client MUST come from the SDK's own httpx module
# (httpx2 on mcp >= 2.0) — see sdk_httpx().
# sse_client has no verify/cert kwargs: an httpx_client_factory forwards the SDK's
# (headers, auth, timeout) and layers TLS on top. The client MUST come from the SDK's
# own httpx module (httpx2 on mcp >= 2.0) — see sdk_httpx().
_httpx_mod = _core.sdk_httpx()
def _mcp_http_client_factory(headers=None, timeout=None, auth=None):
kwargs: dict = {"follow_redirects": True, "verify": ssl_verify,
"timeout": timeout if timeout is not None else _httpx_mod.Timeout(30.0, read=300.0)}
kwargs.update({k: v for k, v in (("headers", headers), ("auth", auth), ("cert", client_cert)) if v is not None})
return _httpx_mod.AsyncClient(**kwargs)
sse_kwargs["httpx_client_factory"] = _mcp_http_client_factory
sse_kwargs["httpx_client_factory"] = lambda headers=None, timeout=None, auth=None: _httpx_mod.AsyncClient(
follow_redirects=True, verify=ssl_verify,
timeout=timeout if timeout is not None else _httpx_mod.Timeout(30.0, read=300.0),
**{k: v for k, v in (("headers", headers), ("auth", auth), ("cert", client_cert)) if v is not None})
return _core.sse_client(**sse_kwargs)
def _streamable_http_transport(self, url: str, headers: dict, connect_timeout: float,
ssl_verify, client_cert, oauth_auth,
strict_cfg_headers: bool, configured_header_names: set):
"""Streamable HTTP context manager (mcp >= 1.24.0: caller-owned httpx client)."""
"""Streamable HTTP context manager: mcp >= 1.24.0 gets a caller-owned httpx client; on the
deprecated API (mcp < 1.24.0) the SDK owns the client."""
if not _core._MCP_NEW_HTTP:
return self._legacy_http_transport(url, headers, connect_timeout, ssl_verify, oauth_auth, strict_cfg_headers)
# Build an explicit AsyncClient matching the SDK's create_mcp_http_client defaults. It
# MUST come from the SDK's httpx module (httpx2 on mcp >= 2.0) since the SDK sends its
# own Request objects through it — see sdk_httpx().
if strict_cfg_headers:
# Fail closed: without an owned client redirects can't be hooked.
raise ImportError(f"MCP server '{self.name}' requires mcp >= 1.24.0 to "
"enforce the portable redirect-header boundary "
"(strict_redirect_headers). Upgrade the mcp package.")
return _core.streamablehttp_client(url, headers=headers, timeout=float(connect_timeout), verify=ssl_verify,
**({"auth": oauth_auth} if oauth_auth is not None else {}))
# Explicit AsyncClient matching the SDK's create_mcp_http_client defaults; MUST come from
# the SDK's httpx module (httpx2 on mcp >= 2.0) since the SDK sends its own Requests through it.
httpx = _core.sdk_httpx()
_strip_auth_on_cross_origin_redirect = _make_redirect_header_stripper(
httpx.URL(url), strict=strict_cfg_headers, configured_header_names=configured_header_names)
client_kwargs: dict = {"follow_redirects": True, "timeout": httpx.Timeout(float(connect_timeout), read=300.0),
"verify": ssl_verify, "event_hooks": {"response": [_strip_auth_on_cross_origin_redirect]}}
if headers:
client_kwargs["headers"] = headers
client_kwargs.update({k: v for k, v in (("auth", oauth_auth), ("cert", client_cert)) if v is not None})
"verify": ssl_verify, **({"headers": headers} if headers else {}),
"event_hooks": {"response": [_strip_auth_on_cross_origin_redirect]},
**{k: v for k, v in (("auth", oauth_auth), ("cert", client_cert)) if v is not None}}
@asynccontextmanager
async def _owned_client_streams():
@@ -418,20 +366,6 @@ class MCPServerTransportMixin:
return _owned_client_streams()
def _legacy_http_transport(self, url: str, headers: dict, connect_timeout: float,
ssl_verify, oauth_auth, strict_cfg_headers: bool):
"""Deprecated API (mcp < 1.24.0): the SDK owns the httpx client."""
if strict_cfg_headers:
# Fail closed: without an owned client we cannot hook redirects, so the
# cross-origin header boundary cannot be enforced.
raise ImportError(f"MCP server '{self.name}' requires mcp >= 1.24.0 to "
"enforce the portable redirect-header boundary "
"(strict_redirect_headers). Upgrade the mcp package.")
http_kwargs: dict = {"headers": headers, "timeout": float(connect_timeout), "verify": ssl_verify}
if oauth_auth is not None:
http_kwargs["auth"] = oauth_auth
return _core.streamablehttp_client(url, **http_kwargs)
async def _run_http(self, config: dict):
"""Run the server using HTTP/StreamableHTTP (or SSE) transport."""
_core._ensure_mcp_sdk()
@@ -441,30 +375,23 @@ class MCPServerTransportMixin:
"Upgrade the mcp package to get HTTP support.")
url = config["url"]
headers = dict(config.get("headers") or {})
# Portable Agent Plugins v1 (strict_redirect_headers): configured headers MUST NOT
# follow a redirect to a different origin. Capture the configured names BEFORE
# client-generated headers are merged in.
strict_cfg_headers = bool(config.get("strict_redirect_headers"))
# Agent Plugins v1 strict_redirect_headers: configured headers MUST NOT follow a cross-origin
# redirect. Capture their names BEFORE client-generated headers are merged in.
configured_header_names = {key.lower() for key in headers}
# Optional per-user identity header; explicit headers of the same name win.
headers = _apply_identity_header(self.name, config, headers)
# Some servers require MCP-Protocol-Version on the initial request; seed it (case-insensitive
# user override wins) from the HANDSHAKE version, not the latest: ``initialize()``'s body speaks
# the handshake era, and a 2026-07-28 header would route it onto the server's per-request-envelope
# ladder, which rejects that body. The header must agree with what the body actually speaks.
# Some servers require MCP-Protocol-Version on the first request; seed it (user override
# wins) from the HANDSHAKE version, not the latest: a 2026-07-28 header would route the
# handshake-era ``initialize()`` body onto the per-request-envelope ladder, which rejects it.
if not any(key.lower() == "mcp-protocol-version" for key in headers):
headers["mcp-protocol-version"] = _core.LATEST_HANDSHAKE_VERSION
connect_timeout = config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT)
ssl_verify = config.get("ssl_verify", True)
client_cert = _resolve_client_cert(self.name, config)
oauth_auth = self._build_oauth_auth(url, config)
common = (url, headers, connect_timeout, config.get("ssl_verify", True), _resolve_client_cert(self.name, config),
self._build_oauth_auth(url, config), bool(config.get("strict_redirect_headers")))
if config.get("transport") == "sse":
transport = self._sse_transport(url, headers, connect_timeout, ssl_verify, client_cert, oauth_auth, strict_cfg_headers)
label = "SSE"
transport, label = self._sse_transport(*common), "SSE"
else:
transport = self._streamable_http_transport(url, headers, connect_timeout, ssl_verify, client_cert,
oauth_auth, strict_cfg_headers, configured_header_names)
transport = self._streamable_http_transport(*common, configured_header_names)
label = "HTTP" if _core._MCP_NEW_HTTP else "legacy HTTP"
return await self._serve_transport(transport, label, float(connect_timeout))
@@ -473,8 +400,7 @@ class MCPServerTransportMixin:
async def _discover_tools(self):
"""Discover tools from the connected session. Capability-gated: prompt-/resource-only
servers raise ``MCPError(-32601)`` on ``tools/list``, which would abort the connection."""
# Fresh transport: re-probe with cheap ``ping`` in case the server gained support
# across the reconnect.
# Fresh transport: re-probe ``ping`` in case the server gained support across the reconnect.
self._ping_unsupported = False
if self.session is None:
return
@@ -482,21 +408,19 @@ class MCPServerTransportMixin:
logger.info("MCP server '%s': does not advertise 'tools' capability — "
"skipping tools/list (prompts/resources remain available)", self.name)
self._tools = []
self._register_discovered_tools_if_needed()
return
async with self._rpc_lock:
self._list_cache_meta = {}
self._tools = await _core._paginate_full_list(
self.session.list_tools, "tools", self.name, cache_meta_out=self._list_cache_meta)
else:
async with self._rpc_lock:
self._list_cache_meta = {}
self._tools = await _core._paginate_full_list(
self.session.list_tools, "tools", self.name, cache_meta_out=self._list_cache_meta)
self._register_discovered_tools_if_needed()
def _register_discovered_tools_if_needed(self) -> None:
"""Publish freshly discovered tools for a registry-owned server if none are registered.
Initial registration normally happens in ``_discover_and_register_server`` after
``start()``. On reconnect, outage handling may clear ``_ready`` and deregister stale
tools; ownership via ``_servers`` authorizes publishing before readiness is restored so a
revival never comes back with zero tools. A server retained after a recoverable initial
failure is likewise owned before its first session, which authorizes its first publication."""
"""Publish freshly discovered tools for a registry-owned server if none are registered
(initial registration normally happens in ``_discover_and_register_server``). On reconnect,
outage handling may clear ``_ready`` and deregister stale tools; ownership via ``_servers``
authorizes publishing before readiness is restored so a revival never comes back with zero
tools — likewise a server retained after a recoverable initial failure."""
if self._registered_tool_names:
return
if not self._ready.is_set():
@@ -504,8 +428,7 @@ class MCPServerTransportMixin:
if _core._servers.get(self.name) is not self:
return
self._registered_tool_names = _core._register_server_tools(self.name, self, self._config)
# A retained initial-failure server that just published tools has recovered: drop
# its stale connect error from status surfaces.
# A retained initial-failure server that just published tools has recovered.
with _core._lock:
if _core._servers.get(self.name) is self:
_core._server_connect_errors.pop(self.name, None)