refactor(tools): collapse defensive layers and compact prose in MCP transport/loop/sampling/errors
This commit is contained in:
@@ -25,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 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:
|
||||
@@ -44,33 +40,29 @@ 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:
|
||||
@@ -90,9 +82,8 @@ 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)
|
||||
@@ -104,9 +95,8 @@ def _classify_mcp_failure(exc: BaseException) -> str:
|
||||
|
||||
|
||||
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}")
|
||||
|
||||
@@ -129,10 +119,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:
|
||||
@@ -166,10 +155,9 @@ def _resolve_client_cert(server_name: str, config: dict):
|
||||
|
||||
|
||||
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
|
||||
@@ -196,9 +184,8 @@ def _resolve_identity_header(server_name: str, config: dict):
|
||||
|
||||
|
||||
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."""
|
||||
"""Merge the identity header into ``headers`` in place; an explicit entry of the same name (any
|
||||
casing) wins — never silently override user config."""
|
||||
resolved = _resolve_identity_header(server_name, config)
|
||||
if resolved is None:
|
||||
return headers
|
||||
@@ -213,11 +200,9 @@ def _apply_identity_header(server_name: str, config: dict, headers: dict) -> dic
|
||||
|
||||
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):
|
||||
@@ -276,7 +261,7 @@ def _format_connect_error(exc: BaseException) -> str:
|
||||
return _sanitize_error("; ".join(deduped[:3]))
|
||||
|
||||
|
||||
# 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
|
||||
|
||||
@@ -291,21 +276,20 @@ 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(
|
||||
@@ -317,45 +301,37 @@ def _get_auth_error_types() -> tuple:
|
||||
|
||||
|
||||
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):
|
||||
if isinstance(exc, _http_status_error_types()):
|
||||
return getattr(exc.response, "status_code", None) == 401
|
||||
return 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")
|
||||
|
||||
# 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.
|
||||
|
||||
Iterative walk 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* a message-less ClosedResourceError."""
|
||||
# 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"))
|
||||
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()
|
||||
@@ -365,11 +341,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 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
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
"""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
|
||||
|
||||
@@ -21,8 +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."""
|
||||
"""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
|
||||
@@ -30,7 +29,7 @@ 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.
|
||||
# Best effort: an unlock/close failure must never propagate out of discovery.
|
||||
with contextlib.suppress(Exception):
|
||||
if os.name == "posix":
|
||||
import fcntl
|
||||
@@ -44,30 +43,29 @@ class _LockCookie:
|
||||
|
||||
|
||||
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."""
|
||||
"""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(fh.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)
|
||||
return True
|
||||
except OSError as e:
|
||||
if e.errno in (errno.EACCES, errno.EAGAIN, errno.EWOULDBLOCK):
|
||||
return False
|
||||
raise
|
||||
return True
|
||||
import portalocker
|
||||
try:
|
||||
portalocker.lock(fh, portalocker.LOCK_EX | portalocker.LOCK_NB)
|
||||
return True
|
||||
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
|
||||
@@ -88,8 +86,8 @@ 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
|
||||
@@ -97,8 +95,8 @@ def _mcp_loop_exception_handler(loop, 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
|
||||
home_override = get_hermes_home_override()
|
||||
@@ -142,9 +140,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
|
||||
|
||||
@@ -154,10 +151,9 @@ def _run_on_mcp_loop(coro_or_factory, timeout: float = 30):
|
||||
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)
|
||||
# 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))
|
||||
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)")
|
||||
@@ -172,23 +168,22 @@ def _run_on_mcp_loop(coro_or_factory, timeout: float = 30):
|
||||
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 "
|
||||
raise TimeoutError(f"MCP call timed out after {time.monotonic() - start_time:.1f}s "
|
||||
f"(configured timeout: {float(timeout):.1f}s)")
|
||||
wait_timeout = min(wait_timeout, remaining)
|
||||
try:
|
||||
return future.result(timeout=wait_timeout)
|
||||
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()
|
||||
|
||||
|
||||
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
|
||||
@@ -204,16 +199,13 @@ 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 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``."""
|
||||
poll_interval = 0.25
|
||||
iterations = max(1, int(max(float(timeout), 0.0) / poll_interval))
|
||||
for i in range(iterations):
|
||||
@@ -231,14 +223,11 @@ def _wait_for_server_session_ready(srv: Any, *, old_session: Any = None, timeout
|
||||
|
||||
|
||||
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."""
|
||||
"""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 hasattr(ready, "clear"):
|
||||
@@ -247,22 +236,23 @@ def _signal_reconnect_and_wait(server_name: str, srv: Any, *, op_description: st
|
||||
if hasattr(reconnect_event, "set"):
|
||||
reconnect_event.set()
|
||||
|
||||
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)
|
||||
|
||||
|
||||
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)
|
||||
_origin._mcp_thread = threading.Thread(
|
||||
target=_origin._mcp_loop.run_forever, name="mcp-event-loop", daemon=True)
|
||||
_origin._mcp_thread.start()
|
||||
|
||||
|
||||
@@ -279,8 +269,7 @@ def _stop_mcp_loop(*, only_if_idle: bool = False) -> bool:
|
||||
_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.
|
||||
# against a closed loop. shutdown_mcp_servers only reaps _servers; everything else ends here.
|
||||
stop_owned_by_loop = False
|
||||
if loop.is_running():
|
||||
from agent.async_utils import safe_schedule_threadsafe
|
||||
@@ -293,7 +282,8 @@ def _stop_mcp_loop(*, only_if_idle: bool = False) -> bool:
|
||||
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)
|
||||
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():
|
||||
|
||||
@@ -13,8 +13,8 @@ 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)
|
||||
|
||||
|
||||
@@ -49,8 +49,8 @@ def _tool_call_dict(tu, index: int) -> dict:
|
||||
|
||||
|
||||
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])
|
||||
tool_results = [b for b in blocks if _tool_use_id(b) is not _MISSING]
|
||||
@@ -75,8 +75,7 @@ def _convert_sampling_message(msg) -> List[dict]:
|
||||
|
||||
|
||||
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)
|
||||
@@ -92,13 +91,10 @@ def _response_total_tokens(response, 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}
|
||||
@@ -129,10 +125,8 @@ 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]:
|
||||
"""MCP SamplingMessages -> OpenAI format (per-block duck-typed dispatch)."""
|
||||
@@ -155,8 +149,7 @@ class SamplingHandler:
|
||||
self.server_name, response.model, _response_total_tokens(response, "?"), *args)
|
||||
|
||||
def _build_tool_use_result(self, choice, response):
|
||||
"""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
|
||||
@@ -197,14 +190,14 @@ class SamplingHandler:
|
||||
model = self._resolve_model(mcp_field(params, "model_preferences", "modelPreferences"))
|
||||
resolved_model = model or self.model_override 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)
|
||||
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). Server-provided tools are forwarded."""
|
||||
"""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)
|
||||
@@ -224,8 +217,7 @@ class SamplingHandler:
|
||||
max_tokens=max_tokens, 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
|
||||
@@ -250,8 +242,7 @@ class SamplingHandler:
|
||||
|
||||
|
||||
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}'."
|
||||
@@ -266,24 +257,21 @@ def _format_elicitation_schema_summary(schema: dict, server_name: str) -> str:
|
||||
|
||||
|
||||
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 state returns decline/cancel, never a silent accept."""
|
||||
"""``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}
|
||||
|
||||
@@ -294,14 +282,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}"}
|
||||
@@ -313,16 +299,15 @@ 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.
|
||||
# URL-mode (OAuth, payment) needs a browser + notifications/elicitation/complete — not implemented.
|
||||
if getattr(params, "mode", "form") == "url":
|
||||
logger.info("MCP server '%s' requested URL-mode elicitation; declining (URL-mode elicitation not implemented)",
|
||||
self.server_name)
|
||||
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 approves without seeing the fields.
|
||||
# ``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])
|
||||
@@ -332,8 +317,7 @@ class ElicitationHandler:
|
||||
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.
|
||||
# Off-thread: inline, the sync consent flow would freeze the MCP loop and every RPC on it.
|
||||
try:
|
||||
answer = await asyncio.wait_for(
|
||||
asyncio.to_thread(invoke_consent), timeout=self.timeout + self._OUTER_TIMEOUT_GRACE_SECONDS)
|
||||
|
||||
@@ -30,8 +30,8 @@ def _is_2xx(resp) -> bool:
|
||||
|
||||
|
||||
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."""
|
||||
"""pgid per live pid, captured while alive (getpgid fails once it exits; the sweep needs it
|
||||
to reach reparented descendants)."""
|
||||
pgids: Dict[int, int] = {}
|
||||
for pid in pids:
|
||||
try:
|
||||
@@ -54,9 +54,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(
|
||||
@@ -76,9 +75,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
|
||||
|
||||
@@ -94,13 +92,12 @@ 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."""
|
||||
"""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 initialize():
|
||||
return asyncio.wait_for(session.initialize(), timeout=connect_timeout)
|
||||
|
||||
@@ -110,10 +107,8 @@ class MCPServerTransportMixin:
|
||||
async def attempt(primary, fallback, should_fallback, log_fmt, *log_extra):
|
||||
try:
|
||||
return await primary()
|
||||
except asyncio.TimeoutError:
|
||||
raise
|
||||
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()
|
||||
@@ -131,13 +126,14 @@ class MCPServerTransportMixin:
|
||||
# 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"),
|
||||
"MCP server '%s': legacy handshake rejected (%s) — retrying via server/discover (2026-07-28 stateless server)")
|
||||
"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."""
|
||||
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 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:
|
||||
@@ -153,8 +149,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:
|
||||
@@ -170,8 +166,7 @@ class MCPServerTransportMixin:
|
||||
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)
|
||||
command, safe_env = _core._resolve_stdio_command(command, _core._build_safe_env(config.get("env")))
|
||||
return command, config.get("args", []), safe_env
|
||||
|
||||
def _track_spawned_children(self, new_pids: Set[int]) -> None:
|
||||
@@ -181,8 +176,7 @@ class MCPServerTransportMixin:
|
||||
for _pid in new_pids:
|
||||
_stdio_pids[_pid] = 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
|
||||
@@ -191,15 +185,14 @@ class MCPServerTransportMixin:
|
||||
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
|
||||
@@ -218,43 +211,37 @@ class MCPServerTransportMixin:
|
||||
"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. AFTER the OSV
|
||||
# preflight so the check inspects the real package.
|
||||
# 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)
|
||||
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.
|
||||
# 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()
|
||||
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)
|
||||
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:
|
||||
@@ -266,22 +253,20 @@ 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
|
||||
client_kwargs: dict = {"verify": ssl_verify, "follow_redirects": True, "timeout": _httpx.Timeout(timeout),
|
||||
**({"cert": client_cert} if client_cert is not None else {})}
|
||||
probe_headers = dict(headers) if headers else {}
|
||||
try:
|
||||
async with _httpx.AsyncClient(**client_kwargs) as client:
|
||||
@@ -289,8 +274,7 @@ class MCPServerTransportMixin:
|
||||
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(
|
||||
@@ -302,12 +286,10 @@ class MCPServerTransportMixin:
|
||||
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 "
|
||||
@@ -316,14 +298,12 @@ class MCPServerTransportMixin:
|
||||
"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
|
||||
@@ -334,9 +314,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:
|
||||
@@ -357,18 +337,15 @@ 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):
|
||||
@@ -386,17 +363,15 @@ class MCPServerTransportMixin:
|
||||
"""Streamable HTTP context manager (mcp >= 1.24.0: caller-owned httpx 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().
|
||||
# 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():
|
||||
@@ -411,15 +386,12 @@ class MCPServerTransportMixin:
|
||||
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.
|
||||
# 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.")
|
||||
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)
|
||||
return _core.streamablehttp_client(url, headers=headers, timeout=float(connect_timeout), verify=ssl_verify,
|
||||
**({"auth": oauth_auth} if oauth_auth is not None else {}))
|
||||
|
||||
async def _run_http(self, config: dict):
|
||||
"""Run the server using HTTP/StreamableHTTP (or SSE) transport."""
|
||||
@@ -430,17 +402,15 @@ 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.
|
||||
# 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.
|
||||
strict_cfg_headers = bool(config.get("strict_redirect_headers"))
|
||||
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)
|
||||
@@ -461,8 +431,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
|
||||
@@ -479,12 +448,11 @@ class MCPServerTransportMixin:
|
||||
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():
|
||||
@@ -492,8 +460,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)
|
||||
|
||||
Reference in New Issue
Block a user