diff --git a/tools/mcp_tool_errors.py b/tools/mcp_tool_errors.py index 3a0219be65..f2ac066da1 100644 --- a/tools/mcp_tool_errors.py +++ b/tools/mcp_tool_errors.py @@ -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: ", - 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: " — 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 diff --git a/tools/mcp_tool_loop.py b/tools/mcp_tool_loop.py index 9492914321..a2c175fd88 100644 --- a/tools/mcp_tool_loop.py +++ b/tools/mcp_tool_loop.py @@ -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(): diff --git a/tools/mcp_tool_sampling.py b/tools/mcp_tool_sampling.py index d8d87e17e5..3c12586bbc 100644 --- a/tools/mcp_tool_sampling.py +++ b/tools/mcp_tool_sampling.py @@ -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) diff --git a/tools/mcp_tool_transport.py b/tools/mcp_tool_transport.py index 66e1e8c0f3..ce65b8b869 100644 --- a/tools/mcp_tool_transport.py +++ b/tools/mcp_tool_transport.py @@ -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)