diff --git a/tools/mcp_tool_errors.py b/tools/mcp_tool_errors.py index 0d769422b1..fe79816d1c 100644 --- a/tools/mcp_tool_errors.py +++ b/tools/mcp_tool_errors.py @@ -1,4 +1,6 @@ -"""MCP connection/transport error classification: URL validation, TLS client certs, identity headers, redirect header stripping, exception-group unwrapping, auth/session-expired/method-not-found detection and connect-error formatting. Split from tools/mcp_tool.py.""" +"""MCP connection/transport error classification: URL validation, TLS client certs, identity +headers, redirect header stripping, exception-group unwrapping, auth/session-expired/ +method-not-found detection and connect-error formatting. Split from tools/mcp_tool.py.""" import asyncio import errno @@ -12,7 +14,6 @@ from tools.mcp_tool_common import _sanitize_error, _core logger = logging.getLogger("tools.mcp_tool") - # Stateless (2026-07-28) servers reject a legacy ``initialize`` with # UnsupportedProtocolVersion (-32022) or plain method-not-found. _JSONRPC_UNSUPPORTED_PROTOCOL_VERSION = -32022 @@ -30,7 +31,7 @@ def _jsonrpc_matches(exc: BaseException, code, codes: tuple, markers: tuple) -> if code in codes: return True msg = str(exc).lower() - return bool(msg) and any(marker in msg for marker in markers) + return any(marker in msg for marker in markers) def _handshake_rejected_as_modern(exc: BaseException) -> bool: @@ -98,16 +99,11 @@ def _classify_mcp_failure(exc: BaseException) -> str: or isinstance(root, (NonMcpEndpointError, InvalidMcpUrlError, FileNotFoundError)) or (isinstance(root, OSError) and getattr(root, "errno", None) == errno.ENOENT) # 401/403 HTTPStatusError that _is_auth_error's type-gate missed (auth types not importable here). - or _response_status(root) in (401, 403) + or getattr(getattr(root, "response", None), "status_code", None) in (401, 403) ) return "permanent" if permanent else "transient" -def _response_status(exc: BaseException): - """``exc.response.status_code`` for httpx-shaped errors, else None.""" - return getattr(getattr(exc, "response", None), "status_code", None) - - def _validate_remote_mcp_url(server_name: str, url: Any) -> str: """The stripped URL if it is a valid http(s) remote MCP URL. Raises InvalidMcpUrlError naming the server for non-strings, missing/other schemes (stdio servers use ``command``, not ``url``), @@ -252,7 +248,6 @@ def _exc_children(exc: BaseException) -> List[BaseException]: def _format_connect_error(exc: BaseException) -> str: """Render nested MCP connection errors into an actionable short message.""" - def _find_missing(current: BaseException) -> Optional[str]: if isinstance(current, FileNotFoundError): if getattr(current, "filename", None): @@ -318,8 +313,7 @@ def _get_auth_error_types() -> tuple: _optional_types("mcp.client.auth", "OAuthFlowError", "OAuthTokenError") + _optional_types("mcp.client.auth", "UnauthorizedError") # older SDKs + _optional_types("tools.mcp_oauth", "OAuthNonInteractiveError") - + list(_http_status_error_types()) - ) + + list(_http_status_error_types())) return _AUTH_ERROR_TYPES @@ -352,16 +346,15 @@ _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.""" + :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. transport_error_types = tuple(_optional_types("anyio", "BrokenResourceError", "ClosedResourceError", "EndOfStream")) - - # Iterative traversal over ``exceptions`` / ``__cause__`` / ``__context__`` with an - # identity-visited set AND a node budget (graphs can be deep or cyclic). Every reachable - # node is inspected so an InterruptedError anywhere overrides transport markers; the chain - # walk matters because SDK wrappers often raise a generic RuntimeError *from* the - # message-less ClosedResourceError. stack: "list[BaseException | None]" = [exc] seen: set[int] = set() transport_error_found = False @@ -377,7 +370,7 @@ def _is_session_expired_error(exc: BaseException) -> bool: # Messages vary across SDK versions and servers: match a narrow allow-list of stable # substrings, not exception type, to avoid false positives. msg = str(current).lower() - if isinstance(current, transport_error_types) or (msg and any(marker in msg for marker in _SESSION_EXPIRED_MARKERS)): + if isinstance(current, transport_error_types) or any(marker in msg for marker in _SESSION_EXPIRED_MARKERS): transport_error_found = True stack.extend(getattr(current, "exceptions", ())) stack.extend((getattr(current, "__cause__", None), getattr(current, "__context__", None))) diff --git a/tools/mcp_tool_loop.py b/tools/mcp_tool_loop.py index 7c3a0f57d0..9492914321 100644 --- a/tools/mcp_tool_loop.py +++ b/tools/mcp_tool_loop.py @@ -8,6 +8,7 @@ from __future__ import annotations import asyncio import concurrent.futures +import contextlib import errno import logging import os @@ -20,11 +21,8 @@ logger = logging.getLogger("tools.mcp_tool") class _LockCookie: - """Holds a cross-process file lock; ``release()`` drops it. - - The file object MUST stay open while the lock is held: both the fcntl and - the portalocker lock are tied to the descriptor's lifetime. - """ + """Holds a cross-process file lock; ``release()`` drops it. The file object MUST stay open + while the lock is held: both the fcntl and the portalocker lock are tied to the descriptor.""" def __init__(self, fh: Any) -> None: self._fh = fh @@ -32,47 +30,37 @@ class _LockCookie: def release(self) -> None: if self._fh is None: return - # Best effort on every step: an unlock/close failure must never - # propagate out of discovery. - try: + # Best effort on every step: an unlock/close failure must never propagate out of discovery. + with contextlib.suppress(Exception): if os.name == "posix": import fcntl fcntl.flock(self._fh.fileno(), fcntl.LOCK_UN) else: import portalocker portalocker.unlock(self._fh) - except Exception: - pass - try: + with contextlib.suppress(Exception): self._fh.close() - except Exception: - pass self._fh = None def _acquire_lock_on_fh(fh: Any) -> bool: - """Non-blocking exclusive lock (fcntl on POSIX, portalocker elsewhere). - - False when another process holds it; unexpected errors propagate so the - caller can treat locking as unavailable. - """ - fd = fh.fileno() + """Non-blocking exclusive lock (fcntl on POSIX, portalocker elsewhere). False when another + process holds it; unexpected errors propagate so the caller can treat locking as unavailable.""" if os.name == "posix": import fcntl try: - fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB) + 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 - else: - import portalocker - try: - portalocker.lock(fh, portalocker.LOCK_EX | portalocker.LOCK_NB) - return True - except portalocker.LockException: - return False + import portalocker + try: + portalocker.lock(fh, portalocker.LOCK_EX | portalocker.LOCK_NB) + return True + except portalocker.LockException: + return False def _try_acquire_mcp_discovery_lock() -> Any: @@ -84,24 +72,15 @@ def _try_acquire_mcp_discovery_lock() -> Any: try: from hermes_constants import get_hermes_home if _origin._MCP_DISCOVERY_LOCK_PATH is None: - _origin._MCP_DISCOVERY_LOCK_PATH = str( - get_hermes_home() / ".mcp-discovery.lock" - ) - lock_path = _origin._MCP_DISCOVERY_LOCK_PATH + _origin._MCP_DISCOVERY_LOCK_PATH = str(get_hermes_home() / ".mcp-discovery.lock") + fh = open(_origin._MCP_DISCOVERY_LOCK_PATH, "w", encoding="utf-8") except Exception: return _core._LOCK_UNAVAILABLE - - try: - fh = open(lock_path, "w", encoding="utf-8") - except Exception: - return _core._LOCK_UNAVAILABLE - try: acquired = _core._acquire_lock_on_fh(fh) except Exception: fh.close() return _core._LOCK_UNAVAILABLE - if acquired: return _core._LockCookie(fh) fh.close() @@ -121,12 +100,7 @@ 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).""" try: - from hermes_constants import ( - get_hermes_home_override, - reset_hermes_home_override, - set_hermes_home_override, - ) - + from hermes_constants import get_hermes_home_override, reset_hermes_home_override, set_hermes_home_override home_override = get_hermes_home_override() except Exception: return coro @@ -146,11 +120,7 @@ def _wrap_with_home_override(coro: "Coroutine") -> "Coroutine": def _wrap_with_dashboard_oauth_flow(coro): """Propagate a dashboard OAuth flow onto the dedicated MCP loop task.""" try: - from tools.mcp_dashboard_oauth import ( - dashboard_oauth_flow, - get_dashboard_oauth_flow, - ) - + from tools.mcp_dashboard_oauth import dashboard_oauth_flow, get_dashboard_oauth_flow flow = get_dashboard_oauth_flow() except Exception: return coro @@ -172,12 +142,9 @@ 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 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.""" from tools.interrupt import is_interrupted from agent.async_utils import safe_schedule_threadsafe @@ -186,58 +153,42 @@ def _run_on_mcp_loop(coro_or_factory, timeout: float = 30): if asyncio.iscoroutine(coro_or_factory): coro_or_factory.close() raise RuntimeError("MCP event loop is not running") - coro = coro_or_factory() if callable(coro_or_factory) else coro_or_factory - - # Tasks created via run_coroutine_threadsafe copy the LOOP thread's - # context, so a per-request profile scope would vanish here; re-establish - # it inside the task's own context. + # Tasks created via run_coroutine_threadsafe copy the LOOP thread's context, so a per-request + # profile scope would vanish here; re-establish it inside the task's own context. coro = _core._wrap_with_home_override(coro) coro = _core._wrap_with_dashboard_oauth_flow(coro) - - future = safe_schedule_threadsafe( - coro, loop, - logger=logger, - log_message="MCP scheduling failed", - ) + future = safe_schedule_threadsafe(coro, loop, logger=logger, log_message="MCP scheduling failed") if future is None: raise RuntimeError("MCP event loop unavailable (failed to schedule)") start_time = time.monotonic() deadline = None if timeout is None else start_time + timeout - while True: if is_interrupted(): future.cancel() raise InterruptedError("User sent a new message") - wait_timeout = 0.1 if deadline is not None: remaining = deadline - time.monotonic() if remaining <= 0: future.cancel() elapsed = time.monotonic() - start_time - raise TimeoutError( - f"MCP call timed out after {elapsed:.1f}s " - f"(configured timeout: {float(timeout):.1f}s)" - ) + raise TimeoutError(f"MCP call timed out after {elapsed:.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 this also fires for the coroutine's own timeout: a + # done future must yield its outcome. if future.done(): return future.result() - continue + def _signal_reconnect(server: Any) -> bool: - """Ask a server task to rebuild its transport, thread-safely. - - Handlers run on caller threads while the event lives on the MCP loop, so - it is set via ``call_soon_threadsafe`` when the loop runs (direct - ``.set()`` otherwise). False when the server has no reconnect machinery. - """ + """Ask a server task to rebuild its transport, thread-safely. 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.""" event = getattr(server, "_reconnect_event", None) if event is None: return False @@ -258,30 +209,20 @@ def reconnect_mcp_server(server_name: str) -> bool: return _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``. - """ +def _wait_for_server_session_ready(srv: Any, *, old_session: Any = None, timeout: float = 15.0) -> bool: + """Poll until the server exposes a usable, ready session. During a reconnect ``srv.session`` + is briefly None or still the stale object; retrying blindly there burns breaker strikes. With + ``old_session`` the observed session must differ from it. Iteration-bounded, not + deadline-bounded: tests freeze ``time.monotonic``.""" poll_interval = 0.25 iterations = max(1, int(max(float(timeout), 0.0) / poll_interval)) for i in range(iterations): session = getattr(srv, "session", None) ready = getattr(srv, "_ready", None) - is_ready = True - if ready is not None and hasattr(ready, "is_set"): - try: - is_ready = bool(ready.is_set()) - except Exception: - is_ready = True + try: + is_ready = bool(ready.is_set()) if hasattr(ready, "is_set") else True + except Exception: + is_ready = True if session is not None and session is not old_session and is_ready: return True if i < iterations - 1: @@ -289,19 +230,10 @@ def _wait_for_server_session_ready( return False -def _signal_reconnect_and_wait( - server_name: str, - srv: Any, - *, - op_description: str, - timeout: float = 15.0, -) -> bool: - """Request a transport rebuild and wait for the fresh session. - - ``_ready`` is cleared on the loop BEFORE ``_reconnect_event`` is set; - otherwise the readiness poll returns immediately and retries against the - same dead session. - """ +def _signal_reconnect_and_wait(server_name: str, srv: Any, *, op_description: str, timeout: float = 15.0) -> bool: + """Request a transport rebuild and wait for the fresh session. ``_ready`` is cleared on the + loop BEFORE ``_reconnect_event`` is set; otherwise the readiness poll returns immediately and + retries against the same dead session.""" loop = _core._mcp_loop if loop is None or not loop.is_running(): return False @@ -309,41 +241,28 @@ def _signal_reconnect_and_wait( def _request_reconnect() -> None: ready = getattr(srv, "_ready", None) - if ready is not None and hasattr(ready, "clear"): + if hasattr(ready, "clear"): ready.clear() reconnect_event = getattr(srv, "_reconnect_event", None) - if reconnect_event is not None and hasattr(reconnect_event, "set"): + if hasattr(reconnect_event, "set"): reconnect_event.set() - logger.info( - "MCP server '%s': %s requesting transport reconnect", - server_name, op_description, - ) + logger.info("MCP server '%s': %s requesting transport reconnect", server_name, op_description) loop.call_soon_threadsafe(_request_reconnect) - return _core._wait_for_server_session_ready( - srv, - old_session=old_session, - timeout=timeout, - ) + return _core._wait_for_server_session_ready(srv, old_session=old_session, timeout=timeout) def _ensure_mcp_loop(): - """Start the background event loop thread if not already running. - - The loop/thread handles live on the ORIGIN module (tests read and reset - ``tools.mcp_tool._mcp_loop``), so they are written there, never here. - """ + """Start the background 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.""" 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() @@ -359,38 +278,29 @@ def _stop_mcp_loop(*, only_if_idle: bool = False) -> bool: _origin._mcp_loop = None _origin._mcp_thread = None if loop is not None: - # Drain before stopping: tasks still suspended when the loop closes - # get resumed by the GC against a closed loop. shutdown_mcp_servers - # only reaps servers held in _servers; everything else ends up here. + # Drain before stopping: tasks still suspended when the loop closes get resumed by the GC + # against a closed loop. shutdown_mcp_servers only reaps servers held in _servers; + # everything else ends up here. stop_owned_by_loop = False if loop.is_running(): from agent.async_utils import safe_schedule_threadsafe future = safe_schedule_threadsafe( - _core._drain_and_stop_mcp_loop(), loop, - logger=logger, - log_message="MCP loop drain: failed to schedule", - log_level=logging.WARNING, - ) + _core._drain_and_stop_mcp_loop(), loop, logger=logger, + log_message="MCP loop drain: failed to schedule", log_level=logging.WARNING) if future is not None: stop_owned_by_loop = True try: future.result(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT + 1) except TimeoutError: - logger.warning( - "Timed out waiting for MCP loop drain after %.1fs", - _core._MCP_LOOP_DRAIN_TIMEOUT + 1, - ) + 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(): try: - loop.run_until_complete( - _core._drain_mcp_loop_tasks(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT) - ) + loop.run_until_complete(_core._drain_mcp_loop_tasks(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT)) except BaseException as exc: logger.warning("Error draining stopped MCP loop tasks: %s", exc) - if not stop_owned_by_loop and loop.is_running(): loop.call_soon_threadsafe(loop.stop) if thread is not None: diff --git a/tools/mcp_tool_sampling.py b/tools/mcp_tool_sampling.py index 598a3bd6e0..d8d87e17e5 100644 --- a/tools/mcp_tool_sampling.py +++ b/tools/mcp_tool_sampling.py @@ -13,9 +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 (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.""" return mcp_field(block, "tool_use_id", "toolUseId", _MISSING) @@ -45,22 +44,15 @@ def _content_part(block) -> Optional[dict]: def _tool_call_dict(tu, index: int) -> dict: args = tu.input - return { - "id": getattr(tu, "id", f"call_{index}"), - "type": "function", - "function": { - "name": tu.name, - "arguments": json.dumps(args, ensure_ascii=False) if isinstance(args, dict) else str(args), - }, - } + return {"id": getattr(tu, "id", f"call_{index}"), "type": "function", "function": { + "name": tu.name, "arguments": json.dumps(args, ensure_ascii=False) if isinstance(args, dict) else str(args)}} def _convert_sampling_message(msg) -> List[dict]: """One MCP SamplingMessage -> OpenAI-format messages (tool results first, then either an assistant tool_calls message or plain content).""" blocks = msg.content_as_list if hasattr(msg, "content_as_list") else ( - msg.content if isinstance(msg.content, list) else [msg.content] - ) + msg.content if isinstance(msg.content, list) else [msg.content]) tool_results = [b for b in blocks if _tool_use_id(b) is not _MISSING] tool_uses = [b for b in blocks if _is_tool_use(b) and _tool_use_id(b) is _MISSING] content_blocks = [b for b in blocks if _tool_use_id(b) is _MISSING and not _is_tool_use(b)] @@ -89,10 +81,8 @@ def _parse_tool_call_arguments(server_name: str, args) -> dict: try: return json.loads(args) except (json.JSONDecodeError, ValueError): - logger.warning( - "MCP server '%s': malformed tool_calls arguments from LLM (wrapping as raw): %.100s", - server_name, args, - ) + logger.warning("MCP server '%s': malformed tool_calls arguments from LLM (wrapping as raw): %.100s", + server_name, args) return {"_raw": args} return args if isinstance(args, dict) else {"_raw": str(args)} @@ -102,15 +92,13 @@ 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``. + """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. - """ + 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.""" _STOP_REASON_MAP = {"stop": "endTurn", "length": "maxTokens", "tool_calls": "toolUse"} _LOG_LEVELS = {"debug": logging.DEBUG, "info": logging.INFO, "warning": logging.WARNING} @@ -147,8 +135,7 @@ class SamplingHandler: return None def _convert_messages(self, params) -> List[dict]: - """Convert MCP SamplingMessages to OpenAI format (``content_as_list`` - when the SDK provides it; per-block duck-typed dispatch).""" + """MCP SamplingMessages -> OpenAI format (per-block duck-typed dispatch).""" return [m for msg in params.messages for m in _convert_sampling_message(msg)] @staticmethod @@ -164,14 +151,12 @@ class SamplingHandler: return self._error(message) def _log_response(self, response, suffix: str = "", *args) -> None: - logger.log( - self.audit_level, "MCP server '%s' sampling response: model=%s, tokens=%s" + suffix, - self.server_name, response.model, _response_total_tokens(response, "?"), *args, - ) + logger.log(self.audit_level, "MCP server '%s' sampling response: model=%s, tokens=%s" + suffix, + self.server_name, response.model, _response_total_tokens(response, "?"), *args) def _build_tool_use_result(self, choice, response): - """Build a CreateMessageResultWithTools from an LLM tool_calls response, - subject to tool-loop governance (``max_tool_rounds``; 0 disables).""" + """CreateMessageResultWithTools from an LLM tool_calls response, subject to tool-loop + governance (``max_tool_rounds``; 0 disables).""" self.metrics["tool_use_count"] += 1 if self.max_tool_rounds == 0: self._tool_loop_count = 0 @@ -180,61 +165,46 @@ class SamplingHandler: if self._tool_loop_count > self.max_tool_rounds: self._tool_loop_count = 0 return self._error( - f"Tool loop limit exceeded for server '{self.server_name}' (max {self.max_tool_rounds} rounds)" - ) + f"Tool loop limit exceeded for server '{self.server_name}' (max {self.max_tool_rounds} rounds)") content_blocks = [ - _core.ToolUseContent( - type="tool_use", id=tc.id, name=tc.function.name, - input=_parse_tool_call_arguments(self.server_name, tc.function.arguments), - ) - for tc in choice.message.tool_calls - ] + _core.ToolUseContent(type="tool_use", id=tc.id, name=tc.function.name, + input=_parse_tool_call_arguments(self.server_name, tc.function.arguments)) + for tc in choice.message.tool_calls] self._log_response(response, ", tool_calls=%d", len(content_blocks)) return _core.CreateMessageResultWithTools( - role="assistant", content=content_blocks, model=response.model, stopReason="toolUse", - ) + role="assistant", content=content_blocks, model=response.model, stopReason="toolUse") def _build_text_result(self, choice, response): - """Build a CreateMessageResult from a normal text response (resets the tool loop).""" + """CreateMessageResult from a normal text response (resets the tool loop).""" self._tool_loop_count = 0 self._log_response(response) return _core.CreateMessageResult( - role="assistant", + role="assistant", model=response.model, content=_core.TextContent(type="text", text=_sanitize_error(choice.message.content or "")), - model=response.model, - stopReason=self._STOP_REASON_MAP.get(choice.finish_reason, "endTurn"), - ) + stopReason=self._STOP_REASON_MAP.get(choice.finish_reason, "endTurn")) def session_kwargs(self) -> dict: """Kwargs to pass to ClientSession for sampling support.""" - return { - "sampling_callback": self, - "sampling_capabilities": _core.SamplingCapability(tools=_core.SamplingToolsCapability()), - } + return {"sampling_callback": self, + "sampling_capabilities": _core.SamplingCapability(tools=_core.SamplingToolsCapability())} def _admit(self, params): - """Rate-limit + allowed_models gate. Returns ``(resolved_model, None)`` - or ``(None, ErrorData)``.""" + """Rate-limit + allowed_models gate. Returns ``(resolved_model, None)`` or ``(None, ErrorData)``.""" if not self._check_rate_limit(): logger.warning("MCP server '%s' sampling rate limit exceeded (%d/min)", self.server_name, self.max_rpm) return None, self._fail( - f"Sampling rate limit exceeded for server '{self.server_name}' ({self.max_rpm} requests/minute)" - ) + f"Sampling rate limit exceeded for server '{self.server_name}' ({self.max_rpm} requests/minute)") model = self._resolve_model(mcp_field(params, "model_preferences", "modelPreferences")) resolved_model = model or self.model_override or "" if self.allowed_models and resolved_model and resolved_model not in self.allowed_models: - logger.warning( - "MCP server '%s' requested model '%s' not in allowed_models", self.server_name, resolved_model, - ) - return None, self._fail( - f"Model '{resolved_model}' not allowed for server " - f"'{self.server_name}'. Allowed: {', '.join(self.allowed_models)}" - ) + logger.warning("MCP server '%s' requested model '%s' not in allowed_models", self.server_name, resolved_model) + return None, self._fail(f"Model '{resolved_model}' not allowed for server " + f"'{self.server_name}'. Allowed: {', '.join(self.allowed_models)}") return resolved_model, None def _build_llm_call(self, params, resolved_model: str) -> Callable[[], object]: - """Translate the sampling params into a zero-arg sync ``call_llm`` thunk - (run off-loop so the MCP loop is not blocked).""" + """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.""" from agent.auxiliary_client import call_llm messages = self._convert_messages(params) @@ -243,29 +213,15 @@ class SamplingHandler: messages.insert(0, {"role": "system", "content": system_prompt}) max_tokens = min(mcp_field(params, "max_tokens", "maxTokens", self.max_tokens_cap), self.max_tokens_cap) temperature = getattr(params, "temperature", None) - # Forward server-provided tools. server_tools = getattr(params, "tools", None) - tools = [ - { - "type": "function", - "function": { - "name": getattr(t, "name", ""), - "description": getattr(t, "description", "") or "", - "parameters": _normalize_mcp_input_schema(mcp_field(t, "input_schema", "inputSchema")), - }, - } - for t in server_tools - ] if server_tools else None - - logger.log( - self.audit_level, - "MCP server '%s' sampling request: model=%s, max_tokens=%d, messages=%d", - self.server_name, resolved_model, max_tokens, len(messages), - ) - return lambda: call_llm( - task="mcp", model=resolved_model or None, messages=messages, temperature=temperature, - max_tokens=max_tokens, tools=tools, timeout=self.timeout, - ) + tools = [{"type": "function", "function": { + "name": getattr(t, "name", ""), "description": getattr(t, "description", "") or "", + "parameters": _normalize_mcp_input_schema(mcp_field(t, "input_schema", "inputSchema"))}} + for t in server_tools] if server_tools else None + logger.log(self.audit_level, "MCP server '%s' sampling request: model=%s, max_tokens=%d, messages=%d", + self.server_name, resolved_model, max_tokens, len(messages)) + return lambda: call_llm(task="mcp", model=resolved_model or None, messages=messages, temperature=temperature, + max_tokens=max_tokens, tools=tools, timeout=self.timeout) async def __call__(self, context, params): """SDK sampling callback (``SamplingFnT``). Returns CreateMessageResult, @@ -280,11 +236,9 @@ class SamplingHandler: return self._fail(f"Sampling LLM call timed out after {self.timeout}s for server '{self.server_name}'") except Exception as exc: return self._fail(f"Sampling LLM call failed: {_sanitize_error(_exc_str(exc))}") - # Empty choices happen on content filtering / provider errors. if not getattr(response, "choices", None): return self._fail(f"LLM returned empty response (no choices) for server '{self.server_name}'") - choice = response.choices[0] self.metrics["requests"] += 1 total_tokens = _response_total_tokens(response, 0) @@ -296,12 +250,11 @@ 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.""" + """Render a flat-object requested_schema as a human-readable field list (names, types, + descriptions) so the user knows what they're approving.""" props = schema.get("properties") if isinstance(schema, dict) else None if not isinstance(props, dict) or not props: return f"Approval requested by MCP server '{server_name}'." - lines = [f"Fields requested by MCP server '{server_name}':"] for field_name, field_spec in props.items(): spec = field_spec if isinstance(field_spec, dict) else {} @@ -313,25 +266,24 @@ 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.""" + """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.""" - # 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. + # Default 5 min mirrors the gateway approval default so async surfaces (Telegram, Slack) + # have time to 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 so the handler stays + # unit-testable in isolation. self.owner = owner self.metrics = {"requests": 0, "accepted": 0, "declined": 0, "errors": 0} @@ -347,10 +299,9 @@ class ElicitationHandler: return _core.ElicitResult(action=action) 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, and gateway-platform detection needs them. + ``Context.run`` executes a context once, so it is copied per elicitation.""" from tools.approval import request_elicitation_consent kwargs = {"timeout_seconds": int(self.timeout), "surface": f"mcp-elicitation/{self.server_name}"} @@ -362,42 +313,34 @@ class ElicitationHandler: async def __call__(self, context, params): """SDK elicitation callback (``ElicitationFnT``). Returns ElicitResult or ErrorData.""" self.metrics["requests"] += 1 - # URL-mode (OAuth, payment) would need a browser + waiting for # notifications/elicitation/complete — not implemented; decline cleanly. if getattr(params, "mode", "form") == "url": - logger.info( - "MCP server '%s' requested URL-mode elicitation; declining (URL-mode elicitation not implemented)", - self.server_name, - ) + logger.info("MCP server '%s' requested URL-mode elicitation; declining (URL-mode elicitation not implemented)", + self.server_name) return self._result("decline", "declined") message = getattr(params, "message", "") or f"MCP server '{self.server_name}' is requesting your approval" - # ``requestedSchema`` on mcp 1.x, ``requested_schema`` on 2.0 (pydantic - # aliases don't apply to attribute access) — read both or the user is - # asked to approve without seeing which fields the server wants. + # ``requestedSchema`` on mcp 1.x, ``requested_schema`` on 2.0 (pydantic aliases don't apply + # to attribute access) — read both or the user approves without seeing the fields. schema = getattr(params, "requestedSchema", None) or getattr(params, "requested_schema", None) or {} description = _format_elicitation_schema_summary(schema, self.server_name) logger.info("MCP server '%s' elicitation request: %s", self.server_name, _sanitize_error(message)[:200]) - # Lazy import avoids import-order coupling with early-bootstrap tools.approval. try: invoke_consent = self._consent_thunk(message, description) 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. + # Offload the sync consent flow to a thread — inline it would freeze the MCP loop and + # every other RPC on this session. try: answer = await asyncio.wait_for( - asyncio.to_thread(invoke_consent), timeout=self.timeout + self._OUTER_TIMEOUT_GRACE_SECONDS, - ) + asyncio.to_thread(invoke_consent), timeout=self.timeout + self._OUTER_TIMEOUT_GRACE_SECONDS) except asyncio.TimeoutError: logger.warning("MCP server '%s' elicitation timed out after %ds", self.server_name, int(self.timeout)) return self._result("cancel", "errors") except Exception as exc: logger.error("MCP server '%s' elicitation failed: %s", self.server_name, exc, exc_info=True) return self._result("decline", "errors") - return self._result(*self._ANSWER_RESULTS.get(answer, ("decline", "declined"))) diff --git a/tools/mcp_tool_transport.py b/tools/mcp_tool_transport.py index 407a30a3e8..604bcd73bd 100644 --- a/tools/mcp_tool_transport.py +++ b/tools/mcp_tool_transport.py @@ -1,4 +1,6 @@ -"""Transport bring-up for MCPServerTask: stdio spawn (OSV preflight, watchdog wrap, child PID ledger), Streamable HTTP / SSE connect (preflight, identity header, client certs, OAuth), protocol negotiation and initial tool discovery. Split from tools/mcp_tool.py.""" +"""Transport bring-up for MCPServerTask: stdio spawn (OSV preflight, watchdog wrap, child PID +ledger), Streamable HTTP / SSE connect (preflight, identity header, client certs, OAuth), +protocol negotiation and initial tool discovery. Split from tools/mcp_tool.py.""" import logging import asyncio @@ -14,12 +16,8 @@ logger = logging.getLogger("tools.mcp_tool") # JSON-RPC ``initialize`` body used by the content-type preflight POST. _PROBE_INITIALIZE_BODY = ( - '{"jsonrpc":"2.0","id":"_probe",' - '"method":"initialize",' - '"params":{"protocolVersion":"2025-03-26",' - '"capabilities":{},' - '"clientInfo":{"name":"hermes-probe",' - '"version":"0.1"}}}' + '{"jsonrpc":"2.0","id":"_probe","method":"initialize","params":{"protocolVersion":"2025-03-26",' + '"capabilities":{},"clientInfo":{"name":"hermes-probe","version":"0.1"}}}' ) @@ -133,13 +131,10 @@ class MCPServerTransportMixin: "(valid: auto, stateless, legacy)", self.name, mode) # mcp 1.x has no server/discover client — nothing to fall back to. return await attempt( - initialize, discover, - lambda exc: _handshake_rejected_as_modern(exc) and hasattr(session, "discover"), - "MCP server '%s': legacy handshake rejected (%s) — " - "retrying via server/discover (2026-07-28 stateless server)") + initialize, discover, lambda exc: _handshake_rejected_as_modern(exc) and hasattr(session, "discover"), + "MCP server '%s': legacy handshake rejected (%s) — retrying via server/discover (2026-07-28 stateless server)") - async def _serve_session(self, session, connect_timeout: float, - label: str = "", mark_lifecycle: bool = False) -> str: + 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 @@ -187,12 +182,11 @@ 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 so startup sweeps can reap orphans after an unclean parent exit. + # Best-effort — never break startup. for _pid in new_pids: try: from hermes_cli.process_identity import register_child - register_child(_pid, "mcp-helper") except Exception: logger.debug("spawn-ledger register_child failed for MCP helper pid %s", _pid, exc_info=True) @@ -225,44 +219,43 @@ 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: 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. 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. - encoding_error_handler="replace", - ) + 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 + # 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). 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. + # Route subprocess stderr to ~/.hermes/logs/mcp-stderr.log so server banners don't land + # on the user's TTY and 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. + # (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. 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: ``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. connect_timeout = float(config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT)) return await self._serve_session(session, connect_timeout, mark_lifecycle=True) finally: @@ -302,11 +295,9 @@ class MCPServerTransportMixin: ct = _content_type_base(resp) if ct and ct not in self._MCP_CONTENT_TYPES and _is_2xx(resp): post_resp = await client.post( - url, + url, content=_PROBE_INITIALIZE_BODY, headers={**probe_headers, "Content-Type": "application/json", - "Accept": "application/json, text/event-stream"}, - content=_PROBE_INITIALIZE_BODY, - ) + "Accept": "application/json, text/event-stream"}) if _is_2xx(post_resp) and _content_type_base(post_resp) in self._MCP_CONTENT_TYPES: resp = post_resp except _httpx.HTTPError: @@ -323,8 +314,7 @@ class MCPServerTransportMixin: f"MCP server '{self.name}' at {url} returned Content-Type '{ct_base}', not an MCP " f"response (expected one of: {', '.join(self._MCP_CONTENT_TYPES)}). The URL most likely " "points at a web page rather than an MCP endpoint — check it resolves to a Streamable " - "HTTP / SSE endpoint (e.g. https://host/mcp, not https://host/)." - ) + "HTTP / SSE endpoint (e.g. https://host/mcp, not https://host/).") def _reconnect_or_reraise_group(self, eg: BaseExceptionGroup) -> str: """Map an SDK transport TaskGroup failure to a clean ``"reconnect"``. HTTP/SSE stream @@ -422,8 +412,8 @@ 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 we cannot hook redirects, so the cross-origin + # header boundary cannot be enforced. raise ImportError(f"MCP server '{self.name}' requires mcp >= 1.24.0 to " "enforce the portable redirect-header boundary " "(strict_redirect_headers). Upgrade the mcp package.") @@ -441,9 +431,9 @@ 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. + # Portable Agent Plugins v1 (strict_redirect_headers): configured headers MUST NOT follow + # a redirect to a different origin. Capture the configured names BEFORE client-generated + # headers are merged in. strict_cfg_headers = bool(config.get("strict_redirect_headers")) configured_header_names = {key.lower() for key in headers} # Optional per-user identity header; explicit headers of the same name win. @@ -458,7 +448,6 @@ class MCPServerTransportMixin: ssl_verify = config.get("ssl_verify", True) client_cert = _resolve_client_cert(self.name, config) oauth_auth = self._build_oauth_auth(url, config) - if config.get("transport") == "sse": transport = self._sse_transport(url, headers, connect_timeout, ssl_verify, client_cert, oauth_auth, strict_cfg_headers) label = "SSE" @@ -504,8 +493,8 @@ 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: drop its + # stale connect error from status surfaces. with _core._lock: if _core._servers.get(self.name) is self: _core._server_connect_errors.pop(self.name, None)