From c1f8af1e86bdced8ede9a6c09e0fffd2dfe100bb Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 13:55:31 -0700 Subject: [PATCH] refactor(tools/mcp): split mcp_tool.py into transport/lifecycle/schema/handlers/... sibling modules; compact watchdog and schema cache --- tests/tools/test_mcp_schema_cache.py | 10 +- tests/tools/test_mcp_stdio_children_dead.py | 4 +- tools/mcp_schema_cache.py | 21 +- tools/mcp_stdio_watchdog.py | 68 +- tools/mcp_tool.py | 7842 ++----------------- tools/mcp_tool_agent.py | 291 + tools/mcp_tool_common.py | 181 + tools/mcp_tool_config.py | 382 + tools/mcp_tool_content.py | 254 + tools/mcp_tool_errors.py | 541 ++ tools/mcp_tool_handlers.py | 785 ++ tools/mcp_tool_health.py | 392 + tools/mcp_tool_lifecycle.py | 342 + tools/mcp_tool_registration.py | 537 ++ tools/mcp_tool_sampling.py | 521 ++ tools/mcp_tool_schema.py | 307 + tools/mcp_tool_transport.py | 676 ++ 17 files changed, 5832 insertions(+), 7322 deletions(-) create mode 100644 tools/mcp_tool_agent.py create mode 100644 tools/mcp_tool_common.py create mode 100644 tools/mcp_tool_config.py create mode 100644 tools/mcp_tool_content.py create mode 100644 tools/mcp_tool_errors.py create mode 100644 tools/mcp_tool_handlers.py create mode 100644 tools/mcp_tool_health.py create mode 100644 tools/mcp_tool_lifecycle.py create mode 100644 tools/mcp_tool_registration.py create mode 100644 tools/mcp_tool_sampling.py create mode 100644 tools/mcp_tool_schema.py create mode 100644 tools/mcp_tool_transport.py diff --git a/tests/tools/test_mcp_schema_cache.py b/tests/tools/test_mcp_schema_cache.py index cc6df9a6d2..ebfd4bea75 100644 --- a/tests/tools/test_mcp_schema_cache.py +++ b/tests/tools/test_mcp_schema_cache.py @@ -43,31 +43,23 @@ class TestCacheRoundTrip: assert entry is not None assert msc.tools_from_cache_entry(entry) == tools assert msc.utility_tools_from_cache_entry(entry) == [] - assert msc.has_cached_entry("srv", "fp1") def test_fingerprint_mismatch_returns_none(self, monkeypatch, tmp_path): self._isolate(monkeypatch, tmp_path) msc.write_cache_entry("srv", "fp1", tools=[], utility_tools=[]) assert msc.get_cached_entry("srv", "OTHER") is None - assert not msc.has_cached_entry("srv", "OTHER") def test_missing_server_returns_none(self, monkeypatch, tmp_path): self._isolate(monkeypatch, tmp_path) assert msc.get_cached_entry("nope", "fp") is None - def test_clear_cache_entry(self, monkeypatch, tmp_path): - self._isolate(monkeypatch, tmp_path) - msc.write_cache_entry("srv", "fp1", tools=[], utility_tools=[]) - msc.clear_cache_entry("srv") - assert msc.get_cached_entry("srv", "fp1") is None - def test_corrupt_cache_file_is_tolerated(self, monkeypatch, tmp_path): self._isolate(monkeypatch, tmp_path) (tmp_path / "cache.json").write_text("{not json", encoding="utf-8") assert msc.get_cached_entry("srv", "fp") is None # And writes recover the file. msc.write_cache_entry("srv", "fp", tools=[], utility_tools=[]) - assert msc.has_cached_entry("srv", "fp") + assert msc.get_cached_entry("srv", "fp") is not None def test_malformed_entry_shapes_are_tolerated(self): assert msc.tools_from_cache_entry({"tools": "nope"}) == [] diff --git a/tests/tools/test_mcp_stdio_children_dead.py b/tests/tools/test_mcp_stdio_children_dead.py index 23ec4a8077..ef54ed4663 100644 --- a/tests/tools/test_mcp_stdio_children_dead.py +++ b/tests/tools/test_mcp_stdio_children_dead.py @@ -105,9 +105,9 @@ def test_watch_ok_probe_does_not_create_unawaited_coroutine(): """ import inspect as _inspect - import tools.mcp_tool as mcp_mod + import tools.mcp_tool_handlers as handlers_mod - src = _inspect.getsource(mcp_mod) + src = _inspect.getsource(handlers_mod) assert "isawaitable(_watch_children())" not in src assert "iscoroutinefunction(_watch_children)" in src diff --git a/tools/mcp_schema_cache.py b/tools/mcp_schema_cache.py index 48bb5ae665..0e3377bd9a 100644 --- a/tools/mcp_schema_cache.py +++ b/tools/mcp_schema_cache.py @@ -83,16 +83,15 @@ def get_cached_entry(server_name: str, fingerprint: str) -> Optional[dict]: return None ttl_ms = entry.get("ttl_ms") written_at = entry.get("written_at") - if isinstance(ttl_ms, (int, float)) and isinstance(written_at, (int, float)): - if (time.time() - written_at) * 1000.0 >= float(ttl_ms): - return None + if ( + isinstance(ttl_ms, (int, float)) + and isinstance(written_at, (int, float)) + and (time.time() - written_at) * 1000.0 >= float(ttl_ms) + ): + return None return entry -def has_cached_entry(server_name: str, fingerprint: str) -> bool: - return get_cached_entry(server_name, fingerprint) is not None - - def write_cache_entry( server_name: str, fingerprint: str, @@ -132,14 +131,6 @@ def write_cache_entry( _save_all(data) -def clear_cache_entry(server_name: str) -> None: - with _cache_lock: - data = _load_all() - if server_name in data: - del data[server_name] - _save_all(data) - - def tools_from_cache_entry(entry: dict) -> List[dict]: """Return cached MCP tool dicts (name, description, inputSchema).""" tools = entry.get("tools") diff --git a/tools/mcp_stdio_watchdog.py b/tools/mcp_stdio_watchdog.py index a39f36d6fe..7c4c1896b7 100644 --- a/tools/mcp_stdio_watchdog.py +++ b/tools/mcp_stdio_watchdog.py @@ -1,38 +1,19 @@ #!/usr/bin/env python3 """Parent-death watchdog supervisor for stdio MCP subprocesses. -Problem this fixes (#TBD): a stdio MCP server (e.g. ``npx -y mcp-remote -``) is spawned as a direct child of the Hermes process. Hermes's own -teardown path (``MCPServerTask.shutdown()`` / ``_kill_orphaned_mcp_children`` -at final exit) reaps it cleanly on a *graceful* exit. But if the spawning -Hermes process dies hard — ``kill -9``, an OS-level crash, a force-quit of -the TUI/desktop app — that teardown code never runs, and the child (plus any -of its own descendants, e.g. mcp-remote's spawned ``node`` process) is -orphaned. macOS has no direct equivalent of Linux's -``prctl(PR_SET_PDEATHSIG)`` to make the kernel auto-kill a child when its -parent dies, so nothing reaps these until the next Hermes startup's opt-in -``_kill_orphaned_mcp_children()`` sweep — which only runs if something calls -it. Repeated ungraceful session restarts can pile up N orphaned processes, -all racing to hold the same upstream SSE session, producing errors like -"Invalid request parameters" / "Received request before initialization was -complete" on the *legitimate* new connection. +If Hermes dies hard (kill -9, crash, force-quit), its graceful teardown never +runs and the stdio child plus its descendants are orphaned; macOS has no +``PR_SET_PDEATHSIG`` equivalent. Piled-up orphans then race the legitimate new +connection for the same upstream session. So the MCP command is spawned via +this supervisor, which: + 1. runs the real command as its own child in a new process group (so the + whole tree can be killpg'd); + 2. passes stdin/stdout/stderr straight through — the MCP stdio protocol + talks over those pipes, so this must be a no-op relay, not a proxy; + 3. polls ``getppid()`` against the recorded parent PID and, when the parent + is gone, SIGTERMs the child's group, waits, then SIGKILLs. -Fix: don't spawn the MCP server command directly. Spawn this supervisor -instead, which: - 1. execs the real command as its own child (own process group via - ``start_new_session``, so it doesn't inherit the supervisor's - controlling terminal weirdly and so we can killpg it cleanly); - 2. transparently passes stdin/stdout/stderr through — the MCP stdio - protocol talks directly over those pipes, so the supervisor must be a - no-op relay, not a bytes-in-the-middle proxy; - 3. runs a background thread that polls the direct POSIX parent identity: - compare current ``getppid()`` against the parent PID recorded when the - wrapper was created; - 4. the instant the original parent is gone, terminates the real child's - process group (SIGTERM, grace period, then SIGKILL) and exits. - -This is intentionally a thin, standard-library-only script so it starts fast -and can't itself become a resource leak. +Standard-library only so it starts fast and cannot itself leak. Usage (see ``tools/mcp_tool.py::_run_stdio``):: @@ -60,13 +41,9 @@ def _is_orphaned(original_ppid: int, getppid=os.getppid) -> bool: def _terminate_process_group(proc: subprocess.Popen) -> None: - """Best-effort SIGTERM-then-SIGKILL of the child's process group. - - This module only ever runs on POSIX (the wrap site in tools/mcp_tool.py - gates on ``os.name == "posix"``), but guard the POSIX-only primitives - anyway so an accidental Windows import/execute degrades to a plain - child kill instead of AttributeError. - """ + """Best-effort SIGTERM-then-SIGKILL of the child's process group. Only runs + on POSIX in practice, but guards the POSIX-only primitives so an accidental + Windows run degrades to a plain child kill instead of AttributeError.""" killpg = getattr(os, "killpg", None) if killpg is None: # windows-footgun: ok — non-POSIX fallback try: @@ -115,9 +92,8 @@ def main(argv: list[str] | None = None) -> int: print("mcp_stdio_watchdog: no command given after '--'", file=sys.stderr) return 2 - # New process group so we can killpg() the whole tree the real command - # may spawn (e.g. mcp-remote's own child `node` process), without - # touching our own group or the (already-gone) original parent's. + # New process group so we can killpg() the whole tree the real command may + # spawn, without touching our own group or the original parent's. proc = subprocess.Popen( real_argv, stdin=sys.stdin, @@ -126,12 +102,10 @@ def main(argv: list[str] | None = None) -> int: start_new_session=True, ) - # Because the real server lives in its OWN process group (above), the - # parent's graceful-shutdown killpg of *our* group no longer reaches it. - # Forward SIGTERM/SIGINT to the child's group so graceful teardown - # (`_kill_orphaned_mcp_children`, shutdown sweeps) still kills a wedged - # server that ignores stdin EOF — otherwise the watchdog wrap would - # invert the bug it fixes. + # The server lives in its OWN group, so the parent's shutdown killpg of + # *our* group no longer reaches it. Forward SIGTERM/SIGINT to the child's + # group so graceful teardown still kills a wedged server that ignores stdin + # EOF — otherwise the wrap would invert the bug it fixes. def _forward_shutdown(signum, frame): # noqa: ARG001 _terminate_process_group(proc) sys.exit(128 + signum) diff --git a/tools/mcp_tool.py b/tools/mcp_tool.py index 01e7d9b3db..5c7f5d61ff 100644 --- a/tools/mcp_tool.py +++ b/tools/mcp_tool.py @@ -1,246 +1,225 @@ #!/usr/bin/env python3 """ -MCP (Model Context Protocol) Client Support +MCP (Model Context Protocol) client: connects to configured MCP servers over +stdio, Streamable HTTP or SSE, discovers their tools and registers them into +the hermes tool registry. The ``mcp`` package is optional; without it this +module is a no-op. -Connects to external MCP servers via stdio, HTTP/StreamableHTTP, or SSE -transport, discovers their tools, and registers them into the hermes-agent -tool registry so the agent can call them like any built-in tool. - -Configuration is read from ~/.hermes/config.yaml under the ``mcp_servers`` key. -The ``mcp`` Python package is optional -- if not installed, this module is a -no-op and logs a debug message. - -Example config:: +Config lives under ``mcp_servers`` in ~/.hermes/config.yaml:: mcp_servers: filesystem: command: "npx" args: ["-y", "@modelcontextprotocol/server-filesystem", "/tmp"] env: {} - timeout: 120 # per-tool-call timeout in seconds (default: 300) - connect_timeout: 60 # initial connection timeout (default: 60) - keepalive_interval: 10 # liveness ping cadence in seconds (default: - # 180). Set below the server's session TTL for - # servers that GC idle sessions quickly (e.g. - # Unreal Engine editor MCP, ~15s). Floored at 5s. - idle_timeout_seconds: 3600 # optional stdio recycle after idle - max_lifetime_seconds: 86400 # optional stdio recycle after age - # The recycle settings may also live under lifecycle: {...}. - # Use 0 to disable either recycle limit. - github: - command: "npx" - args: ["-y", "@modelcontextprotocol/server-github"] - env: - GITHUB_PERSONAL_ACCESS_TOKEN: "ghp_..." - supports_parallel_tool_calls: true # tools from this server may run concurrently + timeout: 120 # per tool call (default 300) + connect_timeout: 60 # initial connect (default 60) + keepalive_interval: 10 # liveness ping; keep below the server's + # session TTL (default 180, floor 5) + idle_timeout_seconds: 3600 # optional stdio recycle (0 = off); may + max_lifetime_seconds: 86400 # also live under lifecycle: {...} + supports_parallel_tool_calls: true remote_api: url: "https://my-mcp-server.example.com/mcp" - headers: - Authorization: "Bearer sk-..." - identity_header: # optional per-user identity header attached - name: "X-User-Id" # to this server's HTTP/SSE requests - value_from: "static" # "static" (default) or "profile" - value: "alice" # required for static; profile mode uses the - # active Hermes profile name - timeout: 180 - skip_preflight: true # bypass the content-type probe for a valid - # Streamable HTTP endpoint that answers HEAD/GET - # with a non-MCP content type but serves real - # MCP over POST. Default: false. + headers: {Authorization: "Bearer sk-..."} + identity_header: {name: "X-User-Id", value_from: "static", value: "alice"} + skip_preflight: true # endpoint answers HEAD/GET with a non-MCP + # content type but serves MCP over POST searxng: url: "http://localhost:8000/sse" - transport: sse # use SSE transport instead of Streamable HTTP - timeout: 180 - connect_timeout: 10 - command: "npx" - args: ["-y", "analysis-server"] - sampling: # server-initiated LLM requests - enabled: true # default: true - model: "gemini-3-flash" # override model (optional) - max_tokens_cap: 4096 # max tokens per request - timeout: 30 # LLM call timeout (seconds) - max_rpm: 10 # max requests per minute - allowed_models: [] # model whitelist (empty = all) - max_tool_rounds: 5 # tool loop limit (0 = disable) - log_level: "info" # audit verbosity + transport: sse + sampling: {enabled: true, model: "gemini-3-flash", max_tokens_cap: 4096, + timeout: 30, max_rpm: 10, allowed_models: [], max_tool_rounds: 5} -Features: - - Stdio transport (command + args) and HTTP/StreamableHTTP transport (url) - - SSE transport (transport: sse) for MCP servers using the SSE protocol - - Automatic reconnection with exponential backoff (up to 5 retries) - - Environment variable filtering for stdio subprocesses (security) - - Credential stripping in error messages returned to the LLM - - Configurable per-server timeouts for tool calls and connections - - Thread-safe architecture with dedicated background event loop - - Sampling support: MCP servers can request LLM completions via - sampling/createMessage (text and tool-use responses) - - Parallel tool call opt-in: per-server ``supports_parallel_tool_calls`` - flag allows concurrent execution of tools from the same server +Architecture: one background event loop (``_mcp_loop``) in a daemon thread; +each server is a long-lived Task on it (``MCPServerTask``) so the transport's +anyio cancel scopes are entered and exited in the same Task. Tool calls are +scheduled onto the loop via ``run_coroutine_threadsafe``. ``_servers`` and the +loop handles are shared with caller threads; every mutation holds ``_lock``. -Architecture: - A dedicated background event loop (_mcp_loop) runs in a daemon thread. - Each MCP server runs as a long-lived asyncio Task on this loop, keeping - its transport context alive. Tool call coroutines are scheduled onto the - loop via ``run_coroutine_threadsafe()``. - - On shutdown, each server Task is signalled to exit its ``async with`` - block, ensuring the anyio cancel-scope cleanup happens in the *same* - Task that opened the connection (required by anyio). - -Thread safety: - _servers and _mcp_loop/_mcp_thread are accessed from both the MCP - background thread and caller threads. All mutations are protected by - _lock so the code is safe regardless of GIL presence (e.g. Python 3.13+ - free-threading). +Module map (all names re-exported here): ``mcp_tool_common`` (pure helpers), +``mcp_tool_schema`` (schema conversion / naming), ``mcp_tool_content`` (result +block rendering), ``mcp_tool_errors`` (failure classification, URL/cert/header +resolution), ``mcp_tool_config`` (config loading, stdio env), ``mcp_tool_sampling`` +(sampling + elicitation handlers), ``mcp_tool_handlers`` (registry handlers and +per-call recovery), ``mcp_tool_registration`` (registry writes), ``mcp_tool_transport`` +/ ``mcp_tool_health`` (MCPServerTask mixins), ``mcp_tool_lifecycle`` (shutdown, +orphan reaping), ``mcp_tool_agent`` (live-agent tool list refresh). """ import asyncio import contextvars import concurrent.futures import errno -import fnmatch import inspect -import json import logging -import math import os -import random -import re -import shutil -import sys +import shutil # noqa: F401 — tests patch ``tools.mcp_tool.shutil.which`` import threading import time -from contextlib import asynccontextmanager -from types import SimpleNamespace -from typing import Callable -from datetime import datetime -from typing import Any, Coroutine, Dict, List, Optional, Set, Tuple -from urllib.parse import urlparse - -from tools.registry import tool_error -from tools.ansi_strip import strip_unicode_tags +from typing import Any, Callable, Coroutine, Dict, List, Optional, Set logger = logging.getLogger(__name__) - -# Hard allocation ceiling for a single MCP text payload (chars). This is the -# FIRST line of defense against a buggy or malicious MCP server returning -# multi-megabyte text: without it the full payload is allocated, JSON-encoded -# and handed downstream before the budget/spillover layer ever sees it -# (#56059). It deliberately sits far ABOVE the budget layer's 50K MCP -# spillover threshold (tools/budget_config.py) so ordinary large results -# reach spillover INTACT — spilled to disk in full, preview in context — -# while only pathological multi-MB floods are lossy-truncated here. -# -# Distilled from #56060 (Stoltemberg), #56072 (AlexFucuson9) and #56511 -# (Tranquil-Flow), which capped at get_max_bytes() (50K) — correct -# protection, but at that level it would truncate before spillover could -# preserve the data. The 40% head / 60% tail split is #56511's shape. -_MCP_HARD_RESULT_CAP_CHARS = 2_000_000 +# Split modules. Every name is re-exported here so ``from tools.mcp_tool import X`` +# and ``mock.patch("tools.mcp_tool.X")`` keep working; the siblings read origin +# state back through ``tools.mcp_tool`` at call time (never by value). +from tools.mcp_tool_common import ( # noqa: F401 + _BACKOFF_JITTER, + _CREDENTIAL_PATTERN, + _DEFAULT_TOOL_TIMEOUT, + _MISSING, + _env_ref_name, + _exc_str, + _get_lifecycle_seconds, + _jittered, + _parse_boolish, + _prepend_path, + _resolve_tool_timeout, + _safe_numeric, + _sanitize_error, + mcp_field, +) +from tools.mcp_tool_schema import ( # noqa: F401 + MCP_TOOL_NAME_PREFIX, + _MCP_INJECTION_PATTERNS, + _MCP_NAME_DELIM, + _UTILITY_CAPABILITY_ATTRS, + _UTILITY_CAPABILITY_METHODS, + _build_utility_schemas, + _convert_mcp_schema, + _normalize_mcp_input_schema, + _normalize_name_filter, + _scan_mcp_description, + matches_name_filter, + mcp_prefixed_tool_name, + sanitize_mcp_name_component, +) +from tools.mcp_tool_content import ( # noqa: F401 + _MCP_HARD_RESULT_CAP_CHARS, + _MCP_RESOURCE_MAX_B64_CHARS, + _MCP_RESOURCE_MAX_BYTES, + _cache_mcp_audio_block, + _cache_mcp_image_block, + _is_reserved_mcp_meta_key, + _mcp_image_extension_for_mime_type, + _mcp_resource_filename, + _render_mcp_resource_block, + _strip_reserved_meta_keys, + _truncate_mcp_text_result, +) +from tools.mcp_tool_errors import ( # noqa: F401 + InvalidMcpUrlError, + NonMcpEndpointError, + _AUTH_ERROR_TYPES, + _EXC_TRAVERSAL_MAX_NODES, + _HTTP_STATUS_ERROR_TYPES, + _JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, + _SESSION_EXPIRED_MARKERS, + _apply_identity_header, + _classify_mcp_failure, + _contains_only_cancellation, + _format_connect_error, + _get_auth_error_types, + _handshake_rejected_as_modern, + _http_status_error_types, + _is_auth_error, + _is_method_not_found_error, + _is_session_expired_error, + _make_redirect_header_stripper, + _resolve_client_cert, + _resolve_identity_header, + _unwrap_exception_group, + _validate_remote_mcp_url, +) +from tools.mcp_tool_config import ( # noqa: F401 + _ENV_VAR_PATTERN, + _SAFE_ENV_KEYS, + _SAFE_ENV_KEYS_CASE_INSENSITIVE, + _build_safe_env, + _context_var_value, + _filter_suspicious_mcp_servers, + _get_mcp_stderr_log, + _interpolate_env_vars, + _load_mcp_config, + _mcp_stderr_log_fh, + _mcp_stderr_log_lock, + _resolve_stdio_command, + _warn_hidden_whitespace, + _whitespace_warned, + _workspace_folder, + _wrap_command_with_watchdog, + _write_stderr_log_header, +) +from tools.mcp_tool_sampling import ( # noqa: F401 + ElicitationHandler, + SamplingHandler, + _format_elicitation_schema_summary, +) +from tools.mcp_tool_handlers import ( # noqa: F401 + _StdioChildExited, + _handle_auth_error_and_retry, + _handle_session_expired_and_retry, + _handle_stdio_child_exited_and_retry, + _interrupted_call_result, + _make_check_fn, + _make_get_prompt_handler, + _make_list_prompts_handler, + _make_list_resources_handler, + _make_read_resource_handler, + _make_tool_handler, + _mark_server_call_started, + _track_inflight_rpc, + _trust_gate_check, +) +from tools.mcp_tool_registration import ( # noqa: F401 + _CachedMCPTool, + _annotation_read_only_hint, + _existing_tool_names, + _forget_mcp_tool_server, + _normalize_server_trust, + _record_tool_trust_metadata, + _register_from_cache_sync, + _register_server_tools, + _select_utility_schemas, + _track_mcp_tool_server, +) +from tools.mcp_tool_lifecycle import ( # noqa: F401 + _NON_MCP_CHILD_CMDLINE_MARKERS, + _drain_and_stop_mcp_loop, + _drain_mcp_loop_tasks, + _filter_mcp_children, + _kill_orphaned_mcp_children, + _orphan_stdio_pid_servers, + _orphan_stdio_pids, + _snapshot_child_pids, + _stdio_pgids, + _stdio_pids, + _stop_mcp_loop_if_idle, + shutdown_mcp_servers, +) +from tools.mcp_tool_agent import ( # noqa: F401 + _agent_tools_lock, + _merge_preserving_prefix, + _reinject_post_build_tools, + persist_agent_tool_names, + refresh_agent_mcp_tools, + reprobe_tool_availability, + restore_agent_tool_prefix, +) +from tools.mcp_tool_transport import MCPServerTransportMixin +from tools.mcp_tool_health import MCPServerHealthMixin -def _truncate_mcp_text_result(text: str, max_chars: int = _MCP_HARD_RESULT_CAP_CHARS) -> str: - """Bound pathological MCP text before it propagates (#56059). - - Results at or under ``max_chars`` pass through unchanged; oversized text - keeps a 40% head / 60% tail split with an omission notice in between. - """ - if len(text) <= max_chars: - return text - head_chars = int(max_chars * 0.4) - tail_chars = max_chars - head_chars - omitted = len(text) - head_chars - tail_chars - return ( - text[:head_chars] - + f"\n\n... [MCP RESULT TRUNCATED - {omitted:,} chars omitted " - f"out of {len(text):,} total] ...\n\n" - + text[-tail_chars:] - ) - -# Upper bound for the OSV malware preflight during stdio MCP startup. The -# check makes a blocking urllib HTTPS call whose own timeout can fail to -# interrupt a stalled SSL handshake, which froze the asyncio event loop and -# blew past the gateway's 15s startup budget (#29184). We run it off the loop -# AND bound it here; the check is fail-open, so a timeout lets startup proceed. -# Set just ABOVE osv_check._TIMEOUT (10s) so the inner socket timeout fires -# first in the normal case; this outer bound only bites when a stalled SSL -# handshake defeats the inner timeout (the #29184 failure mode). +# Wall-clock bound on the (fail-open) OSV malware preflight run off the loop +# before a stdio spawn. Kept just ABOVE osv_check._TIMEOUT (10s) so the inner +# socket timeout normally fires first; this only bites when a stalled SSL +# handshake defeats it (which used to freeze the event loop at startup). _OSV_MALWARE_CHECK_TIMEOUT_S = 12.0 # --------------------------------------------------------------------------- -# Stdio subprocess stderr redirection -# --------------------------------------------------------------------------- -# -# The MCP SDK's ``stdio_client(server, errlog=sys.stderr)`` defaults the -# subprocess stderr stream to the parent process's real stderr, i.e. the -# user's TTY. That means any MCP server we spawn at startup (FastMCP -# banners, slack-mcp-server JSON startup logs, etc.) writes directly onto -# the terminal while prompt_toolkit / Rich is rendering the TUI — which -# corrupts the display and can hang the session. -# -# Instead we redirect every stdio MCP subprocess's stderr into a shared -# per-profile log file (~/.hermes/logs/mcp-stderr.log), tagged with the -# server name so individual servers remain debuggable. -# -# Fallback is os.devnull if opening the log file fails for any reason. - -_mcp_stderr_log_fh: Optional[Any] = None -_mcp_stderr_log_lock = threading.Lock() - - -def _get_mcp_stderr_log() -> Any: - """Return a shared append-mode file handle for MCP subprocess stderr. - - Opened once per process and reused for every stdio server. Must have a - real OS-level file descriptor (``fileno()``) because asyncio's subprocess - machinery wires the child's stderr directly to that fd. Falls back to - ``/dev/null`` if opening the log file fails. - """ - global _mcp_stderr_log_fh - with _mcp_stderr_log_lock: - if _mcp_stderr_log_fh is not None: - return _mcp_stderr_log_fh - try: - from hermes_constants import get_hermes_home - log_dir = get_hermes_home() / "logs" - log_dir.mkdir(parents=True, exist_ok=True) - log_path = log_dir / "mcp-stderr.log" - # Line-buffered so server output lands on disk promptly; errors= - # "replace" tolerates garbled binary output from misbehaving - # servers. - fh = open(log_path, "a", encoding="utf-8", errors="replace", buffering=1) - # Sanity-check: confirm a real fd is available before we commit. - fh.fileno() - _mcp_stderr_log_fh = fh - except Exception as exc: # pragma: no cover — best-effort fallback - logger.debug("Failed to open MCP stderr log, using devnull: %s", exc) - try: - _mcp_stderr_log_fh = open(os.devnull, "w", encoding="utf-8") - except Exception: - # Last resort: the real stderr. Not ideal for TUI users but - # it matches pre-fix behavior. - _mcp_stderr_log_fh = sys.stderr - return _mcp_stderr_log_fh - - -def _write_stderr_log_header(server_name: str) -> None: - """Write a human-readable session marker before launching a server. - - Gives operators a way to find each server's output in the shared - ``mcp-stderr.log`` file without needing per-line prefixes (which would - require a pipe + reader thread and complicate shutdown). - """ - fh = _get_mcp_stderr_log() - try: - ts = datetime.now().strftime("%Y-%m-%d %H:%M:%S") - fh.write(f"\n===== [{ts}] starting MCP server '{server_name}' =====\n") - fh.flush() - except Exception: - pass - -# --------------------------------------------------------------------------- -# Graceful import -- MCP SDK is an optional dependency +# Optional MCP SDK: availability probe now, symbol import on first use # --------------------------------------------------------------------------- _MCP_AVAILABLE = False @@ -254,25 +233,19 @@ _MCP_MESSAGE_HANDLER_SUPPORTED = False _MCP_LOGGING_CALLBACK_SUPPORTED = False _MCP_NEW_HTTP = False sse_client = None -# Conservative fallback for SDK builds that don't export LATEST_PROTOCOL_VERSION. -# Streamable HTTP was introduced by 2025-03-26, so this remains valid for the -# HTTP transport path even on older-but-supported SDK versions. +# Fallback for SDKs that don't export LATEST_PROTOCOL_VERSION (Streamable HTTP +# arrived with 2025-03-26, so this stays valid for the HTTP path). LATEST_PROTOCOL_VERSION = "2025-03-26" -# The newest revision reachable through `ClientSession.initialize()`, which is -# NOT the newest revision the SDK knows about: from 2026-07-28 onward the -# handshake is replaced by a per-request envelope, so `initialize()` keeps -# sending `LATEST_HANDSHAKE_VERSION`. Seeding the MCP-Protocol-Version header -# from LATEST_PROTOCOL_VERSION would advertise a revision the body does not -# speak. Defaults to the handshake fallback for SDKs predating the split. +# Newest revision `ClientSession.initialize()` actually speaks. From 2026-07-28 +# the handshake is replaced by a per-request envelope, so this can be OLDER +# than LATEST_PROTOCOL_VERSION; the MCP-Protocol-Version header must be seeded +# from this one or it advertises a revision the body does not speak. LATEST_HANDSHAKE_VERSION = LATEST_PROTOCOL_VERSION -# The heavy SDK import is LAZY (see _ensure_mcp_sdk): importing `mcp` costs -# ~260ms (mcp.types alone is ~60ms of pydantic model construction), which used -# to be paid at tool-discovery time on EVERY CLI startup even with zero MCP -# servers configured. Availability is decided here with a metadata-only -# find_spec probe (~1ms, no module execution) so every existing -# `if not _MCP_AVAILABLE` gate, test patch, and skipif keeps its exact -# semantics; the symbol import itself happens on first real SDK use. +# Importing `mcp` costs ~260ms, so it is deferred to first real use +# (_ensure_mcp_sdk). Availability is decided here with a metadata-only +# find_spec probe so every `if not _MCP_AVAILABLE` gate / test patch / skipif +# keeps its exact semantics. try: import importlib.util as _importlib_util _MCP_AVAILABLE = _importlib_util.find_spec("mcp") is not None @@ -285,11 +258,9 @@ ClientSession: Any = None _MCP_SDK_IMPORT_ATTEMPTED = False _MCP_SDK_IMPORT_LOCK = threading.Lock() -# SDK symbols that _ensure_mcp_sdk() binds on first use. Module-level -# __getattr__ (PEP 562) below resolves external access to any of these by -# importing the SDK first — so tests doing mock.patch("tools.mcp_tool. -# stdio_client", ...) trigger the import when patch() saves the original, -# and the subsequent mock is never clobbered (_ensure is idempotent). +# SDK symbols bound by _ensure_mcp_sdk(). Module __getattr__ (PEP 562) imports +# the SDK on first external access, so mock.patch("tools.mcp_tool.stdio_client") +# sees a real original and the mock is never clobbered (_ensure is idempotent). _MCP_SDK_LAZY_SYMBOLS = frozenset({ "StdioServerParameters", "stdio_client", "streamablehttp_client", "streamable_http_client", @@ -312,13 +283,11 @@ def __getattr__(name: str): def _ensure_mcp_sdk() -> bool: - """Import the optional ``mcp`` SDK on first use. Returns availability. + """Import the optional ``mcp`` SDK on first use; return availability. - Idempotent and thread-safe. Sets the module-level ``_MCP_*`` flags and - SDK symbol globals exactly as the old import-time block did. Honors a - test-patched ``_MCP_AVAILABLE=False`` (returns False without importing) - and test-installed mock symbols (``ClientSession`` already set → no - re-import, so mocks are never clobbered). + Idempotent and thread-safe. Honors a test-patched ``_MCP_AVAILABLE=False`` + (no import) and pre-installed mock symbols (``ClientSession`` already set + means no re-import, so mocks are never clobbered). """ global _MCP_SDK_IMPORT_ATTEMPTED, _MCP_AVAILABLE, _MCP_HTTP_AVAILABLE global _MCP_SAMPLING_TYPES, _MCP_NOTIFICATION_TYPES, _MCP_ELICITATION_TYPES @@ -343,8 +312,8 @@ def _ensure_mcp_sdk() -> bool: from mcp import ClientSession, StdioServerParameters from mcp.client.stdio import stdio_client _MCP_AVAILABLE = True - # Prefer the non-deprecated API (mcp >= 1.24.0); fall back to the - # deprecated wrapper for older SDK versions. + # mcp >= 1.24 ships streamable_http_client; 2.0 dropped the + # deprecated streamablehttp_client alias. Either one gives HTTP. try: from mcp.client.streamable_http import streamable_http_client _MCP_NEW_HTTP = True @@ -355,18 +324,6 @@ def _ensure_mcp_sdk() -> bool: _MCP_LEGACY_HTTP = True except ImportError: _MCP_LEGACY_HTTP = False - # HTTP support requires EITHER entry point. mcp 2.0 dropped the - # deprecated `streamablehttp_client` alias, so gating on that name - # alone made _run_http raise ImportError for every HTTP and SSE - # server on 2.x before it could reach the `streamable_http_client` - # path. - # - # Reaching it was necessary and not sufficient: that path also - # unpacked the transport as a fixed 3-tuple, which is 1.x's shape. - # On 2.x it raised "not enough values to unpack (expected 3, got - # 2)" and every HTTP/SSE server parked after its retry ladder. - # Only stdio servers kept working, which is why this survived - # review - the common configs are all stdio. _MCP_HTTP_AVAILABLE = _MCP_NEW_HTTP or _MCP_LEGACY_HTTP try: from mcp.types import LATEST_PROTOCOL_VERSION @@ -375,17 +332,15 @@ def _ensure_mcp_sdk() -> bool: try: from mcp.client.session import LATEST_HANDSHAKE_VERSION except ImportError: - # Pre-2.x SDKs make no distinction: the newest revision IS the - # newest handshake revision, so the header and the body agree - # either way. + # Pre-2.x SDKs: newest revision IS the handshake revision. LATEST_HANDSHAKE_VERSION = LATEST_PROTOCOL_VERSION - # SSE transport client (for MCP servers using SSE transport instead of Streamable HTTP) try: from mcp.client.sse import sse_client except ImportError: sse_client = None logger.debug("mcp.client.sse.sse_client not available -- SSE transport disabled") - # Sampling types -- separated so older SDK versions don't break MCP support + # Optional type families are gated separately so an older SDK + # only loses that feature, not MCP support. try: from mcp.types import ( CreateMessageResult, @@ -399,17 +354,11 @@ def _ensure_mcp_sdk() -> bool: _MCP_SAMPLING_TYPES = True except ImportError: logger.debug("MCP sampling types not available -- sampling disabled") - # Elicitation types -- gated separately for the same reason as sampling. - # Added in mcp Python SDK 1.11.0 (Jul 2025); servers use elicitation to - # ask the client for structured input mid-tool-call (e.g. payment - # authorization). Missing types just disable the feature; everything - # else keeps working. try: from mcp.types import ElicitRequestParams, ElicitResult _MCP_ELICITATION_TYPES = True except ImportError: logger.debug("MCP elicitation types not available -- elicitation disabled") - # Notification types for dynamic tool discovery (tools/list_changed) try: from mcp.types import ( ServerNotification, @@ -445,18 +394,12 @@ _SDK_HTTPX_MOD = None def sdk_httpx(): """Return the httpx module the *installed* MCP SDK is built against. - mcp 2.0 moved its HTTP transports and OAuth stack from ``httpx`` to - ``httpx2`` — a separate distribution with the same public API, importable - side by side with Hermes' own pinned ``httpx``. Every object that crosses - the SDK boundary has to come from the module the SDK itself imports: - the ``AsyncClient`` handed to ``streamable_http_client``, the client the - ``sse_client`` factory returns, the ``Request`` built by the SDK's OAuth - metadata helpers, and the exception classes those raise. Mixing the two - fails at the transport layer rather than at import, so resolve it from the - SDK's own transport module instead of inferring it from a version number. - - Returns ``None`` only when neither module is importable, which also means - the SDK import above failed and no caller here can run. + mcp 2.0 moved to ``httpx2`` (same API, separate distribution). Every + object crossing the SDK boundary — the ``AsyncClient`` passed to the + transport, OAuth ``Request`` objects, the exception classes — must come + from the module the SDK itself imports, or it fails at the transport + layer rather than at import. Resolved from the SDK's transport module, not + a version number. ``None`` only when neither module is importable. """ global _SDK_HTTPX_MOD if _SDK_HTTPX_MOD is not None: @@ -469,9 +412,7 @@ def sdk_httpx(): except ImportError: _SDK_HTTPX_MOD = None if _SDK_HTTPX_MOD is None: - # SDK transport module unavailable (or it stopped importing the - # module under a predictable name). Fall back to whichever is - # present, newest first. + # Transport module missing / renamed its import: newest present wins. try: import httpx2 as _fallback except ImportError: @@ -483,60 +424,29 @@ def sdk_httpx(): return _SDK_HTTPX_MOD -_MISSING = object() +def _client_session_accepts(kwarg: str) -> bool: + """Whether this SDK's ``ClientSession.__init__`` takes ``kwarg``. - -def mcp_field(obj, snake: str, camel: str, default=None): - """Read an MCP model field across the 1.x -> 2.x field rename. - - mcp 2.0 renamed every model field to snake_case and kept the camelCase - spelling only as a *serialization* alias — pydantic aliases do not apply - to attribute access, so ``getattr(result, "isError", False)`` returns the - default on 2.x rather than raising. That turns a rename into silent wrong - behaviour: failed tool calls read as successful, tool schemas read as - empty, paginated lists stop after page one. Asking for both spellings - keeps the read correct on either SDK generation, which matters because - ``mcp`` is an optional extra users can install at their own version. + Older SDKs lack ``message_handler`` (no list_changed notifications) and + ``logging_callback`` (server ``notifications/message`` silently dropped). """ - value = getattr(obj, snake, _MISSING) - if value is not _MISSING: - return value - value = getattr(obj, camel, _MISSING) - return default if value is _MISSING else value + if not _MCP_AVAILABLE: + return False + try: + return kwarg in inspect.signature(ClientSession).parameters + except (TypeError, ValueError): + return False def _check_message_handler_support() -> bool: - """Check if ClientSession accepts ``message_handler`` kwarg. - - Inspects the constructor signature for backward compatibility with older - MCP SDK versions that don't support notification handlers. - """ - if not _MCP_AVAILABLE: - return False - try: - return "message_handler" in inspect.signature(ClientSession).parameters - except (TypeError, ValueError): - return False + return _client_session_accepts("message_handler") def _check_logging_callback_support() -> bool: - """Check if ClientSession accepts the ``logging_callback`` kwarg. - - Mirrors ``_check_message_handler_support`` for backward compatibility - with older MCP SDK versions. Without a logging_callback, the SDK's - default handler silently discards every ``notifications/message`` a - server emits, so server-side diagnostics never reach Hermes' logs. - """ - if not _MCP_AVAILABLE: - return False - try: - return "logging_callback" in inspect.signature(ClientSession).parameters - except (TypeError, ValueError): - return False + return _client_session_accepts("logging_callback") # MCP logging levels (RFC 5424 syslog severities) -> Python logging levels. -# Port of anomalyco/opencode#34529's serverLog mapping. _MCP_LOG_LEVEL_MAP = { "debug": logging.DEBUG, "info": logging.INFO, @@ -549,416 +459,51 @@ _MCP_LOG_LEVEL_MAP = { } # --------------------------------------------------------------------------- -# Constants +# Reconnect / keepalive tuning # --------------------------------------------------------------------------- -_DEFAULT_TOOL_TIMEOUT = 300 # seconds for tool calls - - -def _resolve_tool_timeout(config: dict) -> float: - """Per-server tool-call timeout with unified-layer resolution (#85125 2g). - - Precedence: per-server ``mcp_servers..timeout`` (most specific, - always wins) > ``timeouts.mcp.tool_call`` in config.yaml > the historical - default. Values are platform-clamped by ``resolve_timeout`` either way. - Defaults are unchanged: with neither key set this returns 300, exactly - as before. - """ - per_server = config.get("timeout") - if per_server is not None: - return per_server - try: - from agent.deadline import resolve_timeout - - resolved = resolve_timeout("mcp.tool_call", default=_DEFAULT_TOOL_TIMEOUT) - if resolved is not None: - return resolved - except Exception: - logger.debug("mcp.tool_call timeout resolution failed", exc_info=True) - return _DEFAULT_TOOL_TIMEOUT - _DEFAULT_CONNECT_TIMEOUT = 60 # seconds for initial connection per server _MAX_RECONNECT_RETRIES = 5 _MAX_INITIAL_CONNECT_RETRIES = 3 # retries for the very first connection attempt _MAX_BACKOFF_SECONDS = 60 -# While parked (reconnect budget exhausted, tools deregistered) the run task -# wakes on this cadence and attempts one revival probe. Without it a parked -# server is unrevivable: its tools are out of the registry, so no tool call -# can ever reach the circuit-breaker half-open probe or _signal_reconnect. +# Parked servers (budget exhausted, tools deregistered) self-probe on this +# cadence: with no tools registered nothing else can ever revive them. _PARKED_RETRY_INTERVAL = 300 # seconds between parked self-probes _RECYCLED_RECONNECT_TIMEOUT = 15.0 -# How long a tool call waits for a respawned stdio child after its subprocess -# was found dead — a gateway restart kills every MCP stdio child, -# and the next call from a still-live session would otherwise fail for no real -# reason). Bounded: when the wait elapses the call reports the dead transport -# instead of looping, so a genuinely broken server still parks via the -# rapid-drop budget in run() rather than hot-cycling respawns. +# Bounded wait for a respawned stdio child when a call finds it dead (gateway +# restarts kill every MCP child). Bounded so a broken server still parks via +# run()'s rapid-drop budget instead of hot-cycling respawns. _STDIO_RESPAWN_WAIT_SEC = 15.0 -# Jitter applied to reconnect backoff sleeps. Without it, every server that -# lost the same backend retries in lockstep (thundering herd) and log lines -# from N servers land in synchronized bursts. -_BACKOFF_JITTER = 0.2 # +/-20% -def _jittered(seconds: float) -> float: - """Return ``seconds`` with +/-20% uniform jitter, floored at 0.""" - return max(0.0, seconds * random.uniform(1.0 - _BACKOFF_JITTER, - 1.0 + _BACKOFF_JITTER)) - -# Keepalive cadence for HTTP/SSE sessions. The MCP spec lets a server expire -# idle sessions on any TTL it chooses (Streamable HTTP "Session Management"), -# so a client that wants a session to survive idle periods MUST refresh faster -# than that TTL. The default suits long LB/NAT idle windows (commonly -# 300-600s); servers with short session TTLs (e.g. Unreal Engine's editor MCP, -# ~15s) need a smaller ``keepalive_interval`` in their config or every idle -# tool call lands on a dead session and pays the full reconnect path. The floor -# stops a misconfigured tiny interval from busy-looping the keepalive. +# Servers may expire idle sessions on any TTL, so the client MUST ping faster +# than that TTL; servers with short TTLs (~15s) need a smaller configured +# ``keepalive_interval``. The floor stops a tiny interval from busy-looping. _DEFAULT_KEEPALIVE_INTERVAL = 180 # seconds between liveness pings _MIN_KEEPALIVE_INTERVAL = 5 # clamp floor for configured intervals -# Final shutdown gives pending MCP-loop tasks one bounded cancellation cycle -# before closing their owning loop. Cooperative parked/reconnect waiters finish -# immediately; cancellation-resistant tasks must not hang process exit. +# One bounded cancellation cycle for pending loop tasks at final shutdown, so +# cancellation-resistant tasks cannot hang process exit. _MCP_LOOP_DRAIN_TIMEOUT = 3.0 -# Environment variables that are safe to pass to stdio subprocesses -_SAFE_ENV_KEYS = frozenset({ - "PATH", "HOME", "USER", "LANG", "LC_ALL", "TERM", "SHELL", "TMPDIR", -}) - -_SAFE_ENV_KEYS_CASE_INSENSITIVE = frozenset({ - # Windows process/location vars. These are needed by launcher-style tools - # such as Docker Desktop's MCP plugin discovery, and do not carry secrets. - "ALLUSERSPROFILE", - "APPDATA", - "COMMONPROGRAMFILES", - "COMMONPROGRAMFILES(X86)", - "COMMONPROGRAMW6432", - "COMPUTERNAME", - "COMSPEC", - "HOMEDRIVE", - "HOMEPATH", - "LOCALAPPDATA", - "NUMBER_OF_PROCESSORS", - "OS", - "PATHEXT", - "PROCESSOR_ARCHITECTURE", - "PROGRAMDATA", - "PROGRAMFILES", - "PROGRAMFILES(X86)", - "PROGRAMW6432", - "PUBLIC", - "SYSTEMDRIVE", - "SYSTEMROOT", - "TEMP", - "TMP", - "USERDOMAIN", - "USERNAME", - "USERPROFILE", - "WINDIR", -}) - -# Regex for credential patterns to strip from error messages -_CREDENTIAL_PATTERN = re.compile( - r"(?:" - r"ghp_[A-Za-z0-9_]{1,255}" # GitHub PAT - r"|sk-[A-Za-z0-9_]{1,255}" # OpenAI-style key - r"|Bearer\s+\S+" # Bearer token - r"|token=[^\s&,;\"']{1,255}" # token=... - r"|key=[^\s&,;\"']{1,255}" # key=... - r"|API_KEY=[^\s&,;\"']{1,255}" # API_KEY=... - r"|password=[^\s&,;\"']{1,255}" # password=... - r"|secret=[^\s&,;\"']{1,255}" # secret=... - r")", - re.IGNORECASE, -) - -# Pre-compiled pattern for ${VAR_NAME} style env-var interpolation. -# Supports any non-} characters in the variable name (hyphens, dots, etc.) -# so providers like MY-VAR or my.var work correctly. -_ENV_VAR_PATTERN = re.compile(r"\$\{([^}]+)\}") - - -def _env_ref_name(ref: str) -> str: - """Normalize a ``${...}`` reference body into an env-var name. - - Accepts Cursor-style ``${env:VAR}`` in addition to plain ``${VAR}`` by - stripping a leading ``env:`` prefix. The result is the bare variable name - to look up in the secret scope / ``os.environ``. - """ - ref = ref.strip() - if ref.startswith("env:"): - ref = ref[len("env:"):].strip() - return ref - - -def _workspace_folder() -> str: - """Best-effort absolute workspace root for ``${workspaceFolder}``. - - Resolution order: - - 1. ``tools.file_tools._authoritative_workspace_root()`` — the session's - recorded terminal cwd, a registered task/session cwd override, or a - sentinel-free absolute ``$TERMINAL_CWD`` (in that order). - 2. ``os.getcwd()`` as the final fallback when no session anchor exists. - """ - try: - from tools.file_tools import _authoritative_workspace_root - - root = _authoritative_workspace_root() - if root: - return root - except Exception: - pass - return os.getcwd() - - -def _context_var_value(ref: str) -> Optional[str]: - """Resolve Cursor-style context variables in ``${...}`` references. - - Supports the case-sensitive names Cursor's ``mcp.json`` interpolation - understands beyond env vars: ``${userHome}``, ``${workspaceFolder}``, - ``${workspaceFolderBasename}``, ``${pathSeparator}`` and its ``${/}`` - shorthand. Returns ``None`` for anything else so unknown references keep - the existing env-var lookup semantics. - """ - if ref == "userHome": - return os.path.expanduser("~") - if ref == "workspaceFolder": - return _workspace_folder() - if ref == "workspaceFolderBasename": - root = _workspace_folder() - return os.path.basename(root.rstrip("/\\")) or root - if ref in ("pathSeparator", "/"): - return os.sep - return None - - -# --------------------------------------------------------------------------- -# Security helpers -# --------------------------------------------------------------------------- - -def _build_safe_env(user_env: Optional[dict]) -> dict: - """Build a filtered environment dict for stdio subprocesses. - - Only passes through safe baseline variables (PATH, HOME, etc.) and XDG_* - variables from the current process environment, secrets injected by an - external secret source (Bitwarden, 1Password, plugin backends) that - Hermes explicitly tagged during dotenv loading, plus any variables - explicitly specified by the user in the server config. - - This prevents accidentally leaking secrets like API keys, tokens, or - credentials to MCP server subprocesses. Secret-source-injected vars are - an exception: users configured that backend specifically so Hermes and - its subprocesses can consume those credentials without duplicating them - in every MCP server's ``env:`` block. - """ - try: - from hermes_cli.env_loader import get_secret_source - except Exception: # pragma: no cover — early bootstrap/import fallback - get_secret_source = None - env = {} - for key, value in os.environ.items(): - if ( - key in _SAFE_ENV_KEYS - or key.upper() in _SAFE_ENV_KEYS_CASE_INSENSITIVE - or key.startswith("XDG_") - or (get_secret_source is not None and get_secret_source(key)) - ): - env[key] = value - if user_env: - env.update(user_env) - return env - - -def _sanitize_error(text: str) -> str: - """Strip credential-like patterns from error text before returning to LLM. - - Replaces tokens, keys, and other secrets with [REDACTED] to prevent - accidental credential exposure in tool error responses. - """ - return _CREDENTIAL_PATTERN.sub("[REDACTED]", text) - - -def _exc_str(exc: BaseException) -> str: - """Return a non-empty human-readable string for *exc*. - - Some exception classes (e.g. ``anyio.ClosedResourceError``) are raised - without a message argument, so ``str(exc)`` is ``""``. This helper - falls back to ``repr(exc)`` so that error messages shown to the user - and logged to disk always carry *some* diagnostic information. - """ - text = str(exc).strip() - return text if text else repr(exc) - - -# JSON-RPC "method not found" — the error a server returns when it does not -# implement a requested method (e.g. a tool-capable server that never wired up -# the optional ``ping`` utility). -32601 is the JSON-RPC 2.0 spec constant; -# _ensure_mcp_sdk() overrides it from mcp.types when the SDK is loaded (kept -# lazy so this module never triggers the ~260ms `mcp` import at import time). +# JSON-RPC 2.0 "method not found" (e.g. a server without the optional ``ping``). +# _ensure_mcp_sdk() overrides it from mcp.types once the SDK is loaded. _JSONRPC_METHOD_NOT_FOUND = -32601 -# 2026-07-28 stateless servers answering a legacy ``initialize`` reject it -# with one of these: UnsupportedProtocolVersion (-32022, spec-reserved range) -# or plain method-not-found when the handshake methods are gone entirely. -# Structural codes only — checked via _handshake_rejected_as_modern(). -_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION = -32022 - - -def _handshake_rejected_as_modern(exc: BaseException) -> bool: - """True when a failed ``initialize`` signals a 2026-07-28-only server. - - Mirrors :func:`_is_method_not_found_error`'s structural-then-substring - shape (never ``isinstance`` on SDK exception types — the SDK wraps - task-group errors in ``ExceptionGroup`` and symbols drift across - generations; see references/sdk-exceptiongroup-wrapping.md). - """ - err = getattr(exc, "error", None) - code = getattr(err, "code", None) or getattr(exc, "code", None) - if code in (_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, _JSONRPC_METHOD_NOT_FOUND): - return True - msg = str(exc).lower() - if not msg: - return False - return ( - "unsupported protocol version" in msg - or str(_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION) in msg - or _is_method_not_found_error(exc) - ) - - -def _is_method_not_found_error(exc: BaseException) -> bool: - """Return True if *exc* is a JSON-RPC ``method not found`` (-32601). - - ``ping`` is an *optional* MCP utility (spec: "optional ping mechanism"). - A server that doesn't implement it answers a ping with -32601 rather than - an empty result. Structurally inspect ``MCPError.error.code`` first, then - fall back to a substring match so detection survives SDK version drift and - servers that surface the condition as a plain message. - - The substring fallback matters when a server reports method-not-found - without a structural ``-32601`` code (e.g. surfaced as a plain exception - string). Besides the canonical "method not found", many JSON-RPC - implementations phrase it as "Unknown method: " — agentmemory's MCP - server is one such case (#50028). Without matching that phrasing the - ping→list_tools fallback never latches and the keepalive reconnect-loops. - """ - # Structural: mcp.shared.exceptions.MCPError carries ErrorData.code. - err = getattr(exc, "error", None) - code = getattr(err, "code", None) - if code == _JSONRPC_METHOD_NOT_FOUND: - return True - msg = str(exc).lower() - if not msg: - return False - return ( - str(_JSONRPC_METHOD_NOT_FOUND) in msg - or "method not found" in msg - or "unknown method" in msg - or "not found: ping" in msg - ) - - -# --------------------------------------------------------------------------- -# MCP tool description content scanning -# --------------------------------------------------------------------------- - -# Patterns that indicate potential prompt injection in MCP tool descriptions. -# These are WARNING-level — we log but don't block, since false positives -# would break legitimate MCP servers. -_MCP_INJECTION_PATTERNS = [ - (re.compile(r"ignore\s+(all\s+)?previous\s+instructions", re.I), - "prompt override attempt ('ignore previous instructions')"), - (re.compile(r"you\s+are\s+now\s+a", re.I), - "identity override attempt ('you are now a...')"), - (re.compile(r"your\s+new\s+(task|role|instructions?)\s+(is|are)", re.I), - "task override attempt"), - (re.compile(r"system\s*:\s*", re.I), - "system prompt injection attempt"), - (re.compile(r"<\s*(system|human|assistant)\s*>", re.I), - "role tag injection attempt"), - (re.compile(r"do\s+not\s+(tell|inform|mention|reveal)", re.I), - "concealment instruction"), - (re.compile(r"(curl|wget|fetch)\s+https?://", re.I), - "network command in description"), - (re.compile(r"base64\.(b64decode|decodebytes)", re.I), - "base64 decode reference"), - (re.compile(r"exec\s*\(|eval\s*\(", re.I), - "code execution reference"), - (re.compile(r"import\s+(subprocess|os|shutil|socket)", re.I), - "dangerous import reference"), -] - - -def _scan_mcp_description(server_name: str, tool_name: str, description: str) -> List[str]: - """Scan an MCP tool description for prompt injection patterns. - - Returns a list of finding strings (empty = clean). - """ - findings = [] - if not description: - return findings - for pattern, reason in _MCP_INJECTION_PATTERNS: - if pattern.search(description): - findings.append(reason) - if findings: - logger.warning( - "MCP server '%s' tool '%s': suspicious description content — %s. " - "Description: %.200s", - server_name, tool_name, "; ".join(findings), - description, - ) - return findings - - -def _prepend_path(env: dict, directory: str) -> dict: - """Prepend *directory* to env PATH if it is not already present.""" - updated = dict(env or {}) - if not directory: - return updated - - existing = updated.get("PATH", "") - parts = [part for part in existing.split(os.pathsep) if part] - if directory not in parts: - parts = [directory, *parts] - updated["PATH"] = os.pathsep.join(parts) if parts else directory - return updated - - -# Safety cap on nextCursor pagination loops so a misbehaving server that -# returns a cursor forever cannot spin discovery indefinitely. 50 pages at -# the common 50-100 items/page covers thousands of tools/resources/prompts. +# Cap on nextCursor pagination so a server returning a cursor forever cannot +# spin discovery; 50 pages at 50-100 items/page covers thousands of entries. _MCP_LIST_MAX_PAGES = 50 async def _paginate_full_list(list_method, items_attr: str, server_name: str, cache_meta_out: Optional[dict] = None): - """Drain a paginated MCP ``list_*`` call by following ``nextCursor``. + """Drain a paginated ``list_*`` call by following ``nextCursor``. - The MCP spec allows servers to paginate ``tools/list``, - ``resources/list``, and ``prompts/list`` responses via an opaque - ``nextCursor`` token. The Python SDK's ``ClientSession.list_*`` methods - fetch exactly one page per call, so a client that never passes the - cursor back silently sees only the first page — on a paginated server - every tool/resource/prompt past page 1 would be invisible to the agent. - - Args: - list_method: Bound ``session.list_tools`` / ``list_resources`` / - ``list_prompts`` coroutine function. - items_attr: Result attribute holding the page's items - (``"tools"``, ``"resources"``, or ``"prompts"``). - server_name: For log messages. - cache_meta_out: Optional dict that receives the first page's - SEP-2549 cache hints (``ttl_ms``, ``cache_scope``) when the - server provides them (2026-07-28 servers MUST; earlier ones - won't). Callers use ``ttl_ms`` to bound the schema cache. - - Returns: - Combined list of items across all pages. Callers must hold the - server's ``_rpc_lock`` for the duration so pages come from a - consistent snapshot. + The SDK fetches one page per call, so without this every entry past page + 1 would be invisible. ``cache_meta_out`` receives the first page's + SEP-2549 hints (``ttl_ms``, ``cache_scope``) when present. Callers must + hold the server's ``_rpc_lock`` so pages come from a consistent snapshot. """ items: list = [] cursor = None @@ -966,9 +511,7 @@ async def _paginate_full_list(list_method, items_attr: str, server_name: str, if not cursor: result = await list_method() else: - # Cursor continuation differs by SDK generation: mcp 1.x - # accepts ``cursor=``, mcp 2.0 takes ``params=`` (a - # PaginatedRequestParams). Try modern first, fall back. + # mcp 2.0 takes params=PaginatedRequestParams, 1.x takes cursor=. try: _params_cls = getattr(_mcp_types(), "PaginatedRequestParams", None) if _params_cls is not None: @@ -986,8 +529,7 @@ async def _paginate_full_list(list_method, items_attr: str, server_name: str, cache_meta_out["cache_scope"] = _scope items.extend(getattr(result, items_attr, None) or []) cursor = mcp_field(result, "next_cursor", "nextCursor") - # Per the MCP spec the cursor is an opaque string; anything else - # (including mock objects in tests) means "no more pages". + # Cursor is an opaque string; anything else (incl. mocks) = last page. if not isinstance(cursor, str) or not cursor: break else: @@ -1005,1388 +547,17 @@ def _mcp_types(): return _t -def _resolve_stdio_command(command: str, env: dict) -> tuple[str, dict]: - """Resolve a stdio MCP command against the exact subprocess environment. - - This primarily exists to make bare ``npx``/``npm``/``node`` commands work - reliably even when MCP subprocesses run under a filtered PATH. - """ - resolved_command = os.path.expanduser(str(command).strip()) - resolved_env = dict(env or {}) - - if os.sep not in resolved_command: - path_arg = resolved_env["PATH"] if "PATH" in resolved_env else None - which_hit = shutil.which(resolved_command, path=path_arg) - if which_hit is None and sys.platform == "win32" and resolved_env: - # shutil.which(..., path=...) resolves extensions from the PARENT - # process PATHEXT, not the MCP subprocess env — so a config that - # supplies both PATH and PATHEXT can fail to resolve a command - # its own env can find (#56536). Retry with the config's PATHEXT - # (any key casing: PATHEXT / Pathext / pathext) applied. - cfg_pathext = next( - (v for k, v in resolved_env.items() - if k.upper() == "PATHEXT" and isinstance(v, str) and v.strip()), - None, - ) - if cfg_pathext and cfg_pathext != os.environ.get("PATHEXT"): - _saved = os.environ.get("PATHEXT") - try: - os.environ["PATHEXT"] = cfg_pathext - which_hit = shutil.which(resolved_command, path=path_arg) - finally: - if _saved is None: - os.environ.pop("PATHEXT", None) - else: - os.environ["PATHEXT"] = _saved - if which_hit: - resolved_command = which_hit - elif resolved_command in {"npx", "npm", "node"}: - hermes_home = os.path.expanduser( - os.getenv( - "HERMES_HOME", os.path.join(os.path.expanduser("~"), ".hermes") - ) - ) - candidates = [ - os.path.join(hermes_home, "node", "bin", resolved_command), - os.path.join(os.path.expanduser("~"), ".local", "bin", resolved_command), - # /usr/local/bin is the canonical install location for Node on - # Linux from-source builds, the upstream node:bookworm-slim - # image (which the Hermes Docker image copies node + npm + - # corepack from since #4977), and macOS Homebrew on Intel. - # Without this candidate, any MCP server configured with an - # env.PATH that omits /usr/local/bin (a common pattern when - # users hand-author PATH for sandboxing) fails with ENOENT - # at execvp, and a naive symlink workaround into the user's - # PATH only fails one layer deeper because npx's shebang - # re-execs /usr/bin/env node which needs the same directory. - os.path.join(os.sep, "usr", "local", "bin", resolved_command), - ] - for candidate in candidates: - if os.path.isfile(candidate) and os.access(candidate, os.X_OK): - resolved_command = candidate - break - - command_dir = os.path.dirname(resolved_command) - if command_dir: - resolved_env = _prepend_path(resolved_env, command_dir) - - return resolved_command, resolved_env - - -def _wrap_command_with_watchdog(command: str, args: list) -> tuple[str, list]: - """Wrap a stdio MCP server command in the parent-death watchdog supervisor. - - On POSIX, the watchdog records this process's PID and later detects parent - death directly through ``getppid()``. Returns the (command, args) unchanged - on non-POSIX platforms or if the PID cannot be read. - """ - if os.name != "posix": - # Relies on process groups (os.getpgid/os.killpg); no POSIX - # equivalent wired up here yet, matching the existing killpg-based - # orphan cleanup's platform scope (Windows falls back to plain - # os.kill there too). - return command, args - try: - my_pid = os.getpid() - except Exception: - # Never let watchdog bookkeeping failure block a real MCP connection. - return command, args - watchdog_args = [ - os.path.join(os.path.dirname(os.path.abspath(__file__)), "mcp_stdio_watchdog.py"), - "--ppid", str(my_pid), - "--", - command, - *args, - ] - return sys.executable, watchdog_args - - -# --------------------------------------------------------------------------- -# MCP ImageContent block → Hermes MEDIA tag -# --------------------------------------------------------------------------- - - -def _is_reserved_mcp_meta_key(key: str) -> bool: - """Return True if an MCP ``_meta`` key uses a protocol-reserved prefix. - - Per the MCP spec's key-name rules, a prefix is reserved when a - ``modelcontextprotocol`` or ``mcp`` label is followed by at least one - more label (``modelcontextprotocol.io/...``, ``tools.mcp.com/...``). - A trailing reserved word (``com.example.mcp/...``) is a legitimate - vendor namespace and passes through. Ported from - MoonshotAI/kimi-code#2600. - """ - slash = key.find("/") - if slash <= 0: - return False - labels = key[:slash].split(".") - return any( - label in ("modelcontextprotocol", "mcp") and i < len(labels) - 1 - for i, label in enumerate(labels) - ) - - -def _strip_reserved_meta_keys(meta) -> "Optional[Dict[str, Any]]": - """Drop protocol-reserved keys from a tool result's ``_meta`` mapping. - - Returns the filtered dict, or ``None`` when there is nothing - model-facing left (or the input wasn't a mapping). - """ - if not isinstance(meta, dict): - return None - out = {k: v for k, v in meta.items() - if isinstance(k, str) and not _is_reserved_mcp_meta_key(k)} - return out or None - - -def _mcp_image_extension_for_mime_type(mime_type: str) -> str: - """Return a reasonable file extension for an MCP image MIME type.""" - import mimetypes - normalized = (mime_type or "").split(";", 1)[0].strip().lower() - if normalized in {"image/jpeg", "image/jpg"}: - return ".jpg" - return mimetypes.guess_extension(normalized) or ".png" - - -def _cache_mcp_image_block(block) -> str: - """Cache an MCP ``ImageContent`` block to the shared image cache and - return a ``MEDIA:`` tag that Hermes gateways know how to render. - - Returns an empty string when *block* is not an image, when the base64 - payload is malformed, or when the cache helper rejects the bytes (e.g. - non-image MIME masquerading as an image). Errors are logged, not raised: - a single bad block shouldn't kill the tool result, and the caller will - fall through to any text blocks that did parse. - """ - import base64 - - data = getattr(block, "data", None) - mime_type = mcp_field(block, "mime_type", "mimeType") - normalized_mime = str(mime_type or "").split(";", 1)[0].strip().lower() - if data is None or not normalized_mime.startswith("image/"): - return "" - - try: - raw_bytes = base64.b64decode(data) - except (TypeError, ValueError) as exc: - logger.warning("MCP image block decode failed (%s): %s", normalized_mime, exc) - return "" - - try: - from gateway.platforms.base import cache_image_from_bytes - - image_path = cache_image_from_bytes( - raw_bytes, - ext=_mcp_image_extension_for_mime_type(normalized_mime), - ) - except ImportError: - # gateway.platforms.base not importable in this process (e.g. cron - # without gateway deps). Fall back to silently dropping — callers - # get any text blocks that did parse. - logger.debug("MCP image caching skipped — gateway.platforms.base unavailable") - return "" - except Exception as exc: - logger.warning("MCP image block cache failed: %s", exc) - return "" - - return f"MEDIA:{image_path}" - - -# --------------------------------------------------------------------------- -# MCP resource blocks (ResourceLink / EmbeddedResource / AudioContent) -# --------------------------------------------------------------------------- - -# Hard cap on decoded resource bytes materialized from an MCP tool result. -# Prevents a misbehaving server from filling the cache disk via one block. -_MCP_RESOURCE_MAX_BYTES = 50 * 1024 * 1024 - -# Base64 expands raw bytes by ~4/3; reject oversized payloads before decoding -# so a multi-GB blob string is never transiently doubled in memory. -_MCP_RESOURCE_MAX_B64_CHARS = _MCP_RESOURCE_MAX_BYTES * 4 // 3 + 4 - - -def _mcp_resource_filename(uri: str, mime_type: str) -> str: - """Derive a safe display filename for an MCP resource. - - Only the last path segment of the URI is considered, and only as a - *name hint* — `cache_document_from_bytes` re-sanitizes and prefixes it, - so remote path components can't influence the cache location. - """ - import mimetypes - import re as _re - from pathlib import Path - from urllib.parse import urlparse, unquote - - name = "" - if uri: - try: - name = Path(unquote(urlparse(str(uri)).path or "")).name - except (ValueError, TypeError): - name = "" - # Strip control characters (newlines/ANSI escapes from hostile URIs would - # otherwise land in the filename and the transcript marker) and cap the - # length, preserving the extension. - name = _re.sub(r"[\x00-\x1f\x7f]", "", name).strip() - if len(name) > 150: - stem, dot, ext = name.rpartition(".") - if dot and 0 < len(ext) <= 12: - name = stem[: 150 - len(ext) - 1] + "." + ext - else: - name = name[:150] - if not name or name in {".", ".."}: - normalized = (mime_type or "").split(";", 1)[0].strip().lower() - ext = mimetypes.guess_extension(normalized) or ".bin" - name = f"resource{ext}" - return name - - -def _cache_mcp_audio_block(block) -> str: - """Cache an MCP ``AudioContent`` block and return a ``MEDIA:`` tag. - - Returns an empty string when *block* is not audio or on any failure — - same fail-open contract as ``_cache_mcp_image_block``. - """ - import base64 - - data = getattr(block, "data", None) - mime_type = str(mcp_field(block, "mime_type", "mimeType") or "").split(";", 1)[0].strip().lower() - if data is None or not mime_type.startswith("audio/"): - return "" - if len(data) > _MCP_RESOURCE_MAX_B64_CHARS: - return f"[MCP audio resource too large to cache: ~{len(data) * 3 // 4} bytes]" - try: - raw_bytes = base64.b64decode(data) - except (TypeError, ValueError) as exc: - logger.warning("MCP audio block decode failed (%s): %s", mime_type, exc) - return "" - if len(raw_bytes) > _MCP_RESOURCE_MAX_BYTES: - return f"[MCP audio resource too large to cache: {len(raw_bytes)} bytes]" - try: - from gateway.platforms.base import cache_audio_from_bytes - import mimetypes - - ext = ( - {"audio/wav": ".wav", "audio/x-wav": ".wav", "audio/wave": ".wav"}.get(mime_type) - or mimetypes.guess_extension(mime_type) - or ".ogg" - ) - audio_path = cache_audio_from_bytes(raw_bytes, ext=ext) - except ImportError: - logger.debug("MCP audio caching skipped — gateway.platforms.base unavailable") - return "" - except Exception as exc: - logger.warning("MCP audio block cache failed: %s", exc) - return "" - return f"MEDIA:{audio_path}" - - -def _render_mcp_resource_block(block, server_name: str = "") -> str: - """Render an MCP ``ResourceLink`` or ``EmbeddedResource`` block as text. - - - ``EmbeddedResource`` with text contents → the text itself. - - ``EmbeddedResource`` with blob contents → bytes are decoded (size-capped) - and materialized into the Hermes document cache; returns a marker with - the local path so file/terminal tools can consume it. - - ``ResourceLink`` → the URI plus a pointer at the server's read_resource - tool. No network fetch happens here; the link is only readable through - the originating MCP session. - - Returns an empty string for non-resource blocks. Failures are logged and - reported inline rather than silently dropping the block. - """ - block_type = getattr(block, "type", "") - - if block_type == "resource_link" or ( - hasattr(block, "uri") and not hasattr(block, "resource") and block_type != "text" - ): - uri = getattr(block, "uri", None) - if not uri: - return "" - name = getattr(block, "name", "") or "" - mime = mcp_field(block, "mime_type", "mimeType", "") or "" - details = f"uri={uri}" - if name: - details += f", name={name}" - if mime: - details += f", mimeType={mime}" - reader = ( - mcp_prefixed_tool_name(server_name, "read_resource") - if server_name - else "the MCP server's read_resource tool" - ) - return f"[MCP resource link: {details} — fetch it with {reader}]" - - resource = getattr(block, "resource", None) - if resource is None: - return "" - - text = getattr(resource, "text", None) - if text is not None: - return strip_unicode_tags(str(text)) - - blob = getattr(resource, "blob", None) - if blob is None: - return "" - - import base64 - - uri = str(getattr(resource, "uri", "") or "") - mime = str(mcp_field(resource, "mime_type", "mimeType", "") or "") - if len(blob) > _MCP_RESOURCE_MAX_B64_CHARS: - return f"[MCP embedded resource too large to cache: ~{len(blob) * 3 // 4} bytes, uri={uri}]" - try: - raw_bytes = base64.b64decode(blob) - except (TypeError, ValueError) as exc: - logger.warning("MCP embedded resource decode failed (%s): %s", mime or uri, exc) - return f"[MCP embedded resource could not be decoded: {mime or uri}]" - if len(raw_bytes) > _MCP_RESOURCE_MAX_BYTES: - return f"[MCP embedded resource too large to cache: {len(raw_bytes)} bytes, uri={uri}]" - try: - from gateway.platforms.base import cache_document_from_bytes - - path = cache_document_from_bytes(raw_bytes, _mcp_resource_filename(uri, mime)) - except ImportError: - logger.debug("MCP resource caching skipped — gateway.platforms.base unavailable") - return f"[MCP embedded resource received ({len(raw_bytes)} bytes, {mime or 'unknown type'}) but document cache unavailable in this process]" - except Exception as exc: - logger.warning("MCP embedded resource cache failed: %s", exc) - return f"[MCP embedded resource could not be cached: {mime or uri}]" - detail = mime or "unknown type" - return f"[MCP resource saved to {path} ({detail}, {len(raw_bytes)} bytes) — read it with read_file or terminal tools]" - - -# --------------------------------------------------------------------------- -# Remote MCP URL validation -# --------------------------------------------------------------------------- - - -class InvalidMcpUrlError(ValueError): - """Raised when a remote MCP server's ``url`` cannot be parsed as http(s)://. - - Validated once at startup so we fail fast with a clear message instead of - burning through the reconnect-backoff loop on every attempt. (Ported from - anomalyco/opencode#25019.) - """ - - -class NonMcpEndpointError(ConnectionError): - """Raised when an HTTP MCP URL serves a non-MCP response. - - A genuine MCP Streamable-HTTP endpoint answers with ``application/json`` - or ``text/event-stream``. Anything else on a 2xx response (typically - ``text/html`` from a web-app root) means the configured ``url`` points at - the wrong place. This is non-retryable: every attempt returns the same - page, so the reconnect-backoff loop is skipped and the server is reported - failed immediately with an actionable message. - - Subclasses :class:`ConnectionError` so callers that only catch the broad - class still treat it as a connection problem. - """ - - -def _unwrap_exception_group(exc: BaseException) -> BaseException: - """Extract the root-cause exception from anyio TaskGroup wrappers. - - The MCP SDK uses anyio task groups, which wrap errors in - ``BaseExceptionGroup`` / ``ExceptionGroup``. Their ``str()`` is opaque — - "unhandled errors in a TaskGroup (1 sub-exception)" — so log sites must - unwrap to surface the real cause (e.g. ``BrokenPipeError`` on a dead - stdio pipe, "401 Unauthorized" on an auth failure). - - Adapted from :func:`hermes_cli.mcp_config._unwrap_exception_group` with - two extra behaviours needed on the runtime path: - - - **Fatal leaves re-raise.** A ``KeyboardInterrupt`` / ``SystemExit`` - anywhere in the (possibly nested) group must propagate to the - interpreter, never be flattened into a loggable error. - - **Prefer non-cancellation leaves.** When a group carries both a real - error and the ``CancelledError``s that anyio cancellation sprays across - sibling tasks, the real error is the root cause worth logging. - """ - while isinstance(exc, BaseExceptionGroup) and exc.exceptions: - fatal, _rest = exc.split((KeyboardInterrupt, SystemExit)) - if fatal is not None: - # Surface the fatal signal itself, not the wrapper. - leaf: BaseException = fatal - while isinstance(leaf, BaseExceptionGroup) and leaf.exceptions: - leaf = leaf.exceptions[0] - raise leaf - # Prefer a non-cancellation leaf when one exists: cancellation - # noise from sibling tasks should not mask the real error. - chosen = exc.exceptions[0] - for sub in exc.exceptions: - if not _contains_only_cancellation(sub): - chosen = sub - break - exc = chosen - return exc - - -def _contains_only_cancellation(exc: BaseException) -> bool: - """True if ``exc`` is (or a group containing only) CancelledError.""" - if isinstance(exc, BaseExceptionGroup): - return all(_contains_only_cancellation(sub) for sub in exc.exceptions) - return isinstance(exc, asyncio.CancelledError) - - -def _classify_mcp_failure(exc: BaseException) -> str: - """Classify an MCP connection failure as ``'permanent'`` or ``'transient'``. - - Permanent failures are deterministic — every retry hits the same wall, so - burning the retry ladder (and log lines) on them is pure noise; ``run()`` - parks them immediately: - - - auth failures (401/403) — need new credentials, not a retry; - - :class:`NonMcpEndpointError` — the URL serves a web page, not MCP; - - :class:`InvalidMcpUrlError` — unusable config; - - ``FileNotFoundError`` / ``ENOENT`` — the stdio command doesn't exist. - - Everything else (network blips, EOF, ``ClosedResourceError``, transport - TaskGroup drops, timeouts) is transient and keeps the normal - retry-with-backoff ladder. - """ - root = _unwrap_exception_group(exc) - if _is_auth_error(root): - return "permanent" - if isinstance(root, (NonMcpEndpointError, InvalidMcpUrlError)): - return "permanent" - # Stdio command missing: FileNotFoundError, or an OSError carrying ENOENT. - if isinstance(root, FileNotFoundError): - return "permanent" - if isinstance(root, OSError) and getattr(root, "errno", None) == errno.ENOENT: - return "permanent" - # httpx.HTTPStatusError with 401/403 that _is_auth_error's type-gate - # missed (e.g. auth types not importable in this environment). - status = getattr(getattr(root, "response", None), "status_code", None) - if status in (401, 403): - return "permanent" - return "transient" - - -def _validate_remote_mcp_url(server_name: str, url: Any) -> str: - """Return the URL as a string if it's a valid http(s) remote MCP URL. - - Raises :class:`InvalidMcpUrlError` otherwise with a message naming the - offending server, so users can spot the bad entry in their config. - - Accepts: - - ``http://host`` / ``https://host`` with optional port, path, query - - IPv4, IPv6 (bracketed), DNS hostnames - - Rejects: - - Non-string values (``None``, dicts, ints) - - Missing scheme (``example.com/mcp``) - - Non-http(s) schemes (``file://``, ``ws://``, ``stdio:`` — stdio servers - use the ``command`` key, not ``url``) - - Empty host (``http://``, ``https:///path``) - """ - if not isinstance(url, str): - raise InvalidMcpUrlError( - f"Invalid MCP URL for '{server_name}': expected a string, got " - f"{type(url).__name__}" - ) - stripped = url.strip() - if not stripped: - raise InvalidMcpUrlError( - f"Invalid MCP URL for '{server_name}': empty url" - ) - try: - parsed = urlparse(stripped) - except Exception as exc: # urlparse is very permissive — belt and braces - raise InvalidMcpUrlError( - f"Invalid MCP URL for '{server_name}': {stripped!r} ({exc})" - ) from exc - if parsed.scheme.lower() not in {"http", "https"}: - raise InvalidMcpUrlError( - f"Invalid MCP URL for '{server_name}': scheme must be http or " - f"https, got {parsed.scheme!r} ({stripped!r})" - ) - if not parsed.netloc: - raise InvalidMcpUrlError( - f"Invalid MCP URL for '{server_name}': missing host ({stripped!r})" - ) - # ``urlparse`` accepts ``http://:8080`` (empty host, explicit port). - # Reject that — we need a real host. - if not parsed.hostname: - raise InvalidMcpUrlError( - f"Invalid MCP URL for '{server_name}': missing hostname " - f"({stripped!r})" - ) - return stripped - - -def _resolve_client_cert(server_name: str, config: dict): - """Resolve the ``client_cert`` / ``client_key`` config for mTLS. - - Returns whatever ``httpx``'s ``cert=`` parameter accepts, or ``None`` when - no client certificate is configured: - - - ``None`` if neither ``client_cert`` nor ``client_key`` is set. - - A single absolute path string if ``client_cert`` is a string and - ``client_key`` is unset (PEM file with cert + key combined). - - A ``(cert_path, key_path)`` tuple when both are set, or when - ``client_cert`` is a 2-element list/tuple. - - A ``(cert_path, key_path, password)`` tuple when ``client_cert`` is - a 3-element list/tuple — the third element is the key passphrase. - - User paths support ``~`` expansion. Missing files raise ``FileNotFoundError`` - with a server-scoped message so the failure surfaces as a clear setup - error rather than 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: - return None - - def _expand(path: Any, label: str) -> str: - if not isinstance(path, str) or not path.strip(): - raise ValueError( - f"MCP server '{server_name}': {label} must be a non-empty " - f"string path (got {type(path).__name__})" - ) - expanded = os.path.expanduser(path.strip()) - if not os.path.isfile(expanded): - raise FileNotFoundError( - f"MCP server '{server_name}': {label} not found at " - f"{expanded!r}" - ) - return expanded - - # Tuple/list form for client_cert — (cert, key) or (cert, key, password). - if isinstance(raw_cert, (list, tuple)): - if raw_key is not None: - raise ValueError( - f"MCP server '{server_name}': specify either client_cert as " - f"a list [cert, key] OR client_cert + client_key, not both" - ) - if len(raw_cert) == 2: - cert_path = _expand(raw_cert[0], "client_cert[0]") - key_path = _expand(raw_cert[1], "client_cert[1]") - return (cert_path, key_path) - if len(raw_cert) == 3: - cert_path = _expand(raw_cert[0], "client_cert[0]") - key_path = _expand(raw_cert[1], "client_cert[1]") - password = raw_cert[2] - if not isinstance(password, str): - raise ValueError( - f"MCP server '{server_name}': client_cert[2] (key " - f"passphrase) must be a string" - ) - return (cert_path, key_path, password) - raise ValueError( - f"MCP server '{server_name}': client_cert list form must have 2 " - f"or 3 elements (got {len(raw_cert)})" - ) - - # String form for client_cert. - cert_path = _expand(raw_cert, "client_cert") - if raw_key is not None: - key_path = _expand(raw_key, "client_key") - return (cert_path, key_path) - # Single combined PEM file (cert + key in one file). - return cert_path - - -def _resolve_identity_header(server_name: str, config: dict): - """Resolve the optional per-server ``identity_header`` config. - - Config shape (in the server's ``mcp_servers`` entry):: - - identity_header: - name: "X-User-Id" - value_from: "static" # or "profile"; default: static - value: "alice" # required when value_from is static - - Returns a ``(header_name, header_value)`` tuple, or ``None`` when the - key is unset or invalid. Invalid configs warn and are ignored — an - identity header must never break the server connection. ``profile`` - mode resolves the value to the active Hermes profile name once at - connect time; there is no per-call mutation. - """ - raw = config.get("identity_header") - if raw is None: - return None - if not isinstance(raw, dict): - logger.warning( - "MCP server '%s': identity_header must be a mapping with " - "'name' and 'value'/'value_from' keys (got %s) — ignoring", - server_name, type(raw).__name__, - ) - return None - name = raw.get("name") - if not isinstance(name, str) or not name.strip(): - logger.warning( - "MCP server '%s': identity_header requires a non-empty " - "'name' — ignoring", server_name, - ) - return None - value_from = (raw.get("value_from") or "static").strip().lower() - if value_from == "static": - value = raw.get("value") - if not isinstance(value, str) or not value.strip(): - logger.warning( - "MCP server '%s': identity_header with value_from: static " - "requires a non-empty string 'value' — ignoring", - server_name, - ) - return None - return (name.strip(), value) - if value_from == "profile": - from hermes_cli.profiles import get_active_profile_name - return (name.strip(), get_active_profile_name()) - logger.warning( - "MCP server '%s': identity_header value_from must be 'static' or " - "'profile' (got %r) — ignoring", server_name, value_from, - ) - return None - - -def _apply_identity_header(server_name: str, config: dict, headers: dict) -> dict: - """Merge the resolved identity header into ``headers`` (in place). - - An explicit per-server ``headers`` entry with the same name (any - casing) wins — the identity header never silently overrides user - config. - """ - resolved = _resolve_identity_header(server_name, config) - if resolved is None: - return headers - name, value = resolved - if any(key.lower() == name.lower() for key in headers): - logger.debug( - "MCP server '%s': identity_header '%s' already set via explicit " - "headers config — keeping the explicit value", server_name, name, - ) - return headers - headers[name] = value - return headers - - -def _make_redirect_header_stripper( - original_url, - *, - strict: bool = False, - configured_header_names: "set[str] | frozenset[str]" = frozenset(), -): - """Build an httpx response hook that guards cross-origin redirects. - - Always strips ``Authorization`` when a redirect leaves the original - origin. When *strict* is true (portable Agent Plugins v1 packages with - ``strict_redirect_headers``), every *configured* header (lowercase names - in *configured_header_names*) is stripped as well — the v1 spec forbids - forwarding package-configured headers to a different origin without - explicit user authorization. - """ - - async def _strip_on_cross_origin_redirect(response): - if response.is_redirect and response.next_request: - target = response.next_request.url - if (target.scheme, target.host, target.port) != ( - original_url.scheme, original_url.host, original_url.port, - ): - response.next_request.headers.pop("authorization", None) - response.next_request.headers.pop("Authorization", None) - if strict: - for _name in configured_header_names: - while _name in response.next_request.headers: - del response.next_request.headers[_name] - - return _strip_on_cross_origin_redirect - - -def _format_connect_error(exc: BaseException) -> str: - """Render nested MCP connection errors into an actionable short message.""" - - def _find_missing(current: BaseException) -> Optional[str]: - nested = getattr(current, "exceptions", None) - if nested: - for child in nested: - missing = _find_missing(child) - if missing: - return missing - return None - if isinstance(current, FileNotFoundError): - if getattr(current, "filename", None): - return str(current.filename) - match = re.search(r"No such file or directory: '([^']+)'", str(current)) - if match: - return match.group(1) - for attr in ("__cause__", "__context__"): - nested_exc = getattr(current, attr, None) - if isinstance(nested_exc, BaseException): - missing = _find_missing(nested_exc) - if missing: - return missing - return None - - def _flatten_messages(current: BaseException) -> List[str]: - nested = getattr(current, "exceptions", None) - if nested: - flattened: List[str] = [] - for child in nested: - flattened.extend(_flatten_messages(child)) - return flattened - messages = [] - text = str(current).strip() - if text: - messages.append(text) - for attr in ("__cause__", "__context__"): - nested_exc = getattr(current, attr, None) - if isinstance(nested_exc, BaseException): - messages.extend(_flatten_messages(nested_exc)) - return messages or [current.__class__.__name__] - - missing = _find_missing(exc) - if missing: - message = f"missing executable '{missing}'" - if os.path.basename(missing) in {"npx", "npm", "node"}: - message += ( - " (ensure Node.js is installed and PATH includes its bin directory, " - "or set mcp_servers..command to an absolute path and include " - "that directory in mcp_servers..env.PATH)" - ) - return _sanitize_error(message) - - deduped: List[str] = [] - for item in _flatten_messages(exc): - if item not in deduped: - deduped.append(item) - return _sanitize_error("; ".join(deduped[:3])) - - -# --------------------------------------------------------------------------- -# Sampling -- server-initiated LLM requests (MCP sampling/createMessage) -# --------------------------------------------------------------------------- - -def _safe_numeric(value, default, coerce=int, minimum=1): - """Coerce a config value to a numeric type, returning *default* on failure. - - Handles string values from YAML (e.g. ``"10"`` instead of ``10``), - non-finite floats, and values below *minimum*. - """ - try: - result = coerce(value) - if isinstance(result, float) and not math.isfinite(result): - return default - return max(result, minimum) - except (TypeError, ValueError, OverflowError): - return default - - -class SamplingHandler: - """Handles sampling/createMessage requests for a single MCP server. - - .. deprecated-upstream:: MCP 2026-07-28 deprecates the Sampling feature - (SEP-2577, 12-month window; suggested migration is direct LLM-provider - integration server-side). This handler stays fully functional for the - deprecation window because handshake-era servers in the wild still - issue sampling/createMessage — but do NOT grow new capability here; - modern servers use MRTR (``resultType: "input_required"``) instead of - server-initiated requests, which the SDK's session layer handles. - - Each MCPServerTask that has sampling enabled creates one SamplingHandler. - The handler is callable and passed directly to ``ClientSession`` as - the ``sampling_callback``. All state (rate-limit timestamps, metrics, - tool-loop counters) lives on the instance -- no module-level globals. - - The callback is async and runs on the MCP background event loop. The - sync LLM call is offloaded to a thread via ``asyncio.to_thread()`` so - it doesn't block the event loop. - """ - - _STOP_REASON_MAP = {"stop": "endTurn", "length": "maxTokens", "tool_calls": "toolUse"} - - def __init__(self, server_name: str, config: dict): - self.server_name = server_name - self.max_rpm = _safe_numeric(config.get("max_rpm", 10), 10, int) - self.timeout = _safe_numeric(config.get("timeout", 30), 30, float) - self.max_tokens_cap = _safe_numeric(config.get("max_tokens_cap", 4096), 4096, int) - self.max_tool_rounds = _safe_numeric( - config.get("max_tool_rounds", 5), 5, int, minimum=0, - ) - self.model_override = config.get("model") - self.allowed_models = config.get("allowed_models", []) - - _log_levels = {"debug": logging.DEBUG, "info": logging.INFO, "warning": logging.WARNING} - self.audit_level = _log_levels.get( - str(config.get("log_level", "info")).lower(), logging.INFO, - ) - - # Per-instance state - self._rate_timestamps: List[float] = [] - self._tool_loop_count = 0 - self.metrics = {"requests": 0, "errors": 0, "tokens_used": 0, "tool_use_count": 0} - - # -- Rate limiting ------------------------------------------------------- - - def _check_rate_limit(self) -> bool: - """Sliding-window rate limiter. Returns True if request is allowed.""" - now = time.time() - window = now - 60 - self._rate_timestamps[:] = [t for t in self._rate_timestamps if t > window] - if len(self._rate_timestamps) >= self.max_rpm: - return False - self._rate_timestamps.append(now) - return True - - # -- Model resolution ---------------------------------------------------- - - def _resolve_model(self, preferences) -> Optional[str]: - """Config override > server hint > None (use default).""" - if self.model_override: - return self.model_override - if preferences and hasattr(preferences, "hints") and preferences.hints: - for hint in preferences.hints: - if hasattr(hint, "name") and hint.name: - return hint.name - return None - - # -- Message conversion -------------------------------------------------- - - @staticmethod - def _extract_tool_result_text(block) -> str: - """Extract text from a ToolResultContent block.""" - if not hasattr(block, "content") or block.content is None: - return "" - items = block.content if isinstance(block.content, list) else [block.content] - return "\n".join(item.text for item in items if hasattr(item, "text")) - - def _convert_messages(self, params) -> List[dict]: - """Convert MCP SamplingMessages to OpenAI format. - - Uses ``msg.content_as_list`` (SDK helper) so single-block and - list-of-blocks are handled uniformly. Dispatches per block type - with ``isinstance`` on real SDK types when available, falling back - to duck-typing via ``hasattr`` for compatibility. - """ - # The presence of a tool-use id is the discriminator for a tool - # *result* block, so it has to be read under both spellings (see - # mcp_field) — on mcp 2.x a bare ``hasattr(b, "toolUseId")`` is False - # for every block, which silently drops tool results out of the - # conversation and pushes them down the "unsupported block type" path - # below. - def _tool_use_id(block): - return mcp_field(block, "tool_use_id", "toolUseId", _MISSING) - - def _is_tool_use(block): - return hasattr(block, "name") and hasattr(block, "input") - - messages: List[dict] = [] - for msg in params.messages: - blocks = msg.content_as_list if hasattr(msg, "content_as_list") else ( - msg.content if isinstance(msg.content, list) else [msg.content] - ) - - # Separate blocks by kind. - 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) - ] - - # Emit tool result messages (role: tool) - for tr in tool_results: - messages.append({ - "role": "tool", - "tool_call_id": _tool_use_id(tr), - "content": self._extract_tool_result_text(tr), - }) - - # Emit assistant tool_calls message - if tool_uses: - tc_list = [] - for tu in tool_uses: - tc_list.append({ - "id": getattr(tu, "id", f"call_{len(tc_list)}"), - "type": "function", - "function": { - "name": tu.name, - "arguments": json.dumps(tu.input, ensure_ascii=False) if isinstance(tu.input, dict) else str(tu.input), - }, - }) - msg_dict: dict = {"role": msg.role, "tool_calls": tc_list} - # Include any accompanying text - text_parts = [b.text for b in content_blocks if hasattr(b, "text")] - if text_parts: - msg_dict["content"] = "\n".join(text_parts) - messages.append(msg_dict) - elif content_blocks: - # Pure text/image content - if len(content_blocks) == 1 and hasattr(content_blocks[0], "text"): - messages.append({"role": msg.role, "content": content_blocks[0].text}) - else: - parts = [] - for block in content_blocks: - block_mime = mcp_field( - block, "mime_type", "mimeType", _MISSING - ) - if hasattr(block, "text"): - parts.append({"type": "text", "text": block.text}) - elif hasattr(block, "data") and block_mime is not _MISSING: - parts.append({ - "type": "image_url", - "image_url": {"url": f"data:{block_mime};base64,{block.data}"}, - }) - else: - logger.warning( - "Unsupported sampling content block type: %s (skipped)", - type(block).__name__, - ) - if parts: - messages.append({"role": msg.role, "content": parts}) - - return messages - - # -- Error helper -------------------------------------------------------- - - @staticmethod - def _error(message: str, code: int = -1): - """Return ErrorData (MCP spec) or raise as fallback.""" - if _MCP_SAMPLING_TYPES: - return ErrorData(code=code, message=message) - raise Exception(message) - - # -- Response building --------------------------------------------------- - - def _build_tool_use_result(self, choice, response): - """Build a CreateMessageResultWithTools from an LLM tool_calls response.""" - self.metrics["tool_use_count"] += 1 - - # Tool loop governance - if self.max_tool_rounds == 0: - self._tool_loop_count = 0 - return self._error( - f"Tool loops disabled for server '{self.server_name}' (max_tool_rounds=0)" - ) - - self._tool_loop_count += 1 - if self._tool_loop_count > self.max_tool_rounds: - self._tool_loop_count = 0 - return self._error( - f"Tool loop limit exceeded for server '{self.server_name}' " - f"(max {self.max_tool_rounds} rounds)" - ) - - content_blocks = [] - for tc in choice.message.tool_calls: - args = tc.function.arguments - if isinstance(args, str): - try: - parsed = json.loads(args) - except (json.JSONDecodeError, ValueError): - logger.warning( - "MCP server '%s': malformed tool_calls arguments " - "from LLM (wrapping as raw): %.100s", - self.server_name, args, - ) - parsed = {"_raw": args} - else: - parsed = args if isinstance(args, dict) else {"_raw": str(args)} - - content_blocks.append(ToolUseContent( - type="tool_use", - id=tc.id, - name=tc.function.name, - input=parsed, - )) - - logger.log( - self.audit_level, - "MCP server '%s' sampling response: model=%s, tokens=%s, tool_calls=%d", - self.server_name, response.model, - getattr(getattr(response, "usage", None), "total_tokens", "?"), - len(content_blocks), - ) - - return CreateMessageResultWithTools( - 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.""" - self._tool_loop_count = 0 # reset on text response - response_text = choice.message.content or "" - - logger.log( - self.audit_level, - "MCP server '%s' sampling response: model=%s, tokens=%s", - self.server_name, response.model, - getattr(getattr(response, "usage", None), "total_tokens", "?"), - ) - - return CreateMessageResult( - role="assistant", - content=TextContent(type="text", text=_sanitize_error(response_text)), - model=response.model, - stopReason=self._STOP_REASON_MAP.get(choice.finish_reason, "endTurn"), - ) - - # -- Session kwargs helper ----------------------------------------------- - - def session_kwargs(self) -> dict: - """Return kwargs to pass to ClientSession for sampling support.""" - return { - "sampling_callback": self, - "sampling_capabilities": SamplingCapability( - tools=SamplingToolsCapability(), - ), - } - - # -- Main callback ------------------------------------------------------- - - async def __call__(self, context, params): - """Sampling callback invoked by the MCP SDK. - - Conforms to ``SamplingFnT`` protocol. Returns - ``CreateMessageResult``, ``CreateMessageResultWithTools``, or - ``ErrorData``. - """ - # Rate limit - if not self._check_rate_limit(): - logger.warning( - "MCP server '%s' sampling rate limit exceeded (%d/min)", - self.server_name, self.max_rpm, - ) - self.metrics["errors"] += 1 - return self._error( - f"Sampling rate limit exceeded for server '{self.server_name}' " - f"({self.max_rpm} requests/minute)" - ) - - # Resolve model - model = self._resolve_model( - mcp_field(params, "model_preferences", "modelPreferences") - ) - - # Get auxiliary LLM client via centralized router - from agent.auxiliary_client import call_llm - - # Model whitelist check (we need to resolve model before calling) - 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, - ) - self.metrics["errors"] += 1 - return self._error( - f"Model '{resolved_model}' not allowed for server " - f"'{self.server_name}'. Allowed: {', '.join(self.allowed_models)}" - ) - - # Convert messages - messages = self._convert_messages(params) - system_prompt = mcp_field(params, "system_prompt", "systemPrompt") - if system_prompt: - messages.insert(0, {"role": "system", "content": system_prompt}) - - # Build LLM call kwargs - max_tokens = min( - mcp_field(params, "max_tokens", "maxTokens", self.max_tokens_cap), - self.max_tokens_cap, - ) - call_temperature = None - if hasattr(params, "temperature") and params.temperature is not None: - call_temperature = params.temperature - - # Forward server-provided tools - call_tools = None - server_tools = getattr(params, "tools", None) - if server_tools: - call_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 - ] - - 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), - ) - - # Offload sync LLM call to thread (non-blocking) - def _sync_call(): - return call_llm( - task="mcp", - model=resolved_model or None, - messages=messages, - temperature=call_temperature, - max_tokens=max_tokens, - tools=call_tools, - timeout=self.timeout, - ) - - try: - response = await asyncio.wait_for( - asyncio.to_thread(_sync_call), timeout=self.timeout, - ) - except asyncio.TimeoutError: - self.metrics["errors"] += 1 - return self._error( - f"Sampling LLM call timed out after {self.timeout}s " - f"for server '{self.server_name}'" - ) - except Exception as exc: - self.metrics["errors"] += 1 - return self._error( - f"Sampling LLM call failed: {_sanitize_error(_exc_str(exc))}" - ) - - # Guard against empty choices (content filtering, provider errors) - if not getattr(response, "choices", None): - self.metrics["errors"] += 1 - return self._error( - f"LLM returned empty response (no choices) for server " - f"'{self.server_name}'" - ) - - # Track metrics - choice = response.choices[0] - self.metrics["requests"] += 1 - total_tokens = getattr(getattr(response, "usage", None), "total_tokens", 0) - if isinstance(total_tokens, int): - self.metrics["tokens_used"] += total_tokens - - # Dispatch based on response type - if ( - choice.finish_reason == "tool_calls" - and hasattr(choice.message, "tool_calls") - and choice.message.tool_calls - ): - return self._build_tool_use_result(choice, response) - - return self._build_text_result(choice, response) - - -# --------------------------------------------------------------------------- -# Elicitation handler -# --------------------------------------------------------------------------- - -def _format_elicitation_schema_summary(schema: dict, server_name: str) -> str: - """Render a JSON-schema-ish requested_schema to a human-readable field list. - - Elicitation schemas are restricted to a flat object with named top-level - properties. We surface field names, types, and descriptions so the user - can tell what the server is asking for before 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(): - field_type = "" - field_desc = "" - if isinstance(field_spec, dict): - field_type = str(field_spec.get("type", "") or "") - field_desc = str(field_spec.get("description", "") or "") - suffix = f" ({field_type})" if field_type else "" - if field_desc: - lines.append(f" - {field_name}{suffix}: {field_desc}") - else: - lines.append(f" - {field_name}{suffix}") - return "\n".join(lines) - - -class ElicitationHandler: - """Handles ``elicitation/create`` requests for a single MCP server. - - Each ``MCPServerTask`` that has elicitation enabled creates one handler. - The handler is callable and passed directly to ``ClientSession`` as the - ``elicitation_callback`` (added in mcp Python SDK 1.11.0). - - Elicitation lets a server ask the client to collect structured input from - the user mid-tool-call (e.g. payment authorization, OAuth confirmation). - Form-mode elicitations are routed through Hermes' existing approval - system (``tools.approval.prompt_dangerous_approval``), which surfaces - the prompt on whichever surface the active session uses -- CLI, TUI, - Telegram, Slack, etc. URL-mode elicitations are declined as unsupported. - - Failure modes are fail-closed: any timeout, exception, or unexpected - state returns ``decline``/``cancel`` rather than silently accepting. - The server treats this as the user not approving. - """ - - # Outer cap for the approval await. ``prompt_dangerous_approval`` runs - # its own input() timeout via the approval-config value; this is an - # asyncio-side safety net so the MCP event loop never blocks - # indefinitely if the inner timeout machinery is bypassed. - _OUTER_TIMEOUT_GRACE_SECONDS = 5 - - def __init__(self, server_name: str, config: dict, owner: Optional["MCPServerTask"] = None): - self.server_name = server_name - # Per-elicitation timeout. Default 5 min mirrors the gateway approval - # default so users on async surfaces (Telegram, Slack) have time to - # respond before the server gives up. - self.timeout = _safe_numeric(config.get("timeout", 300), 300, float) - # Back-reference to the MCPServerTask so we can read the agent's - # captured contextvars snapshot at elicitation time. Optional so - # the handler stays unit-testable in isolation. - self.owner = owner - self.metrics = { - "requests": 0, - "accepted": 0, - "declined": 0, - "errors": 0, - } - - def session_kwargs(self) -> dict: - """Return kwargs to pass to ClientSession for elicitation support.""" - return {"elicitation_callback": self} - - async def __call__(self, context, params): - """Elicitation callback invoked by the MCP SDK. - - Conforms to ``ElicitationFnT`` protocol. Returns ``ElicitResult`` - or ``ErrorData``. - """ - self.metrics["requests"] += 1 - - # URL-mode elicitations point the user to an external URL for - # sensitive out-of-band flows (OAuth, payment processing). Honouring - # them requires opening a browser to that URL and waiting for the - # server's notifications/elicitation/complete -- out of scope for - # the initial implementation. Decline cleanly so the server does - # not hang. - mode = getattr(params, "mode", "form") - if mode == "url": - logger.info( - "MCP server '%s' requested URL-mode elicitation; " - "declining (URL-mode elicitation not implemented)", - self.server_name, - ) - self.metrics["declined"] += 1 - return ElicitResult(action="decline") - - message = getattr(params, "message", "") or ( - f"MCP server '{self.server_name}' is requesting your approval" - ) - # The SDK model spells this field ``requestedSchema`` on mcp 1.x (the - # pinned version) and ``requested_schema`` on 2.0, which renamed model - # fields to snake_case and kept camelCase only as a serialization - # alias -- and pydantic aliases do not apply to attribute access. A - # single-spelling read therefore returns the ``{}`` default on the - # other generation, and _format_elicitation_schema_summary degrades to - # its generic "Approval requested by ..." line, so the user is asked to - # approve without being told which fields the server wants. - 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: tools.approval is imported very early during process - # bootstrap; matching the lazy pattern used by _fire_approval_hook - # avoids any chance of import-order coupling. - try: - from tools.approval import request_elicitation_consent - except Exception as exc: # pragma: no cover -- defensive - logger.error( - "MCP server '%s' elicitation: approval system unavailable: %s", - self.server_name, exc, - ) - self.metrics["errors"] += 1 - return ElicitResult(action="decline") - - # Offload the sync consent flow to a worker thread. Running it - # inline would freeze the MCP background event loop, blocking every - # other RPC on this session. request_elicitation_consent() routes - # itself to the right surface (gateway notify_cb for Telegram / - # Slack / etc., prompt_dangerous_approval for CLI / TUI) and - # normalizes the answer to one of accept / decline / cancel. - # - # The recv-loop task that fires this callback does NOT inherit - # the agent's contextvars (HERMES_SESSION_PLATFORM etc.). When - # the MCP tool wrapper captured the agent's context onto - # owner._pending_call_context we replay it here via - # contextvars.Context.run so the gateway-platform detection in - # request_elicitation_consent picks up the right session. - captured = getattr(self.owner, "_pending_call_context", None) if self.owner else None - - def _invoke_consent() -> str: - if captured is None: - return request_elicitation_consent( - message, - description, - timeout_seconds=int(self.timeout), - surface=f"mcp-elicitation/{self.server_name}", - ) - # Context.run can only execute a context once — copy to allow - # multiple elicitations within a single tool call. - return captured.copy().run( - request_elicitation_consent, - message, - description, - timeout_seconds=int(self.timeout), - surface=f"mcp-elicitation/{self.server_name}", - ) - - try: - answer = await asyncio.wait_for( - 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), - ) - self.metrics["errors"] += 1 - return ElicitResult(action="cancel") - except Exception as exc: - logger.error( - "MCP server '%s' elicitation failed: %s", - self.server_name, exc, exc_info=True, - ) - self.metrics["errors"] += 1 - return ElicitResult(action="decline") - - if answer == "accept": - self.metrics["accepted"] += 1 - return ElicitResult(action="accept", content={}) - if answer == "cancel": - self.metrics["errors"] += 1 - return ElicitResult(action="cancel") - self.metrics["declined"] += 1 - return ElicitResult(action="decline") - - # --------------------------------------------------------------------------- # Server task -- each MCP server lives in one long-lived asyncio Task # --------------------------------------------------------------------------- -class MCPServerTask: - """Manages a single MCP server connection in a dedicated asyncio Task. +class MCPServerTask(MCPServerTransportMixin, MCPServerHealthMixin): + """One MCP server connection living in one long-lived asyncio Task. - The entire connection lifecycle (connect, discover, serve, disconnect) - runs inside one asyncio Task so that anyio cancel-scopes created by - the transport client are entered and exited in the same Task context. - - Supports both stdio and HTTP/StreamableHTTP transports. + Connect, discover, serve and disconnect all run in that Task so the + transport's anyio cancel scopes are entered and exited in the same Task. + Transport bring-up lives in ``MCPServerTransportMixin``; keepalive, + refresh and liveness in ``MCPServerHealthMixin``. """ __slots__ = ( @@ -2413,11 +584,8 @@ class MCPServerTask: self._task: Optional[asyncio.Task] = None self._ready = asyncio.Event() self._shutdown_event = asyncio.Event() - # Set by tool handlers on auth failure after manager.handle_401() - # confirms recovery is viable. When set, _run_http / _run_stdio - # exit their async-with blocks cleanly (no exception), and the - # outer run() loop re-enters the transport so the MCP session is - # rebuilt with fresh credentials. + # When set, _run_http/_run_stdio exit their async-with cleanly and + # run() re-enters the transport (auth recovery, manual refresh, ...). self._reconnect_event = asyncio.Event() self._tools: list = [] self._error: Optional[Exception] = None @@ -2426,67 +594,48 @@ class MCPServerTask: self._elicitation: Optional[ElicitationHandler] = None self._registered_tool_names: list[str] = [] self._reconnect_retries: int = 0 - # Rapid-drop budget (#62212): a freshly (re)established session is - # UNPROVEN until it demonstrates real health — it survived at least - # one full keepalive interval (keepalive success path) or served at - # least one successful tool call. Only a proven session clears the - # reconnect budget; a transport that flaps right after the handshake - # keeps getting charged and still reaches the park instead of - # hot-cycling respawns forever. + # Rapid-drop budget: a (re)established session is UNPROVEN until it + # survives a full keepalive interval or serves a successful call. Only + # a proven session clears the reconnect budget, so a transport that + # flaps right after the handshake still reaches the park. self._session_proven: bool = False - # Set once tools have ever been registered and never cleared again, - # unlike ``_ready`` (which is cleared on every reconnect cycle). Used - # to tell a genuine first-connection failure from a later reconnect - # failure that merely happens to occur while ``_ready`` is - # momentarily clear — see the ``initial_retries`` ladder in run(). + # Set once tools were ever registered, never cleared (unlike _ready, + # which clears every reconnect cycle): separates a first-connect + # failure from a later reconnect failure in run()'s retry ladders. self._ever_connected: bool = False - # True while parked (reconnect budget exhausted) or after a park, - # until the session proves healthy again — used to log the - # parked→revived transition exactly once. + # True from park until the session proves healthy again; logs the + # parked->revived transition exactly once. self._was_parked: bool = False - # In-flight RPC bookkeeping (#48069 salvage): user-visible requests - # registered while running so a reconnect/shutdown teardown can fail - # them fast instead of orphaning them on a dying transport. + # In-flight RPC tasks, so a reconnect/shutdown teardown can fail them + # fast instead of orphaning them on a dying transport. self._inflight_tasks: set = set() - # True while a deliberate teardown is failing in-flight calls — lets - # _track_inflight_rpc convert the cancel into a retryable error. + # True while a deliberate teardown fails in-flight calls; lets + # _track_inflight_rpc turn the cancel into a retryable error. self._reconnecting: bool = False - # SuspectableBackend state (#81051/#77765/#84132): latched by races - # (teardown-vs-keepalive, auth-lock corruption); verified lazily by - # ensure_healthy() before the next call reuses the connection. + # Latched by races (teardown-vs-keepalive, auth-lock corruption); + # verified lazily by ensure_healthy() before the next call. self._suspect_reason: Optional[str] = None - # Set when a teardown failed >=1 in-flight call: the following - # reconnect is a RACE RECOVERY, not a transport failure, and must not - # charge the rapid-drop budget (a single race must never reach park). + # A teardown that failed >=1 in-flight call makes the next reconnect a + # RACE RECOVERY: it must not charge the rapid-drop budget. self._teardown_race: bool = False # One-time grace: an auth/permanent-classified failure on a previously # PROVEN session gets one suspect+reconnect cycle before the park # ladder applies (single auth-lock corruption must not park). self._permanent_grace_used: bool = False - # PIDs of the stdio subprocess spawned for the current transport - # (captured in _run_stdio). Used to fail in-flight calls FAST when - # the child dies instead of waiting out the full tool timeout - # (#81995). + # Children of the current stdio transport: lets in-flight calls fail + # FAST when the child dies instead of riding out the tool timeout. self._stdio_child_pids: Set[int] = set() self._auth_type: str = "" self._refresh_lock = asyncio.Lock() - # MCP stdio sessions are a single JSON-RPC stream. Some servers emit - # list_changed notifications during startup; if the notification - # handler calls list_tools while a normal tool call is in flight, the - # stream can wedge and the user-visible tool call times out. Serialize - # client-initiated RPCs per server. The lock is also applied to HTTP - # transports for conservative per-server ordering. + # A stdio session is one JSON-RPC stream: a list_tools issued by the + # notification handler while a tool call is in flight can wedge it. + # Serialize client-initiated RPCs per server (HTTP too, for ordering). self._rpc_lock = asyncio.Lock() self._pending_refresh_tasks: set[asyncio.Task] = set() - # contextvars snapshot of the agent task that's currently in - # session.call_tool(). The MCP recv loop dispatches incoming - # elicitation/create requests on a SEPARATE asyncio task whose - # context doesn't inherit HERMES_SESSION_PLATFORM, so the - # elicitation handler has no way to detect the gateway session - # that triggered the call. Capturing the agent's context here - # and replaying it inside the elicitation callback restores - # gateway-platform attribution and routes the approval prompt - # to the right surface (Telegram, Slack, etc.). + # contextvars snapshot of the agent task inside session.call_tool(). + # The SDK dispatches elicitation/create on a separate task that does + # not inherit HERMES_SESSION_PLATFORM; replaying this context in the + # elicitation callback routes the approval prompt to the right surface. self._pending_call_context: Optional[contextvars.Context] = None now = time.monotonic() self._lifecycle_started_at: float = now @@ -2494,597 +643,45 @@ class MCPServerTask: self._idle_timeout_seconds: Optional[float] = None self._max_lifetime_seconds: Optional[float] = None self._recycled_reason: Optional[str] = None - # Captures the ``InitializeResult`` returned by - # ``await session.initialize()`` so downstream code can inspect the - # server's real advertised capabilities (``.capabilities.resources``, - # ``.capabilities.prompts``) instead of assuming every ``ClientSession`` - # method attribute corresponds to a supported server method. See #18051. + # InitializeResult from the handshake: the server's REAL advertised + # capabilities, used instead of assuming every ClientSession method + # maps to a supported server method. self.initialize_result: Optional[Any] = None # SEP-2549 cache hints from the last tools/list (ttl_ms, cache_scope). self._list_cache_meta: dict = {} - # Set True the first time a keepalive ``ping`` returns JSON-RPC - # -32601 (method not found): the server is tool-capable but doesn't - # implement the optional ``ping`` utility. Subsequent keepalives fall - # back to ``list_tools`` (the pre-ping probe) so we neither spam pings - # nor reconnect-loop. Reset on each fresh transport connection. + # Latched when keepalive ``ping`` returns -32601 (optional utility not + # implemented); later keepalives use list_tools instead of + # reconnect-looping. Reset on every fresh transport connection. self._ping_unsupported: bool = False - def _is_http(self) -> bool: - """Check if this server uses HTTP transport.""" - return "url" in self._config + # Content types a real Streamable-HTTP endpoint may return on the initial + # POST/GET; anything else on a 2xx means the URL is not an MCP endpoint. + _MCP_CONTENT_TYPES = ("application/json", "text/event-stream") - def _advertises_tools(self) -> bool: - """Whether the server advertises the ``tools`` capability. - - Per the MCP spec, ``InitializeResult.capabilities.tools`` is non-None - iff the server implements the ``tools/*`` request family. Prompt-only - or resource-only servers omit it, and calling ``tools/list`` against - them raises ``MCPError(-32601 Method not found)`` — which previously - killed the connection during discovery and made every keepalive fail. - (Ported from anomalyco/opencode#31271.) - - Returns True when no capability info was captured (legacy fallback: - preserve the old always-call-list_tools behavior rather than regress - any server that was working before this gate). - """ - init_result = self.initialize_result - caps = getattr(init_result, "capabilities", None) if init_result is not None else None - if caps is None: - return True - return getattr(caps, "tools", None) is not None - - async def _negotiate_session(self, session, connect_timeout: float): - """Negotiate the protocol era with the server and return its result. - - MCP 2026-07-28 replaced the ``initialize``/``initialized`` handshake - with a stateless core: every request is self-describing and clients - MAY probe ``server/discover`` up front (SEP-2575). The SDK exposes - both paths on ``ClientSession`` (``initialize()`` / ``discover()``) - and ``adopt()``s whichever result installs the outbound stamp, so - the rest of this file is era-agnostic. - - Per-server ``protocol`` config key: - - - ``auto`` (default): try the legacy handshake FIRST, and fall back - to ``server/discover`` when the server signals it is modern-only - (``UnsupportedProtocolVersion`` -32022, or ``initialize`` missing - -32601). This is the reverse of the SDK's own discover-first auto - mode, on purpose: nearly every configured/catalog server today - speaks the handshake era, and initialize-first means ZERO extra - round-trips and zero behavior change for all of them, while - stateless-only servers still connect via the fallback. - - ``stateless``: probe ``server/discover`` first (one legacy retry - on MCPError, so a handshake-only server still connects). - - ``legacy``: handshake only, no fallback (escape hatch for servers - that misbehave on unknown methods). - - Both result types expose ``.capabilities``, so downstream gates - (``_advertises_tools``, ``_select_utility_schemas``, the config - probe) work unchanged on either. - """ - mode = str((self._config or {}).get("protocol", "auto")).lower().strip() - if mode in ("stateless", "modern", "2026-07-28"): - try: - return await asyncio.wait_for( - session.discover(), timeout=connect_timeout - ) - except asyncio.TimeoutError: - raise - except asyncio.CancelledError: - raise - except Exception as exc: - logger.info( - "MCP server '%s': server/discover rejected (%s) despite " - "protocol=%s — falling back to the legacy handshake", - self.name, exc, mode, - ) - return await asyncio.wait_for( - session.initialize(), timeout=connect_timeout - ) - if mode in ("legacy", "handshake"): - return await asyncio.wait_for( - session.initialize(), timeout=connect_timeout - ) - if mode != "auto": - logger.warning( - "MCP server '%s': unknown protocol=%r — treating as 'auto' " - "(valid: auto, stateless, legacy)", self.name, mode, - ) - try: - return await asyncio.wait_for( - session.initialize(), timeout=connect_timeout - ) - except asyncio.TimeoutError: - raise - except asyncio.CancelledError: - raise - except Exception as exc: - if not _handshake_rejected_as_modern(exc): - raise - if not hasattr(session, "discover"): - # Legacy SDK generation (mcp 1.x) has no server/discover - # client — nothing to fall back to. - raise - logger.info( - "MCP server '%s': legacy handshake rejected (%s) — " - "retrying via server/discover (2026-07-28 stateless server)", - self.name, exc, - ) - return await asyncio.wait_for( - session.discover(), timeout=connect_timeout - ) - - def _is_recycled_stdio(self) -> bool: - """Return True when a stdio server was intentionally recycled.""" - return not self._is_http() and self._recycled_reason is not None - - def mark_tool_call(self) -> None: - """Record that a user-visible MCP operation is starting.""" - self._last_tool_call_at = time.monotonic() - - def _mark_lifecycle_started(self) -> None: - now = time.monotonic() - self._lifecycle_started_at = now - self._last_tool_call_at = now - self._recycled_reason = None - - def _stdio_recycle_reason(self, now: Optional[float] = None) -> Optional[str]: - """Return the stdio recycle reason if idle/age limits have elapsed.""" - if self._is_http() or self._rpc_lock.locked(): - return None - now = time.monotonic() if now is None else now - if ( - self._max_lifetime_seconds is not None - and now - self._lifecycle_started_at >= self._max_lifetime_seconds - ): - return "max_lifetime_seconds" - if ( - self._idle_timeout_seconds is not None - and now - self._last_tool_call_at >= self._idle_timeout_seconds - ): - return "idle_timeout_seconds" - return None - - def _next_stdio_recycle_deadline(self) -> Optional[float]: - """Return the next monotonic recycle deadline for stdio, if any.""" - if self._is_http() or self._rpc_lock.locked(): - return None - deadlines = [] - if self._max_lifetime_seconds is not None: - deadlines.append(self._lifecycle_started_at + self._max_lifetime_seconds) - if self._idle_timeout_seconds is not None: - deadlines.append(self._last_tool_call_at + self._idle_timeout_seconds) - return min(deadlines) if deadlines else None - - def _mark_stdio_recycled(self, reason: str) -> None: - """Mark a stdio session dormant before its transport finishes closing.""" - self._recycled_reason = reason - self.session = None - - # ----- Dynamic tool discovery (notifications/tools/list_changed) ----- - - async def _refresh_tools_task(self): - """Run a dynamic tool refresh and log failures from background tasks.""" - try: - await self._refresh_tools() - except asyncio.CancelledError: - raise - except Exception: - logger.exception("MCP server '%s': dynamic tool refresh failed", self.name) - - def _schedule_tools_refresh(self) -> asyncio.Task: - """Schedule a background tool refresh and keep it strongly referenced.""" - task = asyncio.create_task(self._refresh_tools_task()) - self._pending_refresh_tasks.add(task) - task.add_done_callback(self._pending_refresh_tasks.discard) - return task - - def _make_logging_callback(self): - """Build a ``logging_callback`` for ``ClientSession``. - - Routes MCP ``notifications/message`` log notifications from the - server into Hermes' logging (agent.log via hermes_logging), tagged - with the server name. Without this, the SDK's default callback - silently discards them, so server-side warnings/errors during a - tool call were invisible. Port of anomalyco/opencode#34529. - """ - async def _on_log(params): - try: - level = _MCP_LOG_LEVEL_MAP.get( - str(getattr(params, "level", "info")).lower(), logging.INFO, - ) - data = getattr(params, "data", None) - if not isinstance(data, str): - try: - data = json.dumps(data, ensure_ascii=False, default=str) - except (TypeError, ValueError): - data = str(data) - # Cap pathological payloads so a chatty/broken server can't - # flood agent.log with megabyte lines. - if len(data) > 2000: - data = data[:2000] + "... [truncated]" - logger_name = getattr(params, "logger", None) - origin = f"{self.name}/{logger_name}" if logger_name else self.name - logger.log(level, "MCP server log [%s]: %s", origin, data) - except Exception: - logger.debug( - "Failed to handle MCP log notification from '%s'", - self.name, exc_info=True, - ) - return _on_log - - def _make_message_handler(self): - """Build a ``message_handler`` callback for ``ClientSession``. - - Dispatches on notification type. Only ``ToolListChangedNotification`` - triggers a refresh; prompt and resource change notifications are - logged as stubs for future work. - """ - async def _handler(message): - try: - if isinstance(message, Exception): - logger.debug("MCP message handler (%s): exception: %s", self.name, message) - return - if _MCP_NOTIFICATION_TYPES and isinstance(message, ServerNotification): - # mcp 2.0 turned ServerNotification from a RootModel into - # a plain union of the concrete notification types, so the - # payload IS the message instead of living under ``.root``. - # ``isinstance`` accepts a union, so the guard above still - # holds on both generations; only the unwrap changes. - # Without this, ``message.root`` raises AttributeError into - # the catch-all below and tools/list_changed refreshes stop - # firing silently. - match getattr(message, "root", message): - case ToolListChangedNotification(): - logger.info( - "MCP server '%s': received tools/list_changed notification", - self.name, - ) - # Some servers (notably mongodb-mcp-server) emit - # tools/list_changed immediately after initialize, - # while the client may already be executing another - # request. Refreshing synchronously inside the SDK - # notification handler can race with that request - # and wedge the stdio JSON-RPC stream, making all - # subsequent tool calls time out. Do the refresh in - # a separate task and let the handler return - # promptly. - self._schedule_tools_refresh() - # Yield one loop tick so tests and short-lived - # notification contexts can observe the scheduled - # refresh without awaiting the full server RPC. - await asyncio.sleep(0) - case PromptListChangedNotification(): - logger.debug("MCP server '%s': prompts/list_changed (ignored)", self.name) - case ResourceListChangedNotification(): - logger.debug("MCP server '%s': resources/list_changed (ignored)", self.name) - case _: - pass - except Exception: - logger.exception("Error in MCP message handler for '%s'", self.name) - return _handler - - async def _refresh_tools(self): - """Re-fetch tools from the server and update the registry. - - Called when the server sends ``notifications/tools/list_changed``. - The lock prevents overlapping refreshes from rapid-fire notifications. - After the initial ``await`` (list_tools), all mutations are synchronous - — atomic from the event loop's perspective. - """ - from tools.registry import registry - - if not self._advertises_tools(): - # A server that doesn't implement tools/* should never send - # tools/list_changed, but guard anyway — calling tools/list - # would raise MCPError(-32601). - return - - async with self._refresh_lock: - # Capture old tool names for change diff - old_tool_names = set(self._registered_tool_names) - - # 1. Fetch current tool list from server (follow nextCursor) - async with self._rpc_lock: - new_mcp_tools = await _paginate_full_list( - self.session.list_tools, "tools", self.name - ) - - # 2. Re-register with fresh tool list. Avoid nuke-and-repave for - # all names: live agent turns may already have tool-call IDs - # pointing at existing handler functions. Replacing entries - # in-place is enough for unchanged names and avoids transient - # "tool not connected" / stale-handler races during startup - # notifications. Tools absent from the fresh list are no longer - # callable, so remove only those stale registry entries first. - toolset_name = f"mcp-{self.name}" - stale_tool_names = old_tool_names - { - mcp_prefixed_tool_name(self.name, tool.name) - for tool in new_mcp_tools - } - for tool_name in stale_tool_names: - # Never let one server's refresh remove a colliding name that - # is currently owned by another server. - if registry.get_toolset_for_tool(tool_name) != toolset_name: - continue - registry.deregister(tool_name, scope=_server_registry_scope(self.name)) - _forget_mcp_tool_server(tool_name) - - # 3. Re-register with the fresh list. The helper may skip names that - # are ambiguous after normalization. - self._tools = new_mcp_tools - registered_names = _register_server_tools( - self.name, self, self._config - ) - - # A previously unique raw name can become ambiguous without changing - # its normalized registry name. In that case the pre-pass above does - # not consider it stale, so remove any old entry that the final, - # collision-checked registration set no longer owns. - registered_name_set = set(registered_names) - for tool_name in old_tool_names - registered_name_set: - if registry.get_toolset_for_tool(tool_name) != toolset_name: - continue - registry.deregister(tool_name, scope=_server_registry_scope(self.name)) - _forget_mcp_tool_server(tool_name) - self._registered_tool_names = registered_names - - # 4. Log what changed (user-visible notification) - new_tool_names = set(self._registered_tool_names) - added = new_tool_names - old_tool_names - removed = old_tool_names - new_tool_names - changes = [] - if added: - changes.append(f"added: {', '.join(sorted(added))}") - if removed: - changes.append(f"removed: {', '.join(sorted(removed))}") - if changes: - logger.warning( - "MCP server '%s': tools changed dynamically — %s. " - "Verify these changes are expected.", - self.name, "; ".join(changes), - ) - else: - logger.info( - "MCP server '%s': dynamically refreshed %d tool(s) (no changes)", - self.name, len(self._registered_tool_names), - ) - - async def _keepalive_probe(self) -> None: - """Exercise the session to detect a stale/expired connection. - - Uses ``ping`` (cheap, transport-agnostic liveness) by default. ``ping`` - is an OPTIONAL MCP utility: a server that doesn't implement it answers - JSON-RPC -32601. The first time that happens we latch - ``_ping_unsupported`` and fall back to the pre-ping probe — capability - permitting, ``list_tools``; otherwise ``ping`` is the only option and - the -32601 propagates (a server advertising neither a working ping nor - tools has no liveness primitive left). The latch resets on each fresh - transport connection so a server that gains ping support after a - reconnect is re-probed with the cheap path. - - Raises on a genuine connection failure so the caller triggers a - reconnect; returns normally when the session is alive. - """ - if not self._ping_unsupported: - try: - await asyncio.wait_for(self.session.send_ping(), timeout=30.0) - return - except Exception as exc: - if _is_method_not_found_error(exc): - # Structural -32601 or "Unknown method" — ping is - # definitively unsupported. - if not self._advertises_tools(): - raise - self._ping_unsupported = True - logger.info( - "MCP server '%s': does not implement the optional " - "'ping' utility (-32601); using 'list_tools' for " - "keepalive on this connection.", - self.name, - ) - elif isinstance(exc, (TimeoutError, asyncio.TimeoutError)) and self._advertises_tools(): - # A server that silently drops ping (no response at all) - # produces a TimeoutError indistinguishable from a dead - # transport. Before declaring it dead, try list_tools as - # a confirmation probe (#97245). If the transport is - # genuinely broken, list_tools will also fail and we - # propagate that failure. - try: - await asyncio.wait_for(self.session.list_tools(), timeout=30.0) - except Exception: - # Both probes failed — genuine liveness failure. - raise exc from None - # Transport alive, ping just isn't answered. Latch the - # fallback so subsequent keepalives skip the 30s wait. - self._ping_unsupported = True - logger.info( - "MCP server '%s': ping timed out but list_tools " - "succeeded — server silently drops ping; using " - "'list_tools' for keepalive on this connection.", - self.name, - ) - return - else: - # Any other error (closed transport, session expired, - # etc.) is a real liveness failure — propagate. - raise - - # Fallback probe for servers without ping support. - await asyncio.wait_for(self.session.list_tools(), timeout=30.0) - - def _mark_session_proven(self) -> None: - """Record that the current session demonstrated real health. - - Called from the keepalive success path (session survived at least one - full keepalive interval) and the tool-call success path. Only then is - the reconnect budget cleared: a handshake that completes but drops - moments later must keep consuming ``_reconnect_retries`` so a flapping - transport still reaches the park instead of respawning forever - (#62212 — 6212 spawns in 63h). - """ - if not self._session_proven: - self._session_proven = True - self._reconnect_retries = 0 - if self._was_parked: - self._was_parked = False - logger.warning( - "MCP server '%s': revived — session healthy again after " - "parking (state: parked → connected)", - self.name, - ) - # A session that just proved healthy on a fresh transport clears - # the one-time permanent-failure grace and any race bookkeeping. - self._permanent_grace_used = False - self._teardown_race = False - - # -- SuspectableBackend contract (agent.deadline) ----------------------- - - def mark_suspect(self, reason: str) -> None: - """Latch a suspicion about this connection. Cheap — no I/O. - - The NEXT call verifies via :meth:`ensure_healthy` and recycles the - transport if the probe fails, instead of the connection silently - staying poisoned until process restart (#81051/#77765/#84132). - """ - if self._suspect_reason is None and reason: - logger.warning( - "MCP server '%s': connection marked suspect (%s); next call " - "will health-check it", - self.name, reason, - ) - self._suspect_reason = reason or None - - async def ensure_healthy(self, timeout: float = 5.0) -> bool: - """Verify a suspect connection before reuse; recycle if dead. - - Returns True when healthy (suspicion cleared). On failure, requests a - reconnect, drops the stale session reference so the caller's normal - no-session path takes over, and returns False. Never raises. - """ - reason = self._suspect_reason - if not reason: - return True - if self.session is None: - # Nothing to verify — the reconnect path owns recovery now. - self._suspect_reason = None - self._reconnect_event.set() - return False - try: - await asyncio.wait_for(self._keepalive_probe(), timeout=timeout) - except Exception as exc: - root = _unwrap_exception_group(exc) - logger.warning( - "MCP server '%s': suspect connection (%s) failed health " - "check (%s: %s) — requesting reconnect (state: suspect → " - "degraded)", - self.name, reason, type(root).__name__, root, - ) - self._suspect_reason = None - self.mark_suspect(f"health check failed after {reason}") - self.session = None - self._ready.clear() - self._reconnect_event.set() - return False - logger.info( - "MCP server '%s': suspect connection passed health check " - "(%s) — clearing suspicion", - self.name, reason, - ) - self._suspect_reason = None - self._mark_session_proven() - return True - - def _fail_inflight_calls(self, reason: str) -> None: - """Cancel every in-flight RPC attached to this connection. - - Called from the lifecycle exits (reconnect/shutdown/recycle) BEFORE - the transport unwinds: the MCP SDK does not always fail pending - requests when its streams close, so without this an in-flight call - would wait out the full tool timeout on a dying transport. Cancelling - at least one task flags the cycle as a teardown race - (``_teardown_race``) so run() treats the following reconnect as - recovery rather than charging the rapid-drop budget. - """ - victims = [t for t in self._inflight_tasks if not t.done()] - if not victims: - return - self._reconnecting = True - self._teardown_race = True - self.mark_suspect(f"{reason} tore down {len(victims)} in-flight call(s)") - for task in victims: - task.cancel() - - def _stdio_children_dead(self) -> bool: - """True when every stdio child we spawned has exited. - - Best-effort: only meaningful for stdio transports with captured PIDs; - returns False (unknown → don't fail fast) otherwise. - """ - pids = getattr(self, "_stdio_child_pids", None) - if not pids or self._is_http(): - return False - try: - import psutil - except ImportError: - return False # unknown → don't fail fast - for pid in pids: - # pid_exists handles Windows without signal-permission noise; a - # probe failure is unknown, not proof that every child exited. - try: - alive = psutil.pid_exists(pid) - except Exception: - return False # unknown → don't fail fast - if alive: - return False # at least one child alive → not all dead - return True # every tracked child has exited - - async def _watch_stdio_children(self) -> None: - """Poll child liveness while a stdio RPC is in flight (#81995). - - Resolves when a tracked child dies; the caller then cancels the RPC - immediately instead of letting it hang for the full tool timeout. - """ - while True: - if self._stdio_children_dead(): - return - await asyncio.sleep(0.25) + @staticmethod + async def _cancel_waiters(*tasks: asyncio.Task) -> None: + for t in tasks: + if not t.done(): + t.cancel() + try: + await t + except (asyncio.CancelledError, Exception): + pass async def _wait_for_lifecycle_event(self) -> str: - """Block until either _shutdown_event or _reconnect_event fires. + """Serve the connection until a lifecycle event; return its kind. - Returns: - "shutdown" if the server should exit the run loop entirely. - "reconnect" if the server should tear down the current MCP - session and re-enter the transport (fresh OAuth - tokens, new session ID, etc.). The reconnect event - is cleared before return so the next cycle starts - with a fresh signal. - "recycle" if a stdio idle/max-lifetime limit elapsed. The - current transport is torn down and restarted lazily - on the next tool call. + ``"shutdown"`` exits the run loop; ``"reconnect"`` tears the session + down and re-enters the transport (event cleared before return); + ``"recycle"`` means a stdio idle/lifetime limit elapsed and the + transport restarts lazily on the next call. Shutdown wins a tie. - Shutdown takes precedence if both events are set simultaneously. - - Periodically sends a lightweight keepalive (``ping``, with a - ``list_tools`` fallback for servers that don't implement the optional - ping utility — see :meth:`_keepalive_probe`) to prevent TCP/session - state from going stale during idle periods (#17003). If the keepalive - fails, triggers a reconnect. - - The cadence is ``keepalive_interval`` from server config (default - :data:`_DEFAULT_KEEPALIVE_INTERVAL`, floored at - :data:`_MIN_KEEPALIVE_INTERVAL`). Servers that GC idle sessions on a - short TTL (e.g. Unreal Engine's editor MCP, ~15s) need an interval - below that TTL, otherwise every idle tool call lands on an - already-expired session and pays the full reconnect path. + Between events a keepalive (``ping``, list_tools fallback) runs every + ``keepalive_interval`` — which must stay below the server's session + TTL — and a failure triggers a reconnect. ``ping`` is a few bytes + regardless of tool count; list_changed notifications still arrive + out-of-band. """ - # Refresh faster than the server's session TTL. ``ping`` (MCP base - # protocol liveness) is used rather than ``list_tools`` so the probe - # stays a few bytes regardless of how many tools the server exposes — - # a ``list_tools`` keepalive against an 830-tool server would pull - # ~1 MB every cycle. Tool-list changes still arrive out-of-band via - # ``notifications/tools/list_changed`` → ``_refresh_tools``. keepalive_interval = max( _MIN_KEEPALIVE_INTERVAL, float(self._config.get("keepalive_interval", _DEFAULT_KEEPALIVE_INTERVAL)), @@ -3117,11 +714,9 @@ class MCPServerTask: self._mark_stdio_recycled(recycle_reason) return "recycle" - # Timeout — no lifecycle event fired. Probe the connection - # to detect stale/expired sessions — but NEVER while an RPC - # is in flight (#48069): the stdio session is a single - # JSON-RPC stream and a concurrent ping/list_tools can wedge - # the in-flight request. A busy server is provably alive. + # Timeout: probe for a stale session — but NEVER while an RPC + # is in flight (a concurrent ping can wedge the single stdio + # stream, and a busy server is provably alive anyway). if self.session: if self._rpc_lock.locked() or any( not t.done() for t in self._inflight_tasks @@ -3145,24 +740,16 @@ class MCPServerTask: ) self._reconnect_event.set() break - # Keepalive succeeded — the session survived a full - # keepalive interval, which is real proof of health. - # Clear the rapid-drop budget (#62212). + # Survived a full keepalive interval: real proof of health. self._mark_session_proven() finally: - for t in (shutdown_task, reconnect_task): - if not t.done(): - t.cancel() - try: - await t - except (asyncio.CancelledError, Exception): - pass + await self._cancel_waiters(shutdown_task, reconnect_task) if self._shutdown_event.is_set(): self._fail_inflight_calls("shutdown") return "shutdown" - # Deliberate teardown: fail any in-flight RPC NOW so it doesn't ride - # the dying transport to the full tool timeout (#48069/#81995). + # Deliberate teardown: fail in-flight RPCs NOW rather than letting + # them ride the dying transport to the full tool timeout. self._fail_inflight_calls("reconnect") self._reconnect_event.clear() return "reconnect" @@ -3170,24 +757,11 @@ class MCPServerTask: async def _wait_for_reconnect_or_shutdown( self, timeout: Optional[float] = None ) -> str: - """Block until a reconnect or shutdown is requested while parked. + """Wait, while parked, for a reconnect request or shutdown. - Used by :meth:`run` after the reconnect budget is exhausted. The - task stays alive (so ``_reconnect_event`` always has a listener) but - does no work until something explicitly asks it to come back — - OAuth recovery, a manual ``/mcp`` refresh — or, when ``timeout`` is - given, until the timeout elapses (a periodic self-probe). The timed - wake matters because parking deregisters this server's tools, so - no tool call can ever reach the circuit-breaker's half-open probe - or ``_signal_reconnect`` — without a self-probe a parked server - would be unrevivable short of a full reload. - - Returns: - ``"shutdown"`` if the server should exit the run loop entirely, - ``"reconnect"`` if it should rebuild the transport (explicit - request or self-probe timeout). The reconnect event is cleared - before returning so the next park cycle starts from a fresh - signal. Shutdown takes precedence. + Returns ``"shutdown"`` or ``"reconnect"`` (explicit request or, with + ``timeout``, the periodic self-probe); the reconnect event is cleared + first. Shutdown wins a tie. """ shutdown_task = asyncio.ensure_future(self._shutdown_event.wait()) reconnect_task = asyncio.ensure_future(self._reconnect_event.wait()) @@ -3198,773 +772,45 @@ class MCPServerTask: timeout=timeout, ) finally: - for t in (shutdown_task, reconnect_task): - if not t.done(): - t.cancel() - try: - await t - except (asyncio.CancelledError, Exception): - pass + await self._cancel_waiters(shutdown_task, reconnect_task) if self._shutdown_event.is_set(): return "shutdown" self._reconnect_event.clear() return "reconnect" - async def _run_stdio(self, config: dict): - """Run the server using stdio transport.""" - if config.get("identity_header") is not None: - # Headers don't exist on stdio transports — warn and ignore so a - # copy-pasted HTTP config block doesn't silently mislead. - logger.warning( - "MCP server '%s': identity_header is only supported on " - "HTTP/SSE transports — ignored for stdio servers", self.name, - ) - if not _ensure_mcp_sdk(): - raise ImportError( - f"MCP server '{self.name}' requires the 'mcp' Python SDK, but " - "it is not installed. Run `hermes setup` to install MCP support, " - "then retry." - ) + async def _park(self, revival_reason: str) -> bool: + """Drop this server's tools and wait for a reconnect request. - command = config.get("command") - args = config.get("args", []) - user_env = config.get("env") - - if not command: - raise ValueError( - f"MCP server '{self.name}' has no 'command' in config" - ) - - safe_env = _build_safe_env(user_env) - command, safe_env = _resolve_stdio_command(command, safe_env) - - # Check package against OSV malware database before spawning. - # Run off the event loop (the urllib HTTPS call is blocking) and bound - # it with a wall-clock timeout so a stalled SSL handshake can't freeze - # MCP discovery / gateway startup (#29184). The check is fail-open, so - # on timeout we log and proceed rather than blocking indefinitely. - # NOTE: must run against the REAL command/args — the watchdog wrap - # below rewrites argv to `python -m tools.mcp_stdio_watchdog …`, - # which would silently turn the preflight into a no-op. - from tools.osv_check import check_package_for_malware - try: - malware_error = await asyncio.wait_for( - asyncio.to_thread(check_package_for_malware, command, args), - timeout=_OSV_MALWARE_CHECK_TIMEOUT_S, - ) - except asyncio.TimeoutError: - logger.warning( - "MCP server '%s': OSV malware preflight timed out after %.0fs " - "(network slow/unreachable) — proceeding without the check.", - self.name, _OSV_MALWARE_CHECK_TIMEOUT_S, - ) - malware_error = None - if malware_error: - raise ValueError( - f"MCP server '{self.name}': {malware_error}" - ) - - # Wrap the real command in a parent-death watchdog supervisor so an - # ungraceful exit of this Hermes process (kill -9, crash, force-quit) - # can't leave the stdio MCP child (and its own descendants, e.g. - # mcp-remote's spawned `node`) running forever. On a clean exit, - # MCPServerTask.shutdown() / _kill_orphaned_mcp_children() still do - # the reaping as before -- this only covers the case where that code - # never gets to run. POSIX-only (relies on process groups); no-op - # elsewhere, matching existing killpg-based cleanup's platform scope. - # Applied AFTER the OSV preflight so the check inspects the real - # package, not the watchdog wrapper. - command, args = _wrap_command_with_watchdog(command, args) - - server_params = StdioServerParameters( - command=command, - args=args, - env=safe_env if safe_env else None, - cwd=config.get("cwd"), - # On Windows, pipe I/O can deliver non-UTF-8 bytes at chunk - # boundaries. Use "replace" to substitute undecodable bytes - # with U+FFFD instead of crashing with UnicodeDecodeError. - encoding_error_handler="replace", - ) - - sampling_kwargs = self._sampling.session_kwargs() if self._sampling else {} - if self._elicitation: - sampling_kwargs.update(self._elicitation.session_kwargs()) - if _MCP_NOTIFICATION_TYPES and _MCP_MESSAGE_HANDLER_SUPPORTED: - sampling_kwargs["message_handler"] = self._make_message_handler() - if _MCP_LOGGING_CALLBACK_SUPPORTED: - sampling_kwargs["logging_callback"] = self._make_logging_callback() - - # Reap any orphaned subprocesses from prior failed connection - # attempts before spawning a new one. Without this, each retry in - # the run() reconnect loop spawns a fresh process pair while the - # previous failed pair lingers — leading to rapid zombie - # accumulation (see #57355, #57228). The unscoped sweep also - # opportunistically reaps orphans left by *other* servers that - # never reconnect; per-server filtering via ``server_name`` remains - # available for scoped call sites. Run in a worker thread: the - # reaper blocks up to 2s (SIGTERM → wait → SIGKILL) when orphans - # exist, which would otherwise stall the shared MCP event loop. - await asyncio.to_thread(_kill_orphaned_mcp_children) - - # Snapshot child PIDs before spawning so we can track the new one. - pids_before = _snapshot_child_pids() - new_pids: set = set() - # Redirect subprocess stderr into a shared log file so MCP servers - # (FastMCP banners, slack-mcp startup JSON, etc.) don't dump onto - # the user's TTY and corrupt the TUI. Preserves debuggability via - # ~/.hermes/logs/mcp-stderr.log. - _write_stderr_log_header(self.name) - _errlog = _get_mcp_stderr_log() - try: - async with stdio_client(server_params, errlog=_errlog) as ( - read_stream, - write_stream, - ): - # Capture the newly spawned subprocess PID for force-kill cleanup. - # Filter out non-MCP children that race into the snapshot window: - # slash_worker and LSP servers (jdtls/pyright/yaml-ls) are spawned - # directly by the gateway without start_new_session, so their pgid - # equals the TUI parent PID. If they leak into _stdio_pgids, the - # shutdown sweep's killpg() kills the TUI parent itself. - # See agent/lsp/client.py for the complementary start_new_session fix. - new_pids = _filter_mcp_children( - _snapshot_child_pids() - pids_before - ) - if new_pids: - # Capture pgid while the child is alive — once it exits we - # can no longer call ``os.getpgid`` on it, and the cleanup - # sweep needs the pgid to reach any reparented descendants - # (e.g. ``claude mcp serve`` spawned by a stdio wrapper). - new_pgids: Dict[int, int] = {} - for _pid in new_pids: - try: - new_pgids[_pid] = os.getpgid(_pid) - except (AttributeError, ProcessLookupError, OSError): - # AttributeError: Windows (os.getpgid is POSIX-only) - # ProcessLookupError: child raced and already exited - pass - with _lock: - for _pid in new_pids: - _stdio_pids[_pid] = self.name - _stdio_pgids.update(new_pgids) - # Positive identity for the machine spawn ledger (#61514): - # record each helper child as (pid, create_time, - # 'mcp-helper', spawner=this process) so startup sweeps - # can reap orphans left after an unclean parent exit. - # Best-effort — never let ledger I/O break MCP 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, - ) - # Track the spawned children on the connection object for - # fast-fail of in-flight calls when the subprocess dies - # (#81995). - self._stdio_child_pids = set(new_pids) - async with ClientSession( - read_stream, write_stream, **sampling_kwargs - ) as session: - # Bound the MCP handshake. A stdio server that never - # completes ``initialize`` (e.g. emits a non-JSON-RPC frame - # and then blocks on stdin) otherwise hangs this coroutine - # forever on the background loop: ``connect_timeout`` only - # bounds the caller's ``.result()`` wait, not the coroutine - # itself. Because the connect never unwinds, the cleanup - # ``finally`` below never runs, so the spawned child and its - # stdio pipes/pidfd leak on every discovery retry — unbounded - # until the gateway hits EMFILE. Timing out here converts the - # hang into a normal failure, letting the ``finally`` reap the - # child. See #59349. - connect_timeout = float( - config.get("connect_timeout", _DEFAULT_CONNECT_TIMEOUT) - ) - self.initialize_result = await self._negotiate_session( - session, connect_timeout - ) - self.session = session - self._mark_lifecycle_started() - await self._discover_tools() - self._ready.set() - self._ever_connected = True - # Session is live again: clear any breaker state from a - # prior outage so the first call after recovery isn't - # gated on a stale consecutive-failure count (#16788). - _reset_server_error(self.name) - # A completed handshake alone is NOT proof of health: a - # flapping transport can handshake fine and drop moments - # later, forever (#62212). The session must prove itself - # (keepalive success or a successful tool call) before the - # reconnect budget is cleared — see _mark_session_proven. - self._session_proven = False - # stdio transport does not use OAuth, but we still honor - # _reconnect_event (e.g. future manual /mcp refresh) for - # consistency with _run_http. - return await self._wait_for_lifecycle_event() - finally: - # Runs on clean exit, exceptions, AND asyncio cancellation. - # If any of the spawned PIDs are still alive, the SDK's - # teardown failed (common when the task is cancelled mid-way - # on Linux, where setsid() children escape the parent cgroup). - # Mark them as orphans so the next cleanup sweep can reap them. - if new_pids: - from gateway.status import _pid_exists - _killpg = getattr(os, "killpg", None) - with _lock: - for _pid in new_pids: - _stdio_pids.pop(_pid, None) - for pid in new_pids: - # ``os.kill(pid, 0)`` is NOT a no-op on Windows - # (bpo-14484). Use the cross-platform check. - pid_alive = _pid_exists(pid) - pgroup_alive = False - pgid = _stdio_pgids.get(pid) - if not pid_alive and pgid is not None and _killpg is not None: - # Direct child exited but descendants may still be - # in its pgroup (e.g. ``claude mcp serve`` spawned - # by an MCP wrapper that exited first). Probe with - # signal 0 — succeeds iff any pgroup member is alive. - try: - _killpg(pgid, 0) - pgroup_alive = True - except (ProcessLookupError, PermissionError, OSError): - pgroup_alive = False - if pid_alive or pgroup_alive: - _orphan_stdio_pids.add(pid) - _orphan_stdio_pid_servers[pid] = self.name - else: - # Nothing left to reap — drop the pgid entry so - # PID-reuse can't surface stale pgroup state later. - _stdio_pgids.pop(pid, None) - - # Content types a real MCP Streamable-HTTP endpoint may return on the - # initial POST/GET. Anything else on a 2xx response means the URL is not - # an MCP endpoint. - _MCP_CONTENT_TYPES = ("application/json", "text/event-stream") - - 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 misconfigured ``mcp_servers..url`` pointed at a plain web app - returns HTML (or some other non-MCP body). The MCP SDK then sits on - the connection for the full ``connect_timeout`` (default 60 s) before - surfacing an opaque ``CancelledError``. A cheap, short-timeout probe - here catches that in ≤ ``timeout`` seconds and raises - :class:`NonMcpEndpointError` with an actionable message. - - Detection is allow-list based: a 2xx response is rejected only when it - carries a definite content type that is NOT one an MCP endpoint uses - (``application/json`` / ``text/event-stream``). When HEAD/GET returns - a non-MCP content type (e.g. ``text/html``), a lightweight JSON-RPC - ``initialize`` POST is attempted before giving up — some servers - (e.g. DocuSeal) serve a web UI on GET but speak Streamable HTTP only - via POST. - - A missing or empty content type, non-2xx status, or any - network/transport error passes through silently — the probe is - strictly best-effort, and the real handshake remains the source of - truth for everything except the unambiguous "this is a web page, - not MCP" case. - - Runs on its own httpx client OUTSIDE the SDK's anyio task group, so the - raised error propagates as itself rather than being wrapped in an - ``ExceptionGroup`` (which is what defeats hooks installed inside the - SDK transport). + The run task must NOT exit: it is the only listener on + ``_reconnect_event``, so returning leaves the server unrevivable for + the life of the process. Parking deregisters the tools, so no call + can reach the breaker probe or ``_signal_reconnect``; the wait is + therefore TIMED (one self-probe per ``_PARKED_RETRY_INTERVAL``), and + an explicit ``_reconnect_event.set()`` wakes it immediately. Returns + True when shutdown was requested instead. """ - try: - import httpx as _httpx - except ImportError: - return # No httpx → skip probe; SDK import would have failed first. - - client_kwargs: dict = { - "verify": ssl_verify, - "follow_redirects": True, - "timeout": _httpx.Timeout(timeout), - } - if client_cert is not None: - client_kwargs["cert"] = client_cert - - probe_headers = dict(headers) if headers else {} - try: - async with _httpx.AsyncClient(**client_kwargs) as client: - # HEAD is cheapest; fall back to GET if the server doesn't - # implement it (405 Method Not Allowed / 501 Not Implemented). - resp = await client.head(url, headers=probe_headers) - if resp.status_code in (405, 501): - resp = await client.get(url, headers=probe_headers) - - # Some MCP servers (e.g. DocuSeal) serve their web UI on - # HEAD/GET but speak Streamable HTTP only via POST. Before - # rejecting the endpoint, try a lightweight JSON-RPC POST - # probe so we don't false-positive on POST-only servers. - ct = ( - resp.headers.get("content-type", "") - .split(";")[0] - .strip() - .lower() - ) - if ( - ct - and ct not in self._MCP_CONTENT_TYPES - and 200 <= resp.status_code < 300 - ): - post_resp = await client.post( - url, - headers={ - **probe_headers, - "Content-Type": "application/json", - "Accept": "application/json, text/event-stream", - }, - content=( - '{"jsonrpc":"2.0","id":"_probe",' - '"method":"initialize",' - '"params":{"protocolVersion":"2025-03-26",' - '"capabilities":{},' - '"clientInfo":{"name":"hermes-probe",' - '"version":"0.1"}}}' - ), - ) - if 200 <= post_resp.status_code < 300: - post_ct = ( - post_resp.headers.get("content-type", "") - .split(";")[0] - .strip() - .lower() - ) - if post_ct in self._MCP_CONTENT_TYPES: - resp = post_resp - except _httpx.HTTPError: - return # DNS/connect/timeout/transport error — let the SDK try. - - # Only judge successful responses. A 4xx/5xx may be an auth challenge - # or a transient error the real handshake handles correctly. - if not (200 <= resp.status_code < 300): - return - - ct_base = resp.headers.get("content-type", "").split(";")[0].strip().lower() - if not ct_base: - return # No content type advertised — don't second-guess the SDK. - if ct_base in self._MCP_CONTENT_TYPES: - return # Looks like a real MCP endpoint. - - raise NonMcpEndpointError( - f"MCP server '{self.name}' at {url} returned Content-Type " - f"'{ct_base}', not an MCP response (expected one of: " - f"{', '.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/)." + self._was_parked = True + self._deregister_tools() + self._reconnect_event.clear() + parked = await self._wait_for_reconnect_or_shutdown( + timeout=_PARKED_RETRY_INTERVAL ) - - def _reconnect_or_reraise_group(self, eg: BaseExceptionGroup) -> str: - """Map an SDK transport TaskGroup failure to a clean ``"reconnect"``. - - Streamable-HTTP / SSE transports run their stream pump inside an anyio - TaskGroup. A transient stream drop (idle timeout, brief backend blip, - server-side TCP close) surfaces as a ``BaseExceptionGroup`` escaping the - transport context manager. Left unwrapped it reaches ``run()``'s error - path, which applies exponential backoff and eventually *parks* the - server for 300s and deregisters its tools — a multi-minute tool outage - for what is usually a sub-second glitch while the POST path stays - healthy (issue #66092). - - Returning ``"reconnect"`` instead lets ``run()`` rebuild the session - immediately with no backoff, no park, and no tool deregistration. - - Re-raise (rather than mask) when the failure is not a transient drop: - - shutdown is in progress (``shutdown()`` sets ``_shutdown_event`` - before it ever cancels the task); - - the group carries a ``KeyboardInterrupt`` / ``SystemExit`` — fatal - signals must propagate to the interpreter, never be converted into - a reconnect; - - the group carries a real ``CancelledError`` (task cancellation must - propagate to asyncio, mirroring the ``run()`` guard for #9930); - - we never reached a live session this attempt (``_ready`` unset) — a - connect/handshake failure SHOULD fall through to ``run()``'s backoff - rather than hot-loop reconnects against a broken endpoint. - """ - if self._shutdown_event.is_set(): - raise eg - fatal, _rest = eg.split((KeyboardInterrupt, SystemExit)) - if fatal is not None: - raise eg - cancelled, _rest = eg.split(asyncio.CancelledError) - if cancelled is not None: - raise eg - if not self._ready.is_set(): - raise eg + if parked == "shutdown": + return True logger.debug( - "MCP server '%s': transport TaskGroup exited after a live session " - "(%r) — reconnecting immediately instead of backing off", - self.name, eg, + "MCP server '%s': attempting revival %s (self-probe or explicit " + "reconnect request); rebuilding transport.", + self.name, revival_reason, ) - return "reconnect" + return False - async def _run_http(self, config: dict): - """Run the server using HTTP/StreamableHTTP transport.""" - _ensure_mcp_sdk() - if not _MCP_HTTP_AVAILABLE: - raise ImportError( - f"MCP server '{self.name}' requires HTTP transport but " - "mcp.client.streamable_http is not available. " - "Upgrade the mcp package to get HTTP support." - ) + async def _prepare_run(self, config: dict) -> bool: + """Bind config, build sampling/elicitation handlers, validate HTTP. - url = config["url"] - headers = dict(config.get("headers") or {}) - # Portable Agent Plugins v1 packages set strict_redirect_headers: - # configured headers are visible package data and MUST NOT be - # forwarded to a different origin through a redirect (spec §7.2.1). - # Capture the configured header names before client-generated - # headers (identity, protocol version) 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 (config-gated; static or - # profile-derived). Explicit headers of the same name win. - headers = _apply_identity_header(self.name, config, headers) - # Some MCP servers require MCP-Protocol-Version on the initial - # initialize request and reject session-less POSTs otherwise. - # Seed it as a client-level default, but treat user overrides as - # case-insensitive so conventional casing is preserved. - # - # Seeded from the HANDSHAKE version, not the latest one: this transport - # connects via `ClientSession.initialize()`, which sends - # LATEST_HANDSHAKE_VERSION (2025-11-25) in the body. Advertising - # 2026-07-28 in the header routes the request onto the server's - # per-request-envelope ladder, which then rejects the legacy body for - # missing its required `params._meta` envelope keys. The header has to - # agree with what the body actually speaks. - if not any(key.lower() == "mcp-protocol-version" for key in headers): - headers["mcp-protocol-version"] = LATEST_HANDSHAKE_VERSION - connect_timeout = config.get("connect_timeout", _DEFAULT_CONNECT_TIMEOUT) - ssl_verify = config.get("ssl_verify", True) - client_cert = _resolve_client_cert(self.name, config) - - # OAuth 2.1 PKCE: route through the central MCPOAuthManager so the - # same provider instance is reused across reconnects, pre-flow - # disk-watch is active, and config-time CLI code paths share state. - # If OAuth setup fails (e.g. non-interactive env without cached - # tokens), re-raise so this server is reported as failed without - # blocking other MCP servers from connecting. - _oauth_auth = None - if self._auth_type == "oauth": - try: - from tools.mcp_oauth_manager import get_manager - _oauth_auth = get_manager().get_or_build_provider( - self.name, url, config.get("oauth"), - ) - except Exception as exc: - logger.warning("MCP OAuth setup failed for '%s': %s", self.name, exc) - raise - - sampling_kwargs = self._sampling.session_kwargs() if self._sampling else {} - if self._elicitation: - sampling_kwargs.update(self._elicitation.session_kwargs()) - if _MCP_NOTIFICATION_TYPES and _MCP_MESSAGE_HANDLER_SUPPORTED: - sampling_kwargs["message_handler"] = self._make_message_handler() - if _MCP_LOGGING_CALLBACK_SUPPORTED: - sampling_kwargs["logging_callback"] = self._make_logging_callback() - - # SSE transport (for MCP servers that implement the SSE transport protocol - # rather than Streamable HTTP). Configure with ``transport: sse`` in the - # mcp_servers entry in config.yaml. - if config.get("transport") == "sse": - if _strict_cfg_headers: - # Portable packages never translate to SSE; if a config - # combines both anyway, fail closed rather than run a - # transport that cannot enforce the redirect boundary. - raise ValueError( - f"MCP server '{self.name}': strict_redirect_headers is " - "not supported on the SSE transport." - ) - if sse_client is None: - 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 governs how long sse_client will wait between - # events on the SSE stream. Using the tool_timeout (default 60s) - # here is wrong: SSE servers commonly hold the stream idle for - # minutes between events, so a 60s read timeout drops the - # connection after the first slow stretch. 300s matches the - # Streamable HTTP code path's httpx read timeout below. Original - # observation from @amiller in PR #5981 (Router Teamwork, - # Supermemory on Cloudflare Workers idle-disconnect at ~60s). - _sse_kwargs: dict = { - "url": url, - "headers": headers or None, - "timeout": float(connect_timeout), - "sse_read_timeout": 300.0, - } - if _oauth_auth is not None: - # Pass OAuth auth through to sse_client so SSE MCP servers - # behind OAuth 2.1 PKCE work. Previously built but never - # forwarded — SSE OAuth would silently fail with 401s. - _sse_kwargs["auth"] = _oauth_auth - if client_cert is not None or ssl_verify is not True: - # SSE transport doesn't expose verify/cert as kwargs, so route - # them through an httpx_client_factory that wraps the SDK's - # defaults (follow_redirects=True) and adds our TLS settings. - # The SDK calls the factory with (headers, auth, timeout); we - # forward all of those and layer verify/cert on top. - # The client MUST come from the SDK's own httpx module - # (httpx2 on mcp >= 2.0) — see sdk_httpx(). - _httpx_mod = sdk_httpx() - - _cert_for_factory = client_cert - _verify_for_factory = ssl_verify - - def _mcp_http_client_factory( - headers=None, timeout=None, auth=None, - ): - kwargs: dict = { - "follow_redirects": True, - "verify": _verify_for_factory, - } - if timeout is not None: - kwargs["timeout"] = timeout - else: - kwargs["timeout"] = _httpx_mod.Timeout(30.0, read=300.0) - if headers is not None: - kwargs["headers"] = headers - if auth is not None: - kwargs["auth"] = auth - if _cert_for_factory is not None: - kwargs["cert"] = _cert_for_factory - return _httpx_mod.AsyncClient(**kwargs) - - _sse_kwargs["httpx_client_factory"] = _mcp_http_client_factory - try: - async with sse_client(**_sse_kwargs) as (read_stream, write_stream): - async with ClientSession( - read_stream, write_stream, **sampling_kwargs - ) as session: - # Bound the handshake — same orphaned-task hang as the - # stdio path (#59349): an endpoint that accepts the - # connection but never answers ``initialize`` parks this - # coroutine forever on the background loop. - self.initialize_result = await self._negotiate_session( - session, float(connect_timeout) - ) - self.session = session - await self._discover_tools() - self._ready.set() - self._ever_connected = True - # Session is live again: clear any breaker state from a - # prior outage so the first call after recovery isn't - # gated on a stale consecutive-failure count (#16788). - _reset_server_error(self.name) - # Unproven until keepalive/tool-call success (#62212). - self._session_proven = False - reason = await self._wait_for_lifecycle_event() - if reason == "reconnect": - logger.info( - "MCP server '%s': reconnect requested — " - "tearing down SSE session", self.name, - ) - except BaseExceptionGroup as _eg: - # SSE transport TaskGroup dropped (idle timeout / stream blip): - # reconnect immediately instead of backoff/park (#66092). - reason = self._reconnect_or_reraise_group(_eg) - return reason - - if _MCP_NEW_HTTP: - # New API (mcp >= 1.24.0): build an explicit AsyncClient matching - # the SDK's own create_mcp_http_client defaults. It has to come - # from the SDK's httpx module (httpx2 on mcp >= 2.0), because the - # SDK sends its own Request objects through this client — see - # sdk_httpx(). - httpx = sdk_httpx() - - _original_url = httpx.URL(url) - - _strip_auth_on_cross_origin_redirect = _make_redirect_header_stripper( - _original_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 - if _oauth_auth is not None: - client_kwargs["auth"] = _oauth_auth - if client_cert is not None: - client_kwargs["cert"] = client_cert - - # Caller owns the client lifecycle — the SDK skips cleanup when - # http_client is provided, so we wrap in async-with. - try: - async with httpx.AsyncClient(**client_kwargs) as http_client: - # Unpacked positionally rather than by fixed arity: mcp - # 1.x yields (read, write, get_session_id) and 2.x yields - # (read, write). This file supports both SDK generations, - # and get_session_id was never used here. - async with streamable_http_client(url, http_client=http_client) as _streams: - read_stream, write_stream = _streams[0], _streams[1] - async with ClientSession(read_stream, write_stream, **sampling_kwargs) as session: - # Bound the handshake (#59349) — see stdio path. - self.initialize_result = await self._negotiate_session( - session, float(connect_timeout) - ) - self.session = session - await self._discover_tools() - self._ready.set() - self._ever_connected = True - # Session is live again: clear any breaker state from - # a prior outage so the first call after recovery - # isn't gated on a stale failure count (#16788). - _reset_server_error(self.name) - # Unproven until keepalive/tool-call success (#62212). - self._session_proven = False - reason = await self._wait_for_lifecycle_event() - if reason == "reconnect": - logger.info( - "MCP server '%s': reconnect requested — " - "tearing down HTTP session", self.name, - ) - except BaseExceptionGroup as _eg: - # Streamable-HTTP transport TaskGroup dropped: reconnect - # immediately instead of backoff/park (#66092). - reason = self._reconnect_or_reraise_group(_eg) - return reason - else: - # Deprecated API (mcp < 1.24.0): manages httpx client internally. - if _strict_cfg_headers: - # Fail closed: without an owned httpx client we cannot hook - # redirects, so the v1 cross-origin header boundary cannot be - # enforced on this SDK version. - 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 - try: - async with streamablehttp_client(url, **_http_kwargs) as ( - read_stream, write_stream, _get_session_id, - ): - async with ClientSession(read_stream, write_stream, **sampling_kwargs) as session: - # Bound the handshake (#59349) — see stdio path. - self.initialize_result = await self._negotiate_session( - session, float(connect_timeout) - ) - self.session = session - await self._discover_tools() - self._ready.set() - self._ever_connected = True - # Session is live again: clear any breaker state from a - # prior outage so the first call after recovery isn't - # gated on a stale consecutive-failure count (#16788). - _reset_server_error(self.name) - # Unproven until keepalive/tool-call success (#62212). - self._session_proven = False - reason = await self._wait_for_lifecycle_event() - if reason == "reconnect": - logger.info( - "MCP server '%s': reconnect requested — " - "tearing down legacy HTTP session", self.name, - ) - except BaseExceptionGroup as _eg: - # Legacy Streamable-HTTP transport TaskGroup dropped: reconnect - # immediately instead of backoff/park (#66092). - reason = self._reconnect_or_reraise_group(_eg) - return reason - - async def _discover_tools(self): - """Discover tools from the connected session. - - Capability-gated: prompt-only / resource-only MCP servers don't - implement ``tools/list``, and calling it raises ``MCPError(-32601)``, - which previously aborted the connection — those servers could never - stay connected for their prompts/resources. Skip the call when the - server doesn't advertise the ``tools`` capability. - (Ported from anomalyco/opencode#31271.) - """ - # Fresh transport connection → re-probe with the cheap ``ping`` path. - # Clears any latch from a prior connection in case the server gained - # ping support across the reconnect. - self._ping_unsupported = False - if self.session is None: - return - if not self._advertises_tools(): - logger.info( - "MCP server '%s': does not advertise 'tools' capability — " - "skipping tools/list (prompts/resources remain available)", - self.name, - ) - self._tools = [] - self._register_discovered_tools_if_needed() - return - async with self._rpc_lock: - self._list_cache_meta = {} - self._tools = await _paginate_full_list( - self.session.list_tools, "tools", self.name, - cache_meta_out=self._list_cache_meta, - ) - self._register_discovered_tools_if_needed() - - def _register_discovered_tools_if_needed(self) -> None: - """Re-register tools after an owned server reconnects if needed. - - Initial registration is performed by ``_discover_and_register_server`` - after ``start()`` completes. During a later reconnect, outage handling - may clear ``_ready`` before discovery and may deregister stale tools. - A managed server can still be identified by its entry in ``_servers``; - publish its freshly discovered tools before transport readiness is - restored so a successful revival cannot come back with zero tools. - A server retained after a recoverable initial failure is likewise - registry-owned before its first successful session, so ownership also - authorizes its first publication. - """ - if self._registered_tool_names: - return - if not self._ready.is_set(): - with _lock: - if _servers.get(self.name) is not self: - return - self._registered_tool_names = _register_server_tools( - self.name, self, self._config - ) - # A retained initial-failure server that just published tools has - # recovered: drop its stale connect error so status surfaces stop - # reporting it as failed. - with _lock: - if _servers.get(self.name) is self: - _server_connect_errors.pop(self.name, None) - - async def run(self, config: dict): - """Long-lived coroutine: connect, discover tools, wait, disconnect. - - Includes automatic reconnection with exponential backoff if the - connection drops unexpectedly (unless shutdown was requested). + Returns False when the server must not start: a bad remote URL or a + non-MCP endpoint (both fail fast, non-retryably, with ``_error`` set + and ``_ready`` fired) instead of burning the reconnect ladder inside + the SDK's httpx layer on every retry. """ self._config = config self.tool_timeout = _resolve_tool_timeout(config) @@ -3972,29 +818,23 @@ class MCPServerTask: self._idle_timeout_seconds = _get_lifecycle_seconds(config, "idle_timeout_seconds") self._max_lifetime_seconds = _get_lifecycle_seconds(config, "max_lifetime_seconds") - # Bind the lazily-imported SDK before reading feature flags below - # (_MCP_SAMPLING_TYPES / _MCP_ELICITATION_TYPES are False until the - # SDK import actually runs). + # The _MCP_*_TYPES flags are False until the lazy SDK import runs. _ensure_mcp_sdk() - # Set up sampling handler if enabled and SDK types are available sampling_config = config.get("sampling", {}) if sampling_config.get("enabled", True) and _MCP_SAMPLING_TYPES: self._sampling = SamplingHandler(self.name, sampling_config) else: self._sampling = None - # Set up elicitation handler if enabled and SDK types are available. - # Servers use elicitation/create to ask the client for structured - # input mid-tool-call (e.g. payment authorization). The handler - # routes those requests through Hermes' approval system. + # elicitation/create lets a server ask for structured input mid-call; + # the handler routes it through Hermes' approval system. elicitation_config = config.get("elicitation", {}) if elicitation_config.get("enabled", True) and _MCP_ELICITATION_TYPES: self._elicitation = ElicitationHandler(self.name, elicitation_config, owner=self) else: self._elicitation = None - # Validate: warn if both url and command are present if "url" in config and "command" in config: logger.warning( "MCP server '%s' has both 'url' and 'command' in config. " @@ -4003,46 +843,48 @@ class MCPServerTask: self.name, ) - # Validate remote URL once, up front. Raising here (rather than - # letting it blow up inside the SDK's httpx layer on every retry) - # means a typo in config.yaml fails fast with a clear error — and - # critically, no reconnect-backoff burn. (Ported from - # anomalyco/opencode#25019.) - if self._is_http(): + if not self._is_http(): + return True + try: + _validate_remote_mcp_url(self.name, config.get("url")) + except InvalidMcpUrlError as exc: + logger.warning("%s", exc) + self._error = exc + self._ready.set() + return False + + # Content-type preflight (Streamable HTTP only; SSE legitimately serves + # text/event-stream): a URL at a web-app root returns HTML and would + # make the SDK hang for the full connect_timeout. Skipped once _ready + # was ever set (endpoint already validated) and for OAuth servers, + # where a token-less probe sees HTML/401 and would block the flow. + if config.get("transport") != "sse" and not config.get("skip_preflight") and not self._ready.is_set() and self._auth_type != "oauth": try: - _validate_remote_mcp_url(self.name, config.get("url")) - except InvalidMcpUrlError as exc: + _probe_headers = dict(config.get("headers") or {}) + await self._preflight_content_type( + config["url"], + headers=_probe_headers, + ssl_verify=config.get("ssl_verify", True), + client_cert=_resolve_client_cert(self.name, config), + ) + except NonMcpEndpointError as exc: logger.warning("%s", exc) self._error = exc self._ready.set() - return + return False + return True - # Pre-flight content-type probe (Streamable HTTP only; SSE is - # exercised by its own client and legitimately serves - # text/event-stream). A URL pointed at a web-app root returns - # HTML, which makes the SDK hang for the full connect_timeout - # before surfacing an opaque CancelledError. Probing here — once, - # outside the SDK task group — fails fast and non-retryably with - # an actionable message, mirroring the URL-validation path above. - # Skip the probe when _ready is already set (reconnect after a - # prior successful connect) — the endpoint was validated once, - # re-probing is a redundant round-trip. Also skip for OAuth servers: - # without a cached token the endpoint returns HTML or 401, which - # would incorrectly block the OAuth flow before it can run. - if config.get("transport") != "sse" and not config.get("skip_preflight") and not self._ready.is_set() and self._auth_type != "oauth": - try: - _probe_headers = dict(config.get("headers") or {}) - await self._preflight_content_type( - config["url"], - headers=_probe_headers, - ssl_verify=config.get("ssl_verify", True), - client_cert=_resolve_client_cert(self.name, config), - ) - except NonMcpEndpointError as exc: - logger.warning("%s", exc) - self._error = exc - self._ready.set() - return + async def run(self, config: dict): + """Long-lived coroutine: connect, discover, serve, reconnect. + + State machine: connecting -> connected -> (degraded -> parked -> + revived)*. Unproven drops and transport errors charge a rapid-drop + budget with jittered exponential backoff; exhausting it (or a + permanent error) parks the server via :meth:`_park` rather than + exiting, so it stays revivable. + """ + if not await self._prepare_run(config): + return self._reconnect_retries = 0 initial_retries = 0 @@ -4054,11 +896,9 @@ class MCPServerTask: lifecycle_reason = await self._run_http(config) else: lifecycle_reason = await self._run_stdio(config) - # Transport returned cleanly. Two cases: - # - _shutdown_event was set: exit the run loop entirely. - # - _reconnect_event was set (auth recovery): loop back and - # rebuild the MCP session with fresh credentials. Do NOT - # touch the retry counters — this is not a failure. + # Clean transport return: shutdown, stdio recycle, or a + # requested rebuild (auth recovery / manual refresh / keepalive + # failure). A rebuild is not a failure for the retry counters. if self._shutdown_event.is_set(): break if lifecycle_reason == "recycle": @@ -4073,30 +913,17 @@ class MCPServerTask: break self._reconnect_event.clear() continue - # Per-cycle reconnect chatter — DEBUG. In the flapping case - # this fires on every rebuild; the WARNINGs live on the - # state transitions. + # Per-cycle chatter stays DEBUG; WARNINGs mark state transitions. logger.debug( "MCP server '%s': reconnecting (OAuth recovery or " "manual refresh)", self.name, ) - # A clean transport return means a session was established and - # then asked to rebuild (auth recovery / manual refresh / - # keepalive failure / transport TaskGroup drop). That alone is - # NOT proof of health: a flapping transport handshakes fine and - # drops moments later, and resetting the budget here let such - # servers respawn forever (#62212 — 6212 spawns in 63h). - # Only clear the consecutive-failure budget once the session - # PROVED healthy — survived >=1 full keepalive interval or - # served >=1 successful tool call (_mark_session_proven). + # A clean return is NOT proof of health (a flapping transport + # handshakes fine and drops moments later). Only a PROVEN + # session clears the budget; a teardown race is recovery, not + # a failure, and must never reach the park on its own. if self._teardown_race and not self._session_proven: - # The previous cycle ended because a teardown cancelled - # in-flight calls (keepalive/refresh race, auth recovery) - # — that is RECOVERY, not a transport failure. Do NOT - # charge the rapid-drop budget: a single race must never - # reach the park (#81051/#77765/#84132). Only genuinely - # repeated unproven drops still exhaust the budget below. logger.info( "MCP server '%s': reconnect after teardown race " "(in-flight calls were failed); not charging the " @@ -4109,8 +936,6 @@ class MCPServerTask: self._reconnect_retries = 0 backoff = 1.0 else: - # Unproven session: charge the rapid-drop budget so a - # flapping transport still reaches the park. self._reconnect_retries += 1 if self._reconnect_retries > _MAX_RECONNECT_RETRIES: logger.warning( @@ -4121,51 +946,26 @@ class MCPServerTask: self.name, _MAX_RECONNECT_RETRIES, _PARKED_RETRY_INTERVAL, ) - self._was_parked = True - self._deregister_tools() - self._reconnect_event.clear() - parked = await self._wait_for_reconnect_or_shutdown( - timeout=_PARKED_RETRY_INTERVAL - ) - if parked == "shutdown": + if await self._park("from parked state"): break - logger.debug( - "MCP server '%s': attempting revival from parked " - "state (self-probe or explicit reconnect request); " - "rebuilding transport.", - self.name, - ) - # One probe attempt per wake — see the exception-path - # park below. + # Budget of one probe per wake, so a still-dead server + # parks again instead of burning 5 rapid retries. self._reconnect_retries = _MAX_RECONNECT_RETRIES backoff = 1.0 - # Reset the session reference and readiness; _run_http/_run_stdio - # will repopulate both on successful re-entry. Leaving - # _ready set here lets handler-side recovery mistake the stale - # pre-reconnect session for a fresh one and retry too early. + # Clear readiness too: a stale _ready lets handler-side + # recovery mistake the old session for a fresh one. self._ready.clear() self.session = None continue except asyncio.CancelledError: - # Task was cancelled (shutdown, gateway restart, explicit - # task.cancel()). Don't treat this as a connection failure — - # CancelledError inherits from BaseException (not Exception) - # in Python 3.11+, so the broad ``except Exception`` below - # would NOT catch it; we'd silently exit the reconnect loop - # and the MCP server would stay dead until Hermes is fully - # restarted. Re-raise so the task's cancellation propagates - # correctly to asyncio's task machinery and ``shutdown()``'s - # ``await self._task`` completes. See #9930. + # Not a connection failure: re-raise so cancellation reaches + # asyncio and shutdown()'s ``await self._task`` completes. self.session = None raise except Exception as exc: self.session = None - # Unwrap anyio TaskGroup wrappers first: str(exc) on a - # BaseExceptionGroup is "unhandled errors in a TaskGroup - # (N sub-exceptions)" — useless in logs, and it hides the - # root cause from the auth/permanence classification below. - # Empty dead-pipe errors still get a name this way - # (e.g. "BrokenPipeError: "). + # Unwrap anyio TaskGroup wrappers: the group's str() is useless + # and hides the root cause from the classification below. root = _unwrap_exception_group(exc) failure_class = _classify_mcp_failure(root) if self._is_recycled_stdio(): @@ -4176,31 +976,15 @@ class MCPServerTask: ) self._recycled_reason = None - # If this is the first connection attempt, retry with backoff - # before giving up. A transient DNS/network blip at startup - # should not permanently kill the server. Gated on - # ``_ever_connected`` rather than ``_ready`` — ``_ready`` is - # cleared on every reconnect cycle (see below), so a server - # that already registered tools once and then dropped would - # otherwise be misclassified as never having connected and - # re-enter this initial-connect ladder (#94654). - # ``_ever_connected`` itself is set once and never cleared. - # (Ported from Kilo Code's MCP resilience fix.) + # Initial-connect ladder: a transient blip at startup must not + # kill the server. Gated on _ever_connected (never cleared), + # not _ready (cleared every reconnect cycle). if not self._ever_connected: if failure_class == "permanent": # Deterministic failure (bad command, non-MCP URL, - # 401/403): every retry hits the same wall. Park - # immediately instead of burning the retry ladder - # and spamming N identical warnings (#65673). - # - # Auth failures park here too rather than returning. - # Returning ends the run task, and with it the only - # listener on ``_reconnect_event`` — so a 401 on the - # very first connect left the server unrevivable for - # the life of the process, even after the user - # re-authenticated with ``hermes mcp login``. Parking - # keeps the task alive so the 300s self-probe (and an - # explicit /mcp refresh) can pick up fresh tokens. + # 401/403): park at once instead of burning the ladder. + # Auth failures park rather than return so the task + # stays alive to pick up fresh tokens later. if _is_auth_error(root): logger.warning( "MCP server '%s' failed initial authentication, " @@ -4219,20 +1003,8 @@ class MCPServerTask: ) self._error = exc self._ready.set() - self._was_parked = True - self._deregister_tools() - self._reconnect_event.clear() - parked = await self._wait_for_reconnect_or_shutdown( - timeout=_PARKED_RETRY_INTERVAL - ) - if parked == "shutdown": + if await self._park("after permanent initial failure"): return - logger.debug( - "MCP server '%s': attempting revival after " - "permanent initial failure (self-probe or explicit " - "reconnect request); rebuilding transport.", - self.name, - ) initial_retries = 0 self._reconnect_retries = 0 backoff = 1.0 @@ -4251,20 +1023,8 @@ class MCPServerTask: ) self._error = exc self._ready.set() - self._was_parked = True - self._deregister_tools() - self._reconnect_event.clear() - parked = await self._wait_for_reconnect_or_shutdown( - timeout=_PARKED_RETRY_INTERVAL - ) - if parked == "shutdown": + if await self._park("after initial connection failures"): return - logger.debug( - "MCP server '%s': attempting revival after initial " - "connection failures (self-probe or explicit " - "reconnect request); rebuilding transport.", - self.name, - ) initial_retries = 0 self._reconnect_retries = 0 backoff = 1.0 @@ -4298,14 +1058,9 @@ class MCPServerTask: return if failure_class == "permanent": - # Auth-lock corruption guard (#81051/#77765/#84132): an - # auth-classified permanent failure on a previously - # PROVEN session is often a transient/ambiguous state - # (OAuth flow lock left corrupt by a raced teardown), - # not truly revoked credentials. Grant ONE - # suspect+reconnect cycle before the park ladder: mark - # the connection suspect so the next call health-checks - # it, and rebuild the transport instead of parking. + # An auth failure on a PROVEN session is often a corrupt + # OAuth lock from a raced teardown, not revoked + # credentials: grant ONE suspect+reconnect cycle first. if ( _is_auth_error(root) and self._session_proven @@ -4328,10 +1083,7 @@ class MCPServerTask: if self._shutdown_event.is_set(): return continue - # A previously-working server now fails deterministically - # (revoked credentials, URL now serving a web page, stdio - # binary uninstalled). Retrying can't help — park - # immediately without burning the retry ladder. + # Deterministic failure on a working server: park now. logger.warning( "MCP server '%s' hit a permanent error, parking " "without retries; will self-probe every %ds " @@ -4339,20 +1091,8 @@ class MCPServerTask: self.name, _PARKED_RETRY_INTERVAL, type(root).__name__, root, ) - self._was_parked = True - self._deregister_tools() - self._reconnect_event.clear() - parked = await self._wait_for_reconnect_or_shutdown( - timeout=_PARKED_RETRY_INTERVAL - ) - if parked == "shutdown": + if await self._park("from parked state (permanent error)"): return - logger.debug( - "MCP server '%s': attempting revival from parked state " - "(permanent error; self-probe or explicit reconnect " - "request); rebuilding transport.", - self.name, - ) self._reconnect_retries = _MAX_RECONNECT_RETRIES backoff = 1.0 continue @@ -4367,41 +1107,12 @@ class MCPServerTask: _PARKED_RETRY_INTERVAL, type(root).__name__, root, ) - # Do NOT return — exiting the task orphans the server: - # nothing would ever listen for _reconnect_event again - # and the server would be permanently wedged for the - # life of the process (#16788). Instead, drop the phantom - # tools from the registry and park. Because parking - # deregisters the tools, no tool call can reach the - # circuit-breaker half-open probe or _signal_reconnect — - # so the park is a TIMED wait: every _PARKED_RETRY_INTERVAL - # we wake and attempt one reconnect ourselves (#57129). - # An explicit _reconnect_event.set() (OAuth recovery, - # manual /mcp refresh) still wakes us immediately. - self._was_parked = True - self._deregister_tools() - self._reconnect_event.clear() - parked = await self._wait_for_reconnect_or_shutdown( - timeout=_PARKED_RETRY_INTERVAL - ) - if parked == "shutdown": + if await self._park("from parked state"): return - logger.debug( - "MCP server '%s': attempting revival from parked state " - "(self-probe or explicit reconnect request); " - "rebuilding transport.", - self.name, - ) - # One probe attempt per wake: budget of 1 so a still-dead - # server parks again for another interval instead of - # burning 5 rapid retries each cycle. self._reconnect_retries = _MAX_RECONNECT_RETRIES backoff = 1.0 continue - # Per-attempt retry chatter stays at DEBUG; state transitions - # (connected->degraded, degraded->parked, parked->revived) - # carry the WARNINGs — one line per transition, not per try. logger.debug( "MCP server '%s' connection lost (attempt %d/%d), " "reconnecting in %.0fs: %s: %s", @@ -4416,8 +1127,7 @@ class MCPServerTask: return finally: self.session = None - # Children of this transport are gone (or about to be); - # stale PIDs must never fast-fail the NEXT transport's calls. + # Stale PIDs must never fast-fail the NEXT transport's calls. self._stdio_child_pids = set() async def start(self, config: dict): @@ -4495,59 +1205,37 @@ class MCPServerTask: return_when=asyncio.FIRST_COMPLETED, ) finally: - for task in (shutdown_task, reconnect_task): - if not task.done(): - task.cancel() - try: - await task - except (asyncio.CancelledError, Exception): - pass + await self._cancel_waiters(shutdown_task, reconnect_task) # --------------------------------------------------------------------------- -# Module-level state +# Module-level state (every mutation under ``_lock``) # --------------------------------------------------------------------------- _servers: Dict[str, MCPServerTask] = {} -# Profile registry scope that owns each live connection (None outside -# multiplex). A multiplexed /reload-mcp tears down only its own profile's -# servers; process shutdown still takes everything. +# Profile registry scope owning each live connection (None outside multiplex): +# a multiplexed /reload-mcp tears down only its own profile's servers. _server_scope_keys: Dict[str, Optional[str]] = {} _server_connecting: set[str] = set() _server_connect_errors: Dict[str, str] = {} -# Lazy MCP startup (#56832): servers whose tools were registered from the -# on-disk schema cache without spawning/connecting. Keyed by server name; -# entries are popped once a real connection is established on first use. +# Lazy startup: servers registered from the on-disk schema cache without +# connecting; popped once a real connection is established on first use. _lazy_server_configs: Dict[str, dict] = {} _lazy_server_fingerprints: Dict[str, str] = {} _lazy_server_tool_names: Dict[str, List[str]] = {} -# Discovery installs a task-local claim before calling ``_connect_server`` so -# it can retain a recoverable parked task without making standalone probe calls -# publish failed servers into module-global ownership. +# Discovery installs a task-local claim around ``_connect_server`` so it can +# retain a recoverable parked task without standalone probe calls publishing +# failed servers into module-global ownership. _connect_server_claim: contextvars.ContextVar[ Optional[Callable[[MCPServerTask], None]] ] = contextvars.ContextVar("mcp_connect_server_claim", default=None) -# Connection-retry cooldown (per-server isolation against restart storms). -# -# A single stdio MCP server that fails to spawn (bad PATH, ``exec: not -# found``, crash-on-start) is never recorded in ``_servers`` -- ``start()`` -# raises and ``_discover_and_register_server`` aborts before the -# ``_servers[name] = server`` line. Without a cooldown, EVERY subsequent -# ``discover_mcp_tools()`` (one per agent worker session, i.e. every few -# seconds) sees the server as "not connected" and re-spawns it from -# scratch. That is the restart storm in #50394: the failing server is -# re-attempted on the shared MCP event loop on every worker session, the -# subprocesses pile up unreaped, and the churn destabilises the healthy -# co-located servers (their tools intermittently surface as -# "Unknown tool"). -# -# Fix: after a failed connection attempt, stamp a monotonic -# ``retry_after`` deadline with exponential backoff. ``register_mcp_servers`` -# skips a server whose cooldown has not elapsed, so a chronically failing -# server is retried on a backoff schedule instead of on every worker -# session -- isolating it from the rest of the bridge. A successful -# connection clears the state. +# Per-server connect cooldown. A server that fails to spawn never reaches +# ``_servers``, so without this every ``discover_mcp_tools()`` (one per worker +# session) would respawn it from scratch — a restart storm whose unreaped +# subprocesses destabilise the healthy co-located servers. Failed attempts +# stamp an exponential-backoff deadline that ``register_mcp_servers`` honours; +# a successful connection clears it. _server_connect_retry_after: Dict[str, float] = {} # name -> monotonic deadline _server_connect_failures: Dict[str, int] = {} # name -> consecutive failures _CONNECT_RETRY_BASE_BACKOFF_SEC = 30.0 @@ -4555,14 +1243,7 @@ _CONNECT_RETRY_MAX_BACKOFF_SEC = 600.0 def _record_connect_failure(server_name: str) -> None: - """Stamp an exponential-backoff cooldown after a failed connect. - - Called (under ``_lock``) when a server fails its discovery/connect - attempt. The cooldown grows geometrically with the consecutive - failure count and is capped at :data:`_CONNECT_RETRY_MAX_BACKOFF_SEC`, - so a permanently-broken server settles into infrequent retries - rather than a tight respawn loop. - """ + """Stamp a geometric, capped retry cooldown after a failed connect (under ``_lock``).""" n = _server_connect_failures.get(server_name, 0) + 1 _server_connect_failures[server_name] = n backoff = min( @@ -4583,55 +1264,26 @@ def _connect_cooldown_active(server_name: str) -> bool: deadline = _server_connect_retry_after.get(server_name) return deadline is not None and time.monotonic() < deadline -# Circuit breaker: consecutive error counts per server. After -# _CIRCUIT_BREAKER_THRESHOLD consecutive failures, the handler returns -# a "server unreachable" message that tells the model to stop retrying, -# preventing the 90-iteration burn loop described in #10447. -# -# State machine: -# closed — error count below threshold; all calls go through. -# open — threshold reached; calls short-circuit until the -# cooldown elapses. -# half-open — cooldown elapsed; the next call is a probe that -# actually hits the session. Probe success → closed. -# Probe failure → reopens (cooldown re-armed). -# -# ``_server_breaker_opened_at`` records the monotonic timestamp when -# the breaker most recently transitioned into the open state. Use the -# ``_bump_server_error`` / ``_reset_server_error`` helpers to mutate -# this state — they keep the count and timestamp in sync. +# Circuit breaker per server: closed (count < threshold) -> open (calls +# short-circuit with a "stop retrying" message until the cooldown elapses) -> +# half-open (next call is a probe; success closes, failure re-arms). Mutate +# only via _bump_server_error / _reset_server_error, which keep the count and +# the open timestamp in sync. _server_error_counts: Dict[str, int] = {} _server_breaker_opened_at: Dict[str, float] = {} _CIRCUIT_BREAKER_THRESHOLD = 3 _CIRCUIT_BREAKER_COOLDOWN_SEC = 60.0 -# --------------------------------------------------------------------------- -# Trust-tier gating state (per-server trust + per-tool readOnlyHint). -# -# ``trust: full | untrusted`` is a per-server key in the MCP server config -# (config.yaml → mcp_servers..trust). On an ``untrusted`` server, -# every WRITE-CAPABLE tool call routes through the existing dangerous- -# approval surface before the RPC fires. A tool is write-capable unless its -# discovery-time ``annotations.readOnlyHint`` is exactly ``True`` -# (missing/malformed annotations fail closed to write-capable). -# -# Security model (read this before changing defaults): -# - ``readOnlyHint`` is a HINT supplied by the server itself. A hostile -# server can lie. That is precisely why the gate is tiered per-server by -# OPERATOR config: on an untrusted server the hint can only ever exempt -# tools the server claims are read-only — the worst a lie buys is -# skipping approval for calls the operator was already warned about when -# they marked the server untrusted. It can never widen access on top of -# the approval a write-capable tool would otherwise need. -# - Default trust for servers with NO ``trust`` key is ``full`` (gate off) -# for backward compatibility — existing configs keep working unchanged. -# Operators opt servers into gating explicitly with ``trust: untrusted``. -# - Any unrecognized ``trust`` value normalizes to ``untrusted`` -# (fail closed): a typo must never silently disable the gate. -# -# Classification happens at CALL TIME from data captured at DISCOVERY — -# no toolset or schema mutation, so the conversation's toolset stays -# byte-stable and prompt caching is preserved. +# Trust-tier gating (``mcp_servers..trust: full | untrusted``). On an +# untrusted server every write-capable call needs user approval before the +# RPC fires; a tool is write-capable unless its discovery-time +# ``annotations.readOnlyHint`` is exactly True (malformed fails closed). +# Security model: readOnlyHint is a server-supplied HINT and a hostile server +# can lie, but on an untrusted server a lie can only skip approval for calls +# the operator was already warned about — never widen access. Missing +# ``trust`` defaults to full (backward compatible); any unrecognized value +# normalizes to untrusted (a typo must never disable the gate). Classified +# at CALL time from DISCOVERY data: no schema mutation, prompt cache intact. _server_trust_levels: Dict[str, str] = {} _tool_read_only_hints: Dict[str, Dict[str, bool]] = {} @@ -4639,124 +1291,8 @@ _TRUST_FULL = "full" _TRUST_UNTRUSTED = "untrusted" -def _normalize_server_trust(value: Any) -> str: - """Normalize a config ``trust`` value to ``full`` or ``untrusted``. - - Missing (None) → ``full`` (backward-compatible default, documented - above). Any string other than the two known tiers → ``untrusted``: - a misspelled tier must fail closed, never silently disable gating. - """ - if value is None: - return _TRUST_FULL - text = str(value).strip().lower() - if text == _TRUST_FULL: - return _TRUST_FULL - if text == _TRUST_UNTRUSTED: - return _TRUST_UNTRUSTED - logger.warning( - "MCP trust: unrecognized trust value %r — treating as 'untrusted' " - "(valid values: full, untrusted)", value, - ) - return _TRUST_UNTRUSTED - - -def _annotation_read_only_hint(mcp_tool: Any) -> bool: - """Return True only when the tool's annotations carry readOnlyHint=True. - - Accepts both SDK annotation objects (attribute access) and plain dicts - (schema-cache JSON). Anything else — missing annotations, missing key, - non-bool truthy values — is False: unknown metadata means the tool must - be treated as write-capable. - """ - annotations = getattr(mcp_tool, "annotations", None) - if annotations is None: - return False - if isinstance(annotations, dict): - hint = annotations.get("readOnlyHint") - else: - hint = getattr(annotations, "readOnlyHint", None) - return hint is True - - -def _record_tool_trust_metadata( - server_name: str, config: dict, tools: List[Any] -) -> None: - """Capture per-server trust and per-tool readOnlyHint at discovery.""" - with _lock: - _server_trust_levels[server_name] = _normalize_server_trust( - (config or {}).get("trust") - ) - hints = _tool_read_only_hints.setdefault(server_name, {}) - for tool in tools: - name = getattr(tool, "name", None) - if name: - hints[name] = _annotation_read_only_hint(tool) - - -def _trust_gate_check(server_name: str, tool_name: str) -> Optional[str]: - """Consult the approval path for write-capable tools on untrusted servers. - - Returns None when the call may proceed, or an error string (already - formatted via ``tool_error``) when the call is blocked. Fail-closed: - approval-system errors block the call. - """ - trust = _server_trust_levels.get(server_name, _TRUST_FULL) - if trust != _TRUST_UNTRUSTED: - return None - if _tool_read_only_hints.get(server_name, {}).get(tool_name) is True: - return None - - # Lazy import mirrors the elicitation handler's pattern: tools.approval - # routes the prompt to whichever surface owns the session (CLI, TUI, - # Telegram, Slack, ...) and normalizes the answer. - try: - from tools.approval import request_elicitation_consent - - answer = request_elicitation_consent( - ( - f"MCP tool '{tool_name}' on UNTRUSTED server " - f"'{server_name}' wants to run. This tool is write-capable " - f"(no readOnlyHint=true annotation) and may modify external " - f"state." - ), - ( - f"Server '{server_name}' is configured 'trust: untrusted'. " - f"Approve to run '{tool_name}' once, or deny to block it." - ), - surface=f"mcp-trust/{server_name}", - ) - except Exception as exc: - logger.error( - "MCP trust gate: approval check failed for %s.%s: %s", - server_name, tool_name, exc, exc_info=True, - ) - return tool_error( - f"MCP tool '{tool_name}' on untrusted server '{server_name}' " - f"was blocked: the approval system was unavailable " - f"(fail-closed)." - ) - - if answer == "accept": - return None - logger.info( - "MCP trust gate: user %s '%s' on untrusted server '%s'", - "cancelled" if answer == "cancel" else "denied", - tool_name, server_name, - ) - return tool_error( - f"The user did not approve running write-capable MCP tool " - f"'{tool_name}' on untrusted server '{server_name}'. The command " - f"was NOT run. Do not retry without explicit user direction." - ) - - def _bump_server_error(server_name: str) -> None: - """Increment the consecutive-failure count for ``server_name``. - - When the count crosses :data:`_CIRCUIT_BREAKER_THRESHOLD`, stamp the - breaker-open timestamp so the cooldown clock starts (or re-starts, - for probe failures in the half-open state). - """ + """Count a failure; at the threshold (re)stamp the breaker-open time.""" n = _server_error_counts.get(server_name, 0) + 1 _server_error_counts[server_name] = n if n >= _CIRCUIT_BREAKER_THRESHOLD: @@ -4764,12 +1300,7 @@ def _bump_server_error(server_name: str) -> None: def _reset_server_error(server_name: str) -> None: - """Fully close the breaker for ``server_name``. - - Clears both the failure count and the breaker-open timestamp. Call - this on any unambiguous success signal (successful tool call, - successful reconnect, manual /mcp refresh). - """ + """Close the breaker on any unambiguous success signal.""" _server_error_counts[server_name] = 0 _server_breaker_opened_at.pop(server_name, None) @@ -4777,14 +1308,9 @@ def _reset_server_error(server_name: str) -> None: def _signal_reconnect(server: Any) -> bool: """Ask a server task to rebuild its transport, thread-safely. - The tool handlers run on caller threads, while the server task and its - ``_reconnect_event`` live on the background MCP loop. Setting an - asyncio.Event from another thread must go through - ``loop.call_soon_threadsafe``; non-async adapters and tests without a - running loop can use a direct ``.set()``. - - Returns True if a reconnect signal was delivered, False if the server - has no reconnect machinery (nothing to revive). + 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: @@ -4816,23 +1342,13 @@ def _wait_for_server_session_ready( old_session: Any = None, timeout: float = 15.0, ) -> bool: - """Wait for an MCP server to expose a usable session. + """Poll until the server exposes a usable, ready session. - Tool handlers run in normal worker threads while the MCP transport lives on - the module's background asyncio loop. During a reconnect there is a short - window where ``srv.session`` is ``None`` (or still points at the stale - session until the lifecycle coroutine has left the transport context). A - handler that blindly retries in that window can burn circuit-breaker strikes - and return ``not connected`` even though the reconnect is already in - progress. - - When ``old_session`` is supplied, require the observed session object to be - different so callers do not mistake the pre-reconnect, stale session for a - fresh one. + 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``. """ - # Iteration-bounded rather than deadline-bounded: several tests (and the - # circuit-breaker cooldown logic) monkeypatch time.monotonic to a frozen - # clock, which would make a monotonic-deadline loop spin forever. poll_interval = 0.25 iterations = max(1, int(max(float(timeout), 0.0) / poll_interval)) for i in range(iterations): @@ -4858,15 +1374,11 @@ def _signal_reconnect_and_wait( op_description: str, timeout: float = 15.0, ) -> bool: - """Ask a live MCP server task to rebuild its transport session. + """Request a transport rebuild and wait for the fresh session. - The important detail is clearing ``_ready`` on the MCP event loop before - setting ``_reconnect_event``. Older code left ``_ready`` set across - reconnects, so the caller's readiness poll could return immediately and - retry against the same dead HTTP/stream session. That was observed as - repeated ``Session terminated`` / ``not connected`` / circuit-breaker - failures in long-lived gateway sessions even though a fresh CLI process - could connect successfully. + ``_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 = _mcp_loop if loop is None or not loop.is_running(): @@ -4893,525 +1405,26 @@ def _signal_reconnect_and_wait( timeout=timeout, ) -# --------------------------------------------------------------------------- -# Auth-failure detection helpers (Task 6 of MCP OAuth consolidation) -# --------------------------------------------------------------------------- - -# Cached tuple of auth-related exception types. Lazy so this module -# imports cleanly when the MCP SDK OAuth module is missing. -_AUTH_ERROR_TYPES: tuple = () -_HTTP_STATUS_ERROR_TYPES: Optional[tuple] = None - - -def _http_status_error_types() -> tuple: - """``HTTPStatusError`` classes that can reach us, from both httpx flavours. - - A 401 can be raised either by the MCP SDK's own HTTP stack (``httpx2`` on - mcp >= 2.0) or by Hermes' pinned ``httpx``, and the two define unrelated - exception classes. Both go in the tuple so ``isinstance`` covers whichever - layer raised. - """ - global _HTTP_STATUS_ERROR_TYPES - if _HTTP_STATUS_ERROR_TYPES is not None: - return _HTTP_STATUS_ERROR_TYPES - found: list = [] - sdk_mod = sdk_httpx() - if sdk_mod is not None: - found.append(sdk_mod.HTTPStatusError) - try: - import httpx - if httpx.HTTPStatusError not in found: - found.append(httpx.HTTPStatusError) - except ImportError: - pass - _HTTP_STATUS_ERROR_TYPES = tuple(found) - return _HTTP_STATUS_ERROR_TYPES - - -def _get_auth_error_types() -> tuple: - """Return a tuple of exception types that indicate MCP OAuth failure. - - Cached after first call. Includes: - - ``mcp.client.auth.OAuthFlowError`` / ``OAuthTokenError`` — raised by - the SDK's auth flow when discovery, refresh, or full re-auth fails. - - ``mcp.client.auth.UnauthorizedError`` (older MCP SDKs) — kept as an - optional import for forward/backward compatibility. - - ``tools.mcp_oauth.OAuthNonInteractiveError`` — raised by our callback - handler when no user is present to complete a browser flow. - - ``HTTPStatusError`` from both httpx flavours — caller must - additionally check ``status_code == 401`` via :func:`_is_auth_error`. - """ - global _AUTH_ERROR_TYPES - if _AUTH_ERROR_TYPES: - return _AUTH_ERROR_TYPES - types: list = [] - try: - from mcp.client.auth import OAuthFlowError, OAuthTokenError - types.extend([OAuthFlowError, OAuthTokenError]) - except ImportError: - pass - try: - # Older MCP SDK variants exported this - from mcp.client.auth import UnauthorizedError # type: ignore - types.append(UnauthorizedError) - except ImportError: - pass - try: - from tools.mcp_oauth import OAuthNonInteractiveError - types.append(OAuthNonInteractiveError) - except ImportError: - pass - types.extend(_http_status_error_types()) - _AUTH_ERROR_TYPES = tuple(types) - return _AUTH_ERROR_TYPES - - -def _is_auth_error(exc: BaseException) -> bool: - """Return True if ``exc`` indicates an MCP OAuth failure. - - ``HTTPStatusError`` is only treated as auth-related when the response - status code is 401. Other HTTP errors fall through to the generic error - path in the tool handlers. - """ - types = _get_auth_error_types() - if not types or not isinstance(exc, types): - return False - status_error_types = _http_status_error_types() - if status_error_types and isinstance(exc, status_error_types): - return getattr(exc.response, "status_code", None) == 401 - return True - - -def _handle_auth_error_and_retry( - server_name: str, - exc: BaseException, - retry_call, - op_description: str, -): - """Attempt auth recovery and one retry; return None to fall through. - - Called by the 5 MCP tool handlers when ``session.()`` raises an - auth-related exception. Workflow: - - 1. Ask :class:`tools.mcp_oauth_manager.MCPOAuthManager.handle_401` if - recovery is viable (i.e., disk has fresh tokens, or the SDK can - refresh in-place). - 2. If yes, set the server's ``_reconnect_event`` so the server task - tears down the current MCP session and rebuilds it with fresh - credentials. Wait briefly for ``_ready`` to re-fire. - 3. Retry the operation once. Return the retry result if it produced - a non-error JSON payload. Otherwise return the ``needs_reauth`` - error dict so the model stops hallucinating manual refresh. - 4. Return None if ``exc`` is not an auth error, signalling the - caller to use the generic error path. - - Args: - server_name: Name of the MCP server that raised. - exc: The exception from the failed tool call. - retry_call: Zero-arg callable that re-runs the tool call, returning - the same JSON string format as the handler. - op_description: Human-readable name of the operation (for logs). - - Returns: - A JSON string if auth recovery was attempted, or None to fall - through to the caller's generic error path. - """ - if not _is_auth_error(exc): - return None - - from tools.mcp_oauth_manager import get_manager - manager = get_manager() - - async def _recover(): - return await manager.handle_401(server_name, None) - - try: - recovered = _run_on_mcp_loop(_recover, timeout=10) - except Exception as rec_exc: - logger.warning( - "MCP OAuth '%s': recovery attempt failed: %s", - server_name, rec_exc, - ) - recovered = False - - if recovered: - with _lock: - srv = _servers.get(server_name) - reconnected = False - if srv is not None and hasattr(srv, "_reconnect_event"): - reconnected = _signal_reconnect_and_wait( - server_name, - srv, - op_description=f"{op_description} after OAuth recovery", - timeout=15, - ) - - # A successful OAuth recovery + transport reconnect is independent - # evidence that the server is viable again, so close the circuit - # breaker here — not only on retry success. Without this, a reconnect - # followed by a failing retry would leave the breaker pinned above - # threshold forever. The post-reset retry still goes through - # _bump_server_error on failure, so a genuinely broken server will - # re-trip the breaker as normal. - if reconnected: - _reset_server_error(server_name) - - try: - result = retry_call() - try: - parsed = json.loads(result) - if "error" not in parsed: - _reset_server_error(server_name) - return result - except (json.JSONDecodeError, TypeError): - _reset_server_error(server_name) - return result - except Exception as retry_exc: - logger.warning( - "MCP %s/%s retry after auth recovery failed: %s", - server_name, op_description, retry_exc, - ) - - # No recovery available, or retry also failed: surface a structured - # needs_reauth error. Bumps the circuit breaker so the model stops - # retrying the tool. - _bump_server_error(server_name) - return tool_error( - f"MCP server '{server_name}' requires re-authentication. " - f"Run `hermes mcp login {server_name}` (or delete the tokens " - f"file under ~/.hermes/mcp-tokens/ and restart). Do NOT retry " - f"this tool — ask the user to re-authenticate.", - needs_reauth=True, - server=server_name, - ) - - -# Substrings (lower-cased match) that indicate the MCP server rejected -# the request because its server-side transport session expired / -# was garbage-collected. The caller's OAuth token is still valid — -# only the transport-layer session state needs rebuilding. See #13383. -_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", -) - -# Upper bound on exception-graph nodes inspected by -# ``_is_session_expired_error``. The visited-identity set already breaks -# cycles across ``exceptions`` / ``__cause__`` / ``__context__``; the -# budget additionally bounds pathological acyclic graphs (e.g. deeply -# chained retries) so classification always terminates promptly. Kept -# comfortably above ``sys.getrecursionlimit()`` so legitimately deep -# wrapper stacks (task-group nesting) are still fully scanned. -_EXC_TRAVERSAL_MAX_NODES = 10_000 - - -def _is_session_expired_error(exc: BaseException) -> bool: - """Return True if ``exc`` looks like an MCP transport session expiry. - - Streamable HTTP MCP servers may garbage-collect server-side session - state while the OAuth token remains valid — idle TTL, server - restart, horizontal-scaling pod rotation, etc. The SDK surfaces - this as a JSON-RPC error whose message contains phrases like - ``"Invalid or expired session"``. This class of failure is - distinct from :func:`_is_auth_error`: re-running the OAuth refresh - flow would be pointless because the access token is fine. What's - needed is a transport reconnect — tear down and rebuild the - ``streamablehttp_client`` + ``ClientSession`` pair, which is - exactly what ``MCPServerTask._reconnect_event`` triggers. - """ - # AnyIO's stream exceptions are commonly message-less. In particular, - # ``str(ClosedResourceError()) == ""``, so marker matching alone misses the - # exact failure emitted by both MCP stdio and HTTP transports. - try: - from anyio import BrokenResourceError, ClosedResourceError, EndOfStream - - transport_error_types = ( - BrokenResourceError, - ClosedResourceError, - EndOfStream, - ) - except ImportError: # pragma: no cover - AnyIO is supplied by the MCP SDK - transport_error_types = () - - # ExceptionGroup trees can be arbitrarily deep or even cyclic when custom - # exceptions expose ``exceptions``, and chained exceptions - # (``raise X from Y`` / implicit ``__context__``) can likewise form - # cycles when handlers re-raise previously seen exceptions. Traverse - # once, iteratively, with an identity-visited set AND a bounded node - # budget so classification can never spin, and inspect every reachable - # node so user interruption always overrides transport markers or types - # found elsewhere in the graph. The chain traversal matters for real - # failures: SDK wrappers often raise a generic RuntimeError *from* the - # message-less ClosedResourceError, leaving the transport signal only - # reachable via ``__cause__``. - stack: "list[BaseException | None]" = [exc] - seen: set[int] = set() - transport_error_found = False - budget = _EXC_TRAVERSAL_MAX_NODES - while stack and budget > 0: - current = stack.pop() - if current is None: - continue - identity = id(current) - if identity in seen: - continue - seen.add(identity) - budget -= 1 - - if isinstance(current, InterruptedError): - return False - if isinstance(current, transport_error_types): - transport_error_found = True - - # Exception messages vary across SDK versions + server - # implementations, so match on a small allow-list of stable - # substrings rather than exception type. Kept narrow to avoid - # false positives on unrelated server errors. - msg = str(current).lower() - if msg and any(marker in msg for marker in _SESSION_EXPIRED_MARKERS): - transport_error_found = True - - stack.extend(getattr(current, "exceptions", ())) - stack.append(getattr(current, "__cause__", None)) - stack.append(getattr(current, "__context__", None)) - - return transport_error_found - - -def _handle_session_expired_and_retry( - server_name: str, - exc: BaseException, - retry_call, - op_description: str, -): - """Trigger a transport reconnect and retry once on session expiry. - - Unlike :func:`_handle_auth_error_and_retry`, this does **not** call - the OAuth manager's ``handle_401`` — the access token is still - valid, only the server-side session state is stale. Setting - ``_reconnect_event`` causes the server task's lifecycle loop to - tear down the current ``streamablehttp_client`` + ``ClientSession`` - and rebuild them, reusing the existing OAuth provider instance. - See #13383. - - Args: - server_name: Name of the MCP server that raised. - exc: The exception from the failed call. - retry_call: Zero-arg callable that re-runs the operation, - returning the same JSON string format as the handler. - op_description: Human-readable name of the operation (logs). - - Returns: - A JSON string if reconnect + retry was attempted and produced - a response, or ``None`` to fall through to the caller's - generic error path (not a session-expired error, no server - record, reconnect didn't ready in time, or retry also failed). - """ - if not _is_session_expired_error(exc): - return None - - with _lock: - srv = _servers.get(server_name) - if srv is None or not hasattr(srv, "_reconnect_event"): - return None - - loop = _mcp_loop - if loop is None or not loop.is_running(): - return None - - logger.info( - "MCP server '%s': %s failed with session-expired error (%s); " - "signalling transport reconnect and retrying once.", - server_name, op_description, exc, - ) - - # Trigger the same reconnect mechanism the OAuth recovery path - # uses, then wait briefly for the new session to come back ready. - if not _signal_reconnect_and_wait( - server_name, - srv, - op_description=op_description, - timeout=15, - ): - logger.warning( - "MCP server '%s': reconnect did not ready within 15s after " - "session-expired error; falling through to error response.", - server_name, - ) - return None - - try: - result = retry_call() - try: - parsed = json.loads(result) - if "error" not in parsed: - _reset_server_error(server_name) - return result - except (json.JSONDecodeError, TypeError): - _reset_server_error(server_name) - return result - except Exception as retry_exc: - logger.warning( - "MCP %s/%s retry after session reconnect failed: %s", - server_name, op_description, retry_exc, - ) - return None - - -class _StdioChildExited(RuntimeError): - """A server's stdio subprocess was gone when (or while) a call ran. - - Deliberately NOT a TimeoutError: nothing timed out — the child was - already dead, usually because a gateway restart killed every MCP stdio - subprocess out from under a still-live agent session. The old wording - ("failing the call fast instead of waiting 300s") sent an investigation - into the remote server for an afternoon; the server was healthy. - - Handled by :func:`_handle_stdio_child_exited_and_retry`, which respawns - and retries the call once before any error reaches the model. - """ - - -def _handle_stdio_child_exited_and_retry( - server_name: str, - exc: Exception, - retry_call, - op_description: str, -): - """Respawn a dead stdio child and retry the call once. - - A gateway restart kills every MCP stdio subprocess. An agent session that - outlives the restart still holds the dead child, so its next tool call - used to fail in 0.00s — before anything reached the network — while the - subprocess was respawned seconds later. Cron runs spanning a restart lost - tool calls this way, silently. - - Why retrying here cannot hot-cycle respawns: this function never spawns - anything. It sets ``_reconnect_event`` (one signal, same as before) and - waits for the server task to publish a fresh session. Spawn frequency - stays governed entirely by ``run()``'s rapid-drop budget, which parks a - transport that keeps dropping without proving healthy (#62212). The retry - is single-shot: a child that dies again immediately reports and stops, - so a genuinely broken server converges on the park instead of looping. - - Returns: - A JSON string when this was a dead-stdio failure (retry result, or a - clean error), or ``None`` when ``exc`` is something else and the - caller should use its generic error path. - """ - if not isinstance(exc, _StdioChildExited): - return None - - with _lock: - srv = _servers.get(server_name) - - reconnected = False - if srv is not None and hasattr(srv, "_reconnect_event"): - logger.info( - "MCP server '%s': %s found the stdio subprocess dead (%s); " - "respawning and retrying once.", - server_name, op_description, exc, - ) - loop = _mcp_loop - if loop is not None and loop.is_running(): - reconnected = _signal_reconnect_and_wait( - server_name, - srv, - op_description=op_description, - timeout=_STDIO_RESPAWN_WAIT_SEC, - ) - else: - # No MCP loop to wait on (non-async adapters, tests) — still ask - # for the respawn so the next call lands on a live transport. - _signal_reconnect(srv) - - if reconnected: - try: - result = retry_call() - except _StdioChildExited as retry_exc: - # Respawned and died again straight away: this is a broken - # server, not a restart artifact. Stop here — run()'s budget - # takes it to the park. - logger.warning( - "MCP server '%s': %s stdio subprocess exited again right " - "after respawn (%s); not retrying further.", - server_name, op_description, retry_exc, - ) - _bump_server_error(server_name) - return tool_error( - f"MCP server '{server_name}' respawned its stdio subprocess " - f"and it exited again immediately. The server is not " - f"starting cleanly — do NOT retry this tool; ask the user to " - f"check the server's command and its stderr log." - ) - except Exception as retry_exc: - logger.warning( - "MCP %s/%s retry after stdio respawn failed: %s", - server_name, op_description, retry_exc, - ) - _bump_server_error(server_name) - return tool_error(_sanitize_error( - f"MCP call failed after respawning the stdio subprocess for " - f"'{server_name}': {type(retry_exc).__name__}: " - f"{_exc_str(retry_exc)}" - )) - try: - parsed = json.loads(result) - if "error" not in parsed: - _reset_server_error(server_name) - else: - _bump_server_error(server_name) - except (json.JSONDecodeError, TypeError): - _reset_server_error(server_name) - return result - - _bump_server_error(server_name) - return tool_error( - f"MCP server '{server_name}' stdio subprocess had exited (this is " - f"not a timeout — the call never reached the server). A respawn was " - f"requested but no fresh session came back within " - f"{_STDIO_RESPAWN_WAIT_SEC:.0f}s. Wait a few seconds before retrying; " - f"if it keeps failing the server is not starting and needs the user." - ) - - -# Exact raw server names whose ``supports_parallel_tool_calls`` config is True. -# Raw identity matters: distinct names such as ``foo-bar`` and ``foo_bar`` both -# sanitize to ``foo_bar`` but must not share policy. +# Raw server names opted into parallel tool calls. Raw identity matters: +# ``foo-bar`` and ``foo_bar`` both sanitize to ``foo_bar`` but must not share +# policy. _parallel_safe_servers: set = set() - -# Exact MCP tool-name provenance. The generated registry name is lossy because -# provider-safe normalization maps punctuation to ``_``. Keep the raw server -# name captured at registration time so policy and capability checks never rely -# on parsing or re-sanitizing the generated name. +# registry tool name -> raw server name, captured at registration. The +# generated name is lossy (punctuation -> ``_``), so never re-parse it. _mcp_tool_server_names: Dict[str, str] = {} # Dedicated event loop running in a background daemon thread. _mcp_loop: Optional[asyncio.AbstractEventLoop] = None _mcp_thread: Optional[threading.Thread] = None - -# Protects _mcp_loop, _mcp_thread, _servers, MCP connection status maps, -# _parallel_safe_servers, _mcp_tool_server_names, and _stdio_pids. +# Guards the loop handles, _servers, the status maps and the PID ledgers. _lock = threading.Lock() def _mcp_registry_scope() -> Optional[str]: - """Registry scope owning MCP registrations made from the current context. + """Registry scope for MCP registrations from the current context. - Under a profile multiplexer each profile's MCP tools live in that - profile's registry overlay (the same overlay its plugins use) so two - profiles' servers never share one process-global slot. Single-profile - processes keep MCP tools process-global (``None``). + Under a profile multiplexer each profile's MCP tools live in its own + registry overlay; single-profile processes stay process-global (None). """ from agent.secret_scope import is_multiplex_active @@ -5425,9 +1438,8 @@ def _mcp_registry_scope() -> Optional[str]: def _server_registry_scope(name: str) -> Optional[str]: """Scope owning server *name*'s tools: recorded at connect, else current. - Teardown paths run on the MCP loop (process exit, reconnect exhaustion), - which does not carry the discovering profile's context, so the scope - captured when the server was adopted into ``_servers`` is authoritative. + Teardown runs on the MCP loop without the discovering profile's context, + so the scope captured at adoption into ``_servers`` is authoritative. """ if name in _server_scope_keys: return _server_scope_keys[name] @@ -5435,26 +1447,21 @@ def _server_registry_scope(name: str) -> Optional[str]: # --------------------------------------------------------------------------- -# Cross-process MCP discovery guard +# Cross-process MCP discovery guard: advisory file lock so gateway + CLI + TUI +# don't all run discovery at once. # --------------------------------------------------------------------------- -# Advisory file lock that prevents N concurrent Hermes processes (e.g. -# gateway + CLI + TUI) from all running MCP discovery simultaneously. -# See issue #62771. _LOCK_UNAVAILABLE: Any = object() # sentinel: locking broken/unavailable _MCP_DISCOVERY_LOCK_PATH: Optional[str] = None # resolved lazily - -# Retry constants for the bounded wait when another process holds the lock. +# Bounded wait when another process holds the lock. _MCP_DISCOVERY_LOCK_MAX_RETRIES: int = 240 _MCP_DISCOVERY_LOCK_RETRY_DELAY_S: float = 0.5 class _LockCookie: - """Holds a cross-process file lock; release() drops it. + """Holds a cross-process file lock; ``release()`` drops it. - On Windows the underlying file handle MUST stay alive while the lock is - held (portalocker keeps the kernel lock on the fd). On POSIX the fcntl - lockdown is similarly tied to the file-descriptor lifetime. We keep the - file object in ``_fh`` and close it on release. + 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. """ def __init__(self, fh: Any) -> None: @@ -5486,13 +1493,10 @@ class _LockCookie: def _acquire_lock_on_fh(fh: Any) -> bool: - """Acquire a non-blocking exclusive lock on an open file handle. + """Non-blocking exclusive lock (fcntl on POSIX, portalocker elsewhere). - Uses ``fcntl.flock`` on POSIX and ``portalocker.lock`` on Windows. - - Returns ``True`` if the lock was acquired, ``False`` if another process - holds it (non-blocking refusal). Raises ``RuntimeError`` on unexpected - errors so the caller can treat lock acquisition as unavailable. + False when another process holds it; unexpected errors propagate so the + caller can treat locking as unavailable. """ fd = fh.fileno() if os.name == "posix": @@ -5514,18 +1518,8 @@ def _acquire_lock_on_fh(fh: Any) -> bool: def _try_acquire_mcp_discovery_lock() -> Any: - """Try to acquire an exclusive cross-process lock for MCP discovery. - - Returns - ------- - _LockCookie - Lock acquired successfully. - None - Another process holds the lock (non-blocking refusal). - _LOCK_UNAVAILABLE - Locking mechanism is broken or unavailable -- caller should run - discovery unguarded. - """ + """Return a ``_LockCookie`` (acquired), ``None`` (held by another process) + or ``_LOCK_UNAVAILABLE`` (locking broken: run discovery unguarded).""" global _MCP_DISCOVERY_LOCK_PATH try: from hermes_constants import get_hermes_home @@ -5550,142 +1544,16 @@ def _try_acquire_mcp_discovery_lock() -> Any: if acquired: return _LockCookie(fh) - else: - fh.close() - return None - - -# PIDs of stdio MCP server subprocesses. Tracked so we can force-kill -# them on shutdown if the graceful cleanup (SDK context-manager teardown) -# fails or times out. PIDs are added after connection and removed on -# normal server shutdown. -_stdio_pids: Dict[int, str] = {} # pid -> server_name - -# PIDs that survived their session context exit (SDK teardown failed to -# terminate them). These are detected in _run_stdio's finally block and -# can be cleaned up asynchronously by _kill_orphaned_mcp_children(). -# Separate from _stdio_pids so cleanup sweeps never race with active -# sessions (e.g. concurrent cron jobs or live user chats). -_orphan_stdio_pids: set = set() -_orphan_stdio_pid_servers: Dict[int, str] = {} - -# Process-group IDs of stdio MCP subprocesses, captured at spawn time. -# The MCP SDK spawns stdio children with ``start_new_session=True`` so each -# direct child becomes its own session/pgroup leader (PGID == its own PID). -# Grandchildren spawned by that child (e.g. a wrapper MCP server that itself -# launches helper subprocesses like ``claude mcp serve``) inherit that PGID -# unless they call ``setsid`` themselves. When the direct child exits, those -# grandchildren reparent to init/systemd-user but keep the original PGID, so -# ``killpg(pgid, sig)`` still reaches them. Tracked separately from -# ``_stdio_pids`` so we retain the PGID even after the direct child has -# exited and been removed from the active map. Empty on Windows -# (``os.getpgid`` is POSIX-only). -_stdio_pgids: Dict[int, int] = {} # pid -> pgid - - -def _snapshot_child_pids() -> set: - """Return a set of current child process PIDs. - - Uses /proc on Linux, falls back to psutil, then empty set. - Used by _run_stdio to identify the subprocess spawned by stdio_client. - """ - my_pid = os.getpid() - - # Linux: read from /proc. ``/proc//task//children`` is - # per-THREAD — a child forked from thread T is listed only under T's - # task dir. stdio_client() spawns from the background MCP loop thread, - # so reading only the main thread's file (``task//children``) - # returned an empty set on every Linux install and left - # ``_stdio_child_pids`` / ``_stdio_pids`` empty: the #81995 dead-child - # fast-fail, the #96452 respawn signal, and the killpg shutdown sweep - # never saw the subprocess. Union the children of every task instead. - try: - task_dir = f"/proc/{my_pid}/task" - tids = os.listdir(task_dir) - found: set = set() - for tid in tids: - try: - with open(f"{task_dir}/{tid}/children", encoding="utf-8") as f: - found.update(int(p) for p in f.read().split() if p.strip()) - except (FileNotFoundError, OSError, ValueError): - # Thread exited between listdir and open — skip it. - continue - return found - except (FileNotFoundError, OSError, ValueError): - pass - - # Fallback: psutil - try: - import psutil - return {c.pid for c in psutil.Process(my_pid).children()} - except Exception: - pass - - return set() - - -# Non-MCP gateway children that can race into the _snapshot_child_pids() delta -# during stdio MCP server spawn. LSP servers and slash_worker now use -# start_new_session=True too; this remains defense-in-depth for any future -# non-MCP child spawn that briefly appears in the MCP snapshot delta. Match -# argv markers instead of argv[0] because Python/Java children begin with the -# interpreter or binary path. -_NON_MCP_CHILD_CMDLINE_MARKERS: tuple[str, ...] = ( - "tui_gateway.slash_worker", - "tui_gateway.entry", - "-dorg.eclipse.equinox.launcher", # jdtls (legacy arg style) - "eclipse.jdt.ls", - "org.eclipse.equinox.launcher_", -) - - -def _filter_mcp_children(pids: set) -> set: - """Remove non-MCP children from a PID snapshot delta. - - _snapshot_child_pids() returns *all* direct children of the gateway. When - a stdio MCP server spawns concurrently with a slash_worker or LSP server - spawn, the delta ``_snapshot_child_pids() - pids_before`` can include - PIDs that are NOT the MCP server. Tracking those PIDs in _stdio_pgids is - catastrophic if a future child lacks start_new_session: its pgid can be the - TUI parent's PID, so the shutdown sweep's killpg() kills the TUI itself. - """ - if not pids: - return pids - try: - import psutil - except ImportError: - # psutil unavailable — keep all PIDs (preserves prior behavior). - return pids - filtered: set = set() - for pid in pids: - try: - argv = psutil.Process(pid).cmdline() - except (psutil.NoSuchProcess, psutil.AccessDenied, OSError): - # Process raced away or is a zombie — skip it; it cannot be the - # MCP server we just spawned and is not safe to track. - continue - if any( - marker in arg - for arg in argv[1:] - for marker in _NON_MCP_CHILD_CMDLINE_MARKERS - ): - continue - filtered.add(pid) - return filtered + fh.close() + return None def _mcp_loop_exception_handler(loop, context): - """Suppress benign 'Event loop is closed' noise during shutdown. - - When the MCP event loop is stopped and closed, httpx/httpcore async - transports may fire __del__ finalizers that call call_soon() on the - dead loop. asyncio catches that RuntimeError and routes it here. - We silence it because the connection is being torn down anyway; all - other exceptions are forwarded to the default handler. - """ + """Suppress the benign 'Event loop is closed' RuntimeError that 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 # benign shutdown race — suppress + return loop.default_exception_handler(context) @@ -5706,13 +1574,8 @@ def _ensure_mcp_loop(): def _wrap_with_home_override(coro: "Coroutine") -> "Coroutine": - """Carry the caller's context-local HERMES_HOME override into ``coro``. - - Returns ``coro`` unchanged when no override is active. Otherwise wraps - it so the override is set inside the coroutine's own (task-local) - context on the MCP loop and reset when it completes — concurrent calls - carrying different 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, @@ -5758,15 +1621,11 @@ def _wrap_with_dashboard_oauth_flow(coro): def _run_on_mcp_loop(coro_or_factory, timeout: float = 30): - """Schedule a coroutine on the MCP event loop and block until done. + """Schedule a coroutine on the MCP loop and block until done. - Accepts either a coroutine object or a zero-arg callable that returns one. - Callers can pass a factory to avoid constructing coroutine objects when - the MCP loop is unavailable (which would otherwise leak the coroutine - frame and emit ``"coroutine was never awaited"`` warnings). - - Poll in short intervals so the calling agent thread can honor user - interrupts while the MCP work is still running on the background loop. + 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 @@ -5780,16 +1639,9 @@ def _run_on_mcp_loop(coro_or_factory, timeout: float = 30): coro = coro_or_factory() if callable(coro_or_factory) else coro_or_factory - # Propagate the context-local HERMES_HOME override onto the MCP loop. - # Tasks scheduled via run_coroutine_threadsafe are created INSIDE the - # loop thread, so they copy the loop thread's context — not the - # scheduling thread's. A per-request profile scope (the dashboard's - # ?profile= endpoints, e.g. the MCP "Test server" probe) would silently - # vanish here: OAuth token stores and any other get_hermes_home() - # resolution inside the coroutine would read the process home instead - # of the selected profile's. Re-establish the override inside the - # task's own context (task-local — concurrent calls carrying different - # scopes don't interfere). No-op when no override is active. + # 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 = _wrap_with_home_override(coro) coro = _wrap_with_dashboard_oauth_flow(coro) @@ -5823,224 +1675,40 @@ def _run_on_mcp_loop(coro_or_factory, timeout: float = 30): try: return future.result(timeout=wait_timeout) except concurrent.futures.TimeoutError: - # On supported Python versions, concurrent.futures.TimeoutError - # aliases the built-in TimeoutError, so result(timeout=...) also - # raises it for a coroutine's own timeout. - # Resolve a done future without a timeout to propagate its stored - # outcome, including completion racing with this polling timeout. + # 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 _interrupted_call_result() -> str: - """Standardized JSON error for a user-interrupted MCP tool call.""" - return tool_error("MCP call interrupted: user sent a new message") - - # --------------------------------------------------------------------------- -# Config loading -# --------------------------------------------------------------------------- - -def _interpolate_env_vars(value): - """Recursively resolve ``${VAR}`` placeholders. - - Both ``${VAR}`` and Cursor-style ``${env:VAR}`` are accepted — the - ``env:`` prefix is stripped so a doc copied from a Cursor / Claude MCP - config resolves the same secret. Cursor's context variables are also - supported (case-sensitive): ``${userHome}``, ``${workspaceFolder}``, - ``${workspaceFolderBasename}``, ``${pathSeparator}`` and ``${/}`` — see - :func:`_context_var_value` / :func:`_workspace_folder` for resolution. - Env refs resolve from the active profile's secret scope when multiplexing - is on (so an MCP server config's ``${API_KEY}`` picks up the routed - profile's value, not the process-global ``os.environ`` which may hold - another profile's), falling back to ``os.environ`` otherwise. Unset vars - keep the literal placeholder, as before. - """ - from agent.secret_scope import get_secret as _get_secret - - if isinstance(value, str): - def _replace(m): - ctx = _context_var_value(m.group(1).strip()) - if ctx is not None: - return ctx - name = _env_ref_name(m.group(1)) - return _get_secret(name, m.group(0)) or m.group(0) - return _ENV_VAR_PATTERN.sub(_replace, value) - if isinstance(value, dict): - return {k: _interpolate_env_vars(v) for k, v in value.items()} - if isinstance(value, list): - return [_interpolate_env_vars(v) for v in value] - return value - - -# (server_name, dotted key path) pairs already warned about — see -# _warn_hidden_whitespace(); config loads happen on every discovery pass. -_whitespace_warned: Set[Tuple[str, str]] = set() - - -def _warn_hidden_whitespace(server_name: str, config: dict) -> List[str]: - """Warn about MCP config string values with hidden leading/trailing whitespace. - - A token pasted with a trailing newline or a URL copied with a leading - space produces opaque auth/connect failures (the server rejects the - credential, TLS/DNS fails on ``"example.com "``), and the whitespace is - invisible when eyeballing config.yaml. Inspired by Claude Code v2.1.219, - which added the same startup warning for its MCP config values. - - Advisory only — values are never mutated (whitespace could theoretically - be intentional in an arg). Returns the list of dotted key paths flagged, - for testability. Values themselves are never logged (they are often - secrets); only the key path is named. Each (server, key path) is warned - about once per process — ``_load_mcp_config()`` runs on every discovery/ - status call and repeating the warning would be noise. - """ - flagged: List[str] = [] - - def _walk(value: Any, path: str) -> None: - if isinstance(value, str): - if value != value.strip(): - flagged.append(path) - elif isinstance(value, dict): - for k, v in value.items(): - _walk(v, f"{path}.{k}" if path else str(k)) - elif isinstance(value, list): - for i, v in enumerate(value): - _walk(v, f"{path}[{i}]") - - _walk(config, "") - for key_path in flagged: - dedupe_key = (server_name, key_path) - if dedupe_key in _whitespace_warned: - continue - _whitespace_warned.add(dedupe_key) - logger.warning( - "MCP server '%s': config value '%s' has hidden leading or " - "trailing whitespace — this often causes authentication or " - "connection failures. Check for stray spaces/newlines in " - "config.yaml (or the referenced env var).", - server_name, - key_path, - ) - return flagged - - -def _filter_suspicious_mcp_servers(servers: Dict[str, dict]) -> Dict[str, dict]: - """Drop exfiltration-shaped MCP configs before any stdio spawn path.""" - try: - from hermes_cli.mcp_security import validate_mcp_server_entry as _validate_mcp_server_entry - except Exception: - _validate_mcp_server_entry: Callable[[str, dict[str, Any]], list[str]] | None = None - - if _validate_mcp_server_entry is None: - return servers - - safe_servers = {} - for name, cfg in servers.items(): - if not isinstance(cfg, dict): - safe_servers[name] = cfg - continue - issues = _validate_mcp_server_entry(name, cfg) - if issues: - logger.warning( - "Skipping suspicious MCP server '%s': %s", - name, - "; ".join(issues), - ) - continue - safe_servers[name] = cfg - return safe_servers - - -def _load_mcp_config() -> Dict[str, dict]: - """Read ``mcp_servers`` from the Hermes config file. - - Returns a dict of ``{server_name: server_config}`` or empty dict. - Server config can contain either ``command``/``args``/``env`` for stdio - transport or ``url``/``headers`` for HTTP transport, plus optional - ``timeout``, ``connect_timeout``, and ``auth`` overrides. - - ``${ENV_VAR}`` placeholders in string values are resolved from - ``os.environ`` (which includes ``~/.hermes/.env`` loaded at startup). - """ - try: - from hermes_cli.config import load_config - from utils import env_var_enabled as _env_enabled - - if _env_enabled("HERMES_SAFE_MODE"): - return {} - config = load_config() - servers = config.get("mcp_servers") - if not isinstance(servers, dict): - servers = {} - # Ensure .env vars are available for interpolation - try: - from hermes_cli.env_loader import load_hermes_dotenv - load_hermes_dotenv() - except Exception: - pass - safe_servers: Dict[str, dict] = {} - for name, cfg in _filter_suspicious_mcp_servers(servers).items(): - interpolated = _interpolate_env_vars(cfg) - if isinstance(interpolated, dict): - _warn_hidden_whitespace(name, interpolated) - safe_servers[name] = interpolated - try: - from hermes_cli.plugins import discover_plugins, get_plugin_manager - - discover_plugins() - portable = get_plugin_manager().get_portable_mcp_servers() - for name, cfg in _filter_suspicious_mcp_servers(portable).items(): - if name in safe_servers: - logger.warning( - "Portable MCP server '%s' conflicts with native config; skipping", - name, - ) - continue - safe_servers[name] = dict(cfg) - except Exception: - logger.debug("Failed to load portable MCP servers", exc_info=True) - return safe_servers - except Exception as exc: - logger.debug("Failed to load MCP config: %s", exc) - return {} - - -# --------------------------------------------------------------------------- -# Server connection helper +# Connecting, lazy start, discovery # --------------------------------------------------------------------------- async def _connect_server(name: str, config: dict) -> MCPServerTask: - """Create an MCPServerTask, start it, and return when ready. + """Create an MCPServerTask, start it and return once ready. - The server Task keeps the connection alive in the background. - Call ``server.shutdown()`` (on the same event loop) to tear it down. - - Raises: - ValueError: if required config keys are missing. - ImportError: if HTTP transport is needed but not available. - Exception: on connection or initialization failure. + Tear it down with ``server.shutdown()`` on the same loop. Raises on bad + config, missing HTTP support, or connect/initialize failure. """ server = MCPServerTask(name) claim = _connect_server_claim.get() claim_token = None if claim is not None: claim(server) - # ``start()`` creates the long-lived run task by copying this context. - # The ownership callback is only for this connection attempt; do not - # retain its discovery closure for the server's lifetime. + # The run task copies this context; the claim is for this attempt + # only, so don't retain the discovery closure for the server's life. claim_token = _connect_server_claim.set(None) try: await server.start(config) except asyncio.CancelledError: - # start() already cancels/reaps server._task on external cancellation - # (see the comment there) -- awaiting a redundant shutdown() inside a - # cancelled context would only risk swallowing the cancellation. + # start() already reaps server._task; a shutdown() here could + # swallow the cancellation. raise except BaseException: - # Discovery owns claimed tasks and decides whether a failed start is a - # live recoverable park or a terminal failure. Standalone probes have - # no revival owner, so they must reap their failed task locally. + # Discovery owns claimed tasks (recoverable park vs terminal failure); + # standalone probes have no revival owner and must reap locally. if claim is None: try: await server.shutdown() @@ -6056,10 +1724,6 @@ async def _connect_server(name: str, config: dict) -> MCPServerTask: return server -# --------------------------------------------------------------------------- -# Handler / check-fn factories -# --------------------------------------------------------------------------- - def _request_lazy_reconnect(server_name: str, server: MCPServerTask) -> bool: """Wake a recycled stdio server and wait briefly for a fresh session.""" if not server._is_recycled_stdio(): @@ -6095,23 +1759,16 @@ def _request_lazy_reconnect(server_name: str, server: MCPServerTask) -> bool: def _resolve_server_lazy(name: str, config: dict) -> bool: - """True when this server defers spawn/connect until first tool use. - - Gated per-server by ``mcp_servers..lazy`` in config (default OFF), - following the same per-server key pattern as ``idle_timeout_seconds``. - Design from #56832 (Vansh5632). - """ + """True when ``mcp_servers..lazy`` defers connect to first tool use (default off).""" return _parse_boolish(config.get("lazy", False), default=False) def _ensure_lazy_server_connected(server_name: str) -> bool: - """Connect a lazily-registered MCP server on demand (sync, blocks caller). + """Connect a lazily-registered server on demand (sync; blocks the caller). - Composes with the existing connect machinery: respects the per-server - connect cooldown (#50394), the ``_server_connecting`` dedup set, and - routes through ``_discover_and_register_server`` so parked/recycle/ - cooldown bookkeeping stays in one place. Returns True when a live - session is available afterwards. + Honours the connect cooldown and the ``_server_connecting`` dedup set and + routes through ``_discover_and_register_server`` so park/recycle/cooldown + bookkeeping stays in one place. True when a live session exists after. """ with _lock: server = _servers.get(server_name) @@ -6157,9 +1814,8 @@ def _ensure_lazy_server_connected(server_name: str) -> bool: live_names = set( getattr(server, "_registered_tool_names", []) or [] ) - # Stale-cache reconciliation: the cached manifest may advertise tools - # the live server no longer serves. Deregister those phantoms so the - # model stops seeing tools that can never succeed. + # The cached manifest may advertise tools the live server no longer + # serves; deregister those phantoms. phantom_names = [n for n in cached_names if n not in live_names] if phantom_names: from tools.registry import registry @@ -6177,12 +1833,8 @@ def _ensure_lazy_server_connected(server_name: str) -> bool: def _get_connected_server_for_call(server_name: str) -> Optional[MCPServerTask]: - """Return a connected server, lazily reconnecting recycled stdio state. - - Also the single first-use connect point for lazy (schema-cache - registered) servers, so raw tool calls AND the resource/prompt utility - handlers all trigger the deferred spawn (#56832). - """ + """Return a connected server; the single first-use connect point for lazy + servers and the wake-up point for recycled stdio ones.""" with _lock: server = _servers.get(server_name) is_lazy = server_name in _lazy_server_configs @@ -6198,1585 +1850,11 @@ def _get_connected_server_for_call(server_name: str) -> Optional[MCPServerTask]: return server -def _mark_server_call_started(server: Any) -> None: - """Record a user-visible MCP operation when the server supports it.""" - mark_tool_call = getattr(server, "mark_tool_call", None) - if callable(mark_tool_call): - mark_tool_call() - - -@asynccontextmanager -async def _track_inflight_rpc(server: Any, server_name: str, op: str): - """Register the running RPC on the server so teardown can fail it fast. - - Every user-visible request family wraps its RPC in this context - (#48069 salvage). If a deliberate reconnect/shutdown teardown cancels - the task (``_fail_inflight_calls`` sets ``_reconnecting`` first), the - cancel is converted into a clean retryable RuntimeError instead of a raw - CancelledError; external cancels (caller timeout, user interrupt) - propagate unchanged. - """ - inflight = getattr(server, "_inflight_tasks", None) - task = asyncio.current_task() - if task is not None and inflight is not None: - # Test doubles may pass a bare SimpleNamespace; tracking is then - # simply skipped (fast-fail teardown is a production-connection - # feature, not something a fake needs). - inflight.add(task) - try: - yield - except asyncio.CancelledError: - if getattr(server, "_reconnecting", False): - raise RuntimeError( - f"MCP {op} on '{server_name}' was aborted by a reconnect " - f"teardown; retry the request on the rebuilt session" - ) from None - raise - finally: - if task is not None and inflight is not None: - inflight.discard(task) - - -def _ensure_healthy_or_recycle(server: Any, server_name: str) -> None: - """Health-check a suspect connection before its next call (#85125 3b). - - Implements the SuspectableBackend cheap-mark/lazy-verify contract at the - dispatch boundary: a connection latched as suspect by a race or an auth - error is probed once; a failed probe recycles it so the call below hits - the normal reconnect path. A HEALTHY connection is never recycled here. - """ - if not getattr(server, "_suspect_reason", None): - return - with _lock: - loop = _mcp_loop - if loop is None or not loop.is_running(): - return # no background loop — nothing to verify against - try: - healthy = bool(_run_on_mcp_loop(server.ensure_healthy, timeout=15.0)) - except Exception as exc: # never let the probe break dispatch - logger.debug( - "MCP server '%s': suspect health check errored: %s", - server_name, exc, - ) - healthy = False - if not healthy: - _signal_reconnect(server) - - -def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float): - """Return a sync handler that calls an MCP tool via the background loop. - - The handler conforms to the registry's dispatch interface: - ``handler(args_dict, **kwargs) -> str`` - """ - - def _handler(args: dict, **kwargs) -> str: - # Trust-tier gate (security boundary): write-capable tools on - # servers configured ``trust: untrusted`` must be approved by the - # user before ANY transport work happens — including the lazy - # first-use spawn below. A denied call never touches the server. - gate_error = _trust_gate_check(server_name, tool_name) - if gate_error is not None: - return gate_error - - # Circuit breaker: if this server has failed too many times - # consecutively, short-circuit with a clear message so the model - # stops retrying and uses alternative approaches (#10447). - # - # Once the cooldown elapses, the breaker transitions to - # half-open: we let the *next* call through as a probe. On - # success the success-path below resets the breaker; on - # failure the error paths below bump the count again, which - # re-stamps the open-time via _bump_server_error (re-arming - # the cooldown). - if _server_error_counts.get(server_name, 0) >= _CIRCUIT_BREAKER_THRESHOLD: - opened_at = _server_breaker_opened_at.get(server_name, 0.0) - age = time.monotonic() - opened_at - if age < _CIRCUIT_BREAKER_COOLDOWN_SEC: - remaining = max(1, int(_CIRCUIT_BREAKER_COOLDOWN_SEC - age)) - return tool_error( - f"MCP server '{server_name}' is unreachable after " - f"{_server_error_counts[server_name]} consecutive " - f"failures. Auto-retry available in ~{remaining}s. " - f"Do NOT retry this tool yet — use alternative " - f"approaches or ask the user to check the MCP server." - ) - # Cooldown elapsed → fall through as a half-open probe. - - server = _get_connected_server_for_call(server_name) - if not server: - _bump_server_error(server_name) - return tool_error(f"MCP server '{server_name}' is not connected") - - if not server.session: - # No live session. A reconnect may already be completing (the - # transport swaps in a fresh session object asynchronously) — - # wait briefly before treating this as a failure, so a - # transient reconnect window doesn't burn a circuit-breaker - # strike (#26892). - if _wait_for_server_session_ready( - server, timeout=min(5.0, float(tool_timeout or 5.0)), - ): - pass # Fresh session arrived; proceed below. - else: - # Still down — the server task is reconnecting, or it has - # exhausted its retry budget and parked (e.g. a dead stdio - # subprocess). Probing here would write into a dead/absent - # transport and re-arm the breaker forever (#16788). Instead, - # ask the (always-present) server task to rebuild the - # transport — which respawns a dead stdio subprocess — and - # return a clean "reconnecting" error so the model backs off - # without burning iterations. The breaker resets once the - # fresh session initializes (_run_stdio/_run_http call - # _reset_server_error). - _bump_server_error(server_name) - if _signal_reconnect(server): - return tool_error( - f"MCP server '{server_name}' transport is down; " - f"reconnect requested. Do NOT retry this tool " - f"immediately — give it a few seconds to come back." - ) - return tool_error(f"MCP server '{server_name}' is not connected") - - async def _call(): - _mark_server_call_started(server) - async with server._rpc_lock, _track_inflight_rpc( - server, server_name, f"tools/call {tool_name}" - ): - # Snapshot the agent's context so an elicitation callback - # triggered during this call (fired on the MCP recv loop - # task, which doesn't inherit our contextvars) can replay - # it and detect the gateway platform / session for routing. - server._pending_call_context = contextvars.copy_context() - try: - # Fast-fail (#81995): a stdio subprocess that is already - # dead must not own this call slot — fail immediately - # instead of waiting out the full tool timeout on a - # transport nobody will ever answer. - _stdio_dead = getattr(server, "_stdio_children_dead", None) - # callable() + real-bool result: MagicMock attributes return - # truthy Mocks, which would spuriously trip the fast-fail. - if ( - callable(_stdio_dead) - and isinstance(_stdio_dead_result := _stdio_dead(), bool) - and _stdio_dead_result - ): - # Dead children but stale server.session, so the - # transport-down path above never fired. Hand this to - # the handler's respawn-and-retry path — - # it is not a timeout, and a gateway restart that - # killed the child must not cost the caller a call. - raise _StdioChildExited( - f"MCP stdio subprocess for '{server_name}' had " - f"already exited when the call was dispatched" - ) - _call_coro = server.session.call_tool(tool_name, arguments=args) - _watch_children = getattr(server, "_watch_stdio_children", None) - _watch_ok = ( - _watch_children is not None - and inspect.iscoroutinefunction(_watch_children) - and asyncio.iscoroutine(_call_coro) - ) - if not _watch_ok: - # Stubbed sessions (MagicMock in tests) return a - # non-awaitable, or there is no child-watcher to race - # against: plain await is exactly the pre-#81995 - # semantics. - result = ( - await _call_coro - if asyncio.iscoroutine(_call_coro) - else _call_coro - ) - else: - # Fast-fail machinery (#81995): the RPC races a - # stdio-children watcher so a dead subprocess fails - # the call immediately instead of riding out the full - # tool timeout. - rpc_task = asyncio.ensure_future(_call_coro) - watch_task = asyncio.ensure_future(_watch_children()) - try: - done, _pending = await asyncio.wait( - {rpc_task, watch_task}, - return_when=asyncio.FIRST_COMPLETED, - ) - if watch_task in done and not rpc_task.done(): - rpc_task.cancel() - # Same stale-session problem as the pre-call - # gate above: the subprocess died mid-call but - # nothing clears server.session, so without a - # reconnect the server would stay dead until - # the idle keepalive probe notices. The - # handler's respawn-and-retry path owns the - # reconnect signal. - raise _StdioChildExited( - f"MCP stdio subprocess for " - f"'{server_name}' exited mid-call" - ) - result = await rpc_task - finally: - watch_task.cancel() - if not rpc_task.done(): - rpc_task.cancel() - await asyncio.gather( - rpc_task, watch_task, return_exceptions=True - ) - finally: - server._pending_call_context = None - # The RPC round-trip completed — the session is demonstrably - # healthy at the transport level (even if the tool itself - # returned isError). Clear the rapid-drop budget (#62212). - _mark_proven = getattr(server, "_mark_session_proven", None) - if _mark_proven is not None: - _mark_proven() - # MCP CallToolResult has .content (list of content blocks) and - # .is_error (.isError before mcp 2.0) - if mcp_field(result, "is_error", "isError", False): - error_text = "" - for block in (result.content or []): - if getattr(block, "text", None): - error_text += block.text - continue - # EmbeddedResource blocks inside error payloads carry - # their text under .resource.text — previously dropped, - # leaving a bare "MCP tool returned an error". - res_text = getattr(getattr(block, "resource", None), "text", None) - if res_text: - error_text += str(res_text) - return tool_error(_sanitize_error( - _truncate_mcp_text_result( - error_text or "MCP tool returned an error" - ) - )) - - # Collect text from content blocks. MCP tool results can also - # include ImageContent blocks (screenshot / Blockbench / Playwright - # etc.); cache those via the gateway's image-cache helper so they - # flow through Hermes' MEDIA: tag convention and out to messaging - # adapters that render images natively. Without this, image blocks - # were silently dropped and the agent got an empty response. - # - # Distilled from #17915 (c3115644151) and #10848 (gnanirahulnutakki), - # both too stale to cherry-pick. #10848's approach (integrate with - # Hermes' MEDIA tag + cache_image_from_bytes) was the cleaner of - # the two — plugs into existing infrastructure. - parts: List[str] = [] - for block in (result.content or []): - if hasattr(block, "text") and block.text: - parts.append(strip_unicode_tags(block.text)) - continue - image_tag = _cache_mcp_image_block(block) - if image_tag: - parts.append(image_tag) - continue - audio_tag = _cache_mcp_audio_block(block) - if audio_tag: - parts.append(audio_tag) - continue - # ResourceLink / EmbeddedResource blocks (PDFs, archives, - # office docs, ...). Previously these were silently dropped, - # so document-oriented MCP tools appeared to return metadata - # only (enterprise customer report, 2026-07). - resource_text = _render_mcp_resource_block(block, server_name) - if resource_text: - parts.append(resource_text) - continue - # Benign empty renders (empty text blocks, empty text - # resources, audio in a process without the gateway cache) - # aren't data loss — log at debug. Warn only for genuinely - # unrecognized block shapes. - block_type = getattr(block, "type", None) or type(block).__name__ - if block_type in {"text", "resource", "audio", "image"}: - logger.debug( - "MCP %s: content block type %r rendered empty", - server_name, block_type, - ) - else: - logger.warning( - "MCP %s: dropping unsupported content block type %r", - server_name, block_type, - ) - text_result = "\n".join(parts) if parts else "" - - # Hard-cap pathological payloads before they propagate (#56059); - # ordinary large results pass untouched to the spillover layer. - text_result = _truncate_mcp_text_result(text_result) - - # Combine content + structuredContent when both are present. - # MCP spec: content is model-oriented (text), structuredContent - # is machine-oriented (JSON metadata). For an AI agent, content - # is the primary payload; structuredContent supplements it. - # - # Server-level `_meta` is also surfaced (ported from - # MoonshotAI/kimi-code#2596): servers return namespaced metadata - # there (validated contracts, browser-handoff payloads, ...) that - # was previously invisible to the agent. Protocol-reserved keys - # are dropped first (kimi-code#2600) — per the MCP spec's key-name - # rules a prefix is reserved when a `modelcontextprotocol` or - # `mcp` label is followed by at least one more label (e.g. - # `modelcontextprotocol.io/...`, `tools.mcp.com/...`); those carry - # host/protocol plumbing, not model-facing data. Unprefixed and - # vendor-namespaced keys (`com.example.mcp/...`) pass through — - # their semantics belong to the server. - structured = mcp_field(result, "structured_content", "structuredContent") - # Cap structuredContent too — a malicious server could flood - # context via a multi-MB JSON payload (#56059). When the - # serialized form exceeds the hard cap, replace it with the - # truncated string (head + tail preserved) so it degrades - # gracefully instead of flooding downstream. - if structured is not None: - try: - _structured_json = json.dumps(structured, ensure_ascii=False, default=str) - except (TypeError, ValueError): - _structured_json = None - if _structured_json is not None and len(_structured_json) > _MCP_HARD_RESULT_CAP_CHARS: - structured = _truncate_mcp_text_result(_structured_json) - meta = _strip_reserved_meta_keys(mcp_field(result, "meta", "meta")) - if structured is not None or meta is not None: - payload: Dict[str, Any] = {} - if text_result: - payload["result"] = text_result - if structured is not None: - if text_result: - payload["structuredContent"] = structured - else: - payload["result"] = structured - if meta is not None: - payload["_meta"] = meta - if "result" not in payload: - payload["result"] = text_result - try: - return json.dumps(payload, ensure_ascii=False) - except (TypeError, ValueError): - # Non-serializable metadata: drop the extras rather than - # failing the whole tool call. - return json.dumps({"result": text_result}, ensure_ascii=False) - return json.dumps({"result": text_result}, ensure_ascii=False) - - def _call_once(): - return _run_on_mcp_loop(_call, timeout=tool_timeout) - - try: - result = _call_once() - # Check if the MCP tool itself returned an error - try: - parsed = json.loads(result) - if "error" in parsed: - _bump_server_error(server_name) - else: - _reset_server_error(server_name) # success — reset - except (json.JSONDecodeError, TypeError): - _reset_server_error(server_name) # non-JSON = success - return result - except InterruptedError: - return _interrupted_call_result() - except Exception as exc: - # Dead stdio child: respawn and retry once before any - # error reaches the model — a gateway restart kills every MCP - # subprocess, and the call it lands on is not really a failure. - recovered = _handle_stdio_child_exited_and_retry( - server_name, exc, _call_once, - f"tools/call {tool_name}", - ) - if recovered is not None: - return recovered - - # Auth-specific recovery path: consult the manager, signal - # reconnect if viable, retry once. Returns None to fall - # through for non-auth exceptions. - recovered = _handle_auth_error_and_retry( - server_name, exc, _call_once, - f"tools/call {tool_name}", - ) - if recovered is not None: - return recovered - - # Transport session expiry (#13383): same reconnect flow - # but skips OAuth recovery because the access token is - # still valid — only the server-side session is stale. - recovered = _handle_session_expired_and_retry( - server_name, exc, _call_once, - f"tools/call {tool_name}", - ) - if recovered is not None: - return recovered - - _bump_server_error(server_name) - logger.error( - "MCP tool %s/%s call failed: %s", - server_name, tool_name, exc, - ) - return tool_error(_sanitize_error( - f"MCP call failed: {type(exc).__name__}: {_exc_str(exc)}" - )) - - return _handler - - -def _make_list_resources_handler(server_name: str, tool_timeout: float): - """Return a sync handler that lists resources from an MCP server.""" - - def _handler(args: dict, **kwargs) -> str: - server = _get_connected_server_for_call(server_name) - if not server or not server.session: - return tool_error(f"MCP server '{server_name}' is not connected") - - async def _call(): - _mark_server_call_started(server) - async with server._rpc_lock: - all_resources = await _paginate_full_list( - server.session.list_resources, "resources", server_name - ) - resources = [] - for r in all_resources: - entry = {} - if hasattr(r, "uri"): - entry["uri"] = str(r.uri) - if hasattr(r, "name"): - entry["name"] = r.name - if hasattr(r, "description") and r.description: - entry["description"] = r.description - # Key stays camelCase — this dict is the tool's own JSON - # output shape, not an SDK model. - _mime = mcp_field(r, "mime_type", "mimeType") - if _mime: - entry["mimeType"] = _mime - resources.append(entry) - return json.dumps({"resources": resources}, ensure_ascii=False) - - def _call_once(): - return _run_on_mcp_loop(_call, timeout=tool_timeout) - - try: - return _call_once() - except InterruptedError: - return _interrupted_call_result() - except Exception as exc: - recovered = _handle_auth_error_and_retry( - server_name, exc, _call_once, "resources/list", - ) - if recovered is not None: - return recovered - recovered = _handle_session_expired_and_retry( - server_name, exc, _call_once, "resources/list", - ) - if recovered is not None: - return recovered - logger.error( - "MCP %s/list_resources failed: %s", server_name, exc, - ) - return tool_error(_sanitize_error( - f"MCP call failed: {type(exc).__name__}: {_exc_str(exc)}" - )) - - return _handler - - -def _make_read_resource_handler(server_name: str, tool_timeout: float): - """Return a sync handler that reads a resource by URI from an MCP server.""" - - def _handler(args: dict, **kwargs) -> str: - server = _get_connected_server_for_call(server_name) - if not server or not server.session: - return tool_error(f"MCP server '{server_name}' is not connected") - - uri = args.get("uri") - if not uri: - return tool_error("Missing required parameter 'uri'") - - async def _call(): - _mark_server_call_started(server) - async with server._rpc_lock: - result = await server.session.read_resource(uri) - # read_resource returns ReadResourceResult with .contents list - parts: List[str] = [] - contents = result.contents if hasattr(result, "contents") else [] - for block in contents: - if getattr(block, "text", None) is not None: - parts.append(strip_unicode_tags(block.text)) - elif getattr(block, "blob", None) is not None: - # Materialize binary resource contents into the document - # cache instead of discarding them (same contract as - # EmbeddedResource blocks in tool results). - rendered = _render_mcp_resource_block( - SimpleNamespace(type="resource", resource=block), - server_name, - ) - parts.append(rendered or f"[binary data, {len(block.blob)} bytes]") - return json.dumps({"result": "\n".join(parts) if parts else ""}, ensure_ascii=False) - - def _call_once(): - return _run_on_mcp_loop(_call, timeout=tool_timeout) - - try: - return _call_once() - except InterruptedError: - return _interrupted_call_result() - except Exception as exc: - recovered = _handle_auth_error_and_retry( - server_name, exc, _call_once, "resources/read", - ) - if recovered is not None: - return recovered - recovered = _handle_session_expired_and_retry( - server_name, exc, _call_once, "resources/read", - ) - if recovered is not None: - return recovered - logger.error( - "MCP %s/read_resource failed: %s", server_name, exc, - ) - return tool_error(_sanitize_error( - f"MCP call failed: {type(exc).__name__}: {_exc_str(exc)}" - )) - - return _handler - - -def _make_list_prompts_handler(server_name: str, tool_timeout: float): - """Return a sync handler that lists prompts from an MCP server.""" - - def _handler(args: dict, **kwargs) -> str: - server = _get_connected_server_for_call(server_name) - if not server or not server.session: - return tool_error(f"MCP server '{server_name}' is not connected") - - async def _call(): - _mark_server_call_started(server) - async with server._rpc_lock: - all_prompts = await _paginate_full_list( - server.session.list_prompts, "prompts", server_name - ) - prompts = [] - for p in all_prompts: - entry = {} - if hasattr(p, "name"): - entry["name"] = p.name - if hasattr(p, "description") and p.description: - entry["description"] = p.description - if hasattr(p, "arguments") and p.arguments: - entry["arguments"] = [ - { - "name": a.name, - **({"description": a.description} if hasattr(a, "description") and a.description else {}), - **({"required": a.required} if hasattr(a, "required") else {}), - } - for a in p.arguments - ] - prompts.append(entry) - return json.dumps({"prompts": prompts}, ensure_ascii=False) - - def _call_once(): - return _run_on_mcp_loop(_call, timeout=tool_timeout) - - try: - return _call_once() - except InterruptedError: - return _interrupted_call_result() - except Exception as exc: - recovered = _handle_auth_error_and_retry( - server_name, exc, _call_once, "prompts/list", - ) - if recovered is not None: - return recovered - recovered = _handle_session_expired_and_retry( - server_name, exc, _call_once, "prompts/list", - ) - if recovered is not None: - return recovered - logger.error( - "MCP %s/list_prompts failed: %s", server_name, exc, - ) - return tool_error(_sanitize_error( - f"MCP call failed: {type(exc).__name__}: {_exc_str(exc)}" - )) - - return _handler - - -def _make_get_prompt_handler(server_name: str, tool_timeout: float): - """Return a sync handler that gets a prompt by name from an MCP server.""" - - def _handler(args: dict, **kwargs) -> str: - server = _get_connected_server_for_call(server_name) - if not server or not server.session: - return tool_error(f"MCP server '{server_name}' is not connected") - - name = args.get("name") - if not name: - return tool_error("Missing required parameter 'name'") - arguments = args.get("arguments", {}) - - async def _call(): - _mark_server_call_started(server) - async with server._rpc_lock: - result = await server.session.get_prompt(name, arguments=arguments) - # GetPromptResult has .messages list - messages = [] - for msg in (result.messages if hasattr(result, "messages") else []): - entry = {} - if hasattr(msg, "role"): - entry["role"] = msg.role - if hasattr(msg, "content"): - content = msg.content - if hasattr(content, "text"): - entry["content"] = strip_unicode_tags(content.text) - elif isinstance(content, str): - entry["content"] = strip_unicode_tags(content) - else: - entry["content"] = strip_unicode_tags(str(content)) - messages.append(entry) - resp = {"messages": messages} - if hasattr(result, "description") and result.description: - resp["description"] = result.description - return json.dumps(resp, ensure_ascii=False) - - def _call_once(): - return _run_on_mcp_loop(_call, timeout=tool_timeout) - - try: - return _call_once() - except InterruptedError: - return _interrupted_call_result() - except Exception as exc: - recovered = _handle_auth_error_and_retry( - server_name, exc, _call_once, "prompts/get", - ) - if recovered is not None: - return recovered - recovered = _handle_session_expired_and_retry( - server_name, exc, _call_once, "prompts/get", - ) - if recovered is not None: - return recovered - logger.error( - "MCP %s/get_prompt failed: %s", server_name, exc, - ) - return tool_error(_sanitize_error( - f"MCP call failed: {type(exc).__name__}: {_exc_str(exc)}" - )) - - return _handler - - -def _make_check_fn(server_name: str): - """Return a check function that verifies the MCP connection is alive.""" - - def _check() -> bool: - with _lock: - server = _servers.get(server_name) - if server is not None and ( - server.session is not None or server._is_recycled_stdio() - ): - return True - # Lazy (schema-cache registered) servers are available: the - # first real call spawns/connects them (#56832). - return server_name in _lazy_server_configs - - return _check - - -# --------------------------------------------------------------------------- -# Discovery & registration -# --------------------------------------------------------------------------- - -def _normalize_mcp_input_schema(schema: dict | None) -> dict: - """Normalize MCP input schemas for LLM tool-calling compatibility. - - MCP servers can emit plain JSON Schema with ``definitions`` / - ``#/definitions/...`` references. Kimi / Moonshot rejects that form and - requires local refs to point into ``#/$defs/...`` instead. Normalize the - common draft-07 shape here so MCP tool schemas remain portable across - OpenAI-compatible providers. - - Additional MCP-server robustness repairs applied recursively: - - * Missing or ``null`` ``type`` on an object-shaped node is coerced to - ``"object"`` (some servers omit it). See PR #4897. - * When an ``object`` node lacks ``properties``, an empty ``properties`` - dict is added so ``required`` entries don't dangle. - * ``required`` arrays are pruned to only names that exist in - ``properties``; otherwise Google AI Studio / Gemini 400s with - ``property is not defined``. See PR #4651. - * MCP/Pydantic optional fields commonly arrive as - ``anyOf: [{...}, {"type": "null"}], default: null``. Anthropic rejects - nullable branches in tool input schemas, so nullable unions are collapsed - to the non-null branch and optionality remains represented solely by the - parent object's ``required`` list. - - All repairs are provider-agnostic and ideally produce a schema valid on - OpenAI, Anthropic, Gemini, and Moonshot in one pass. - """ - if not schema: - return {"type": "object", "properties": {}} - - def _rewrite_local_refs(node): - """Walk the schema, promoting legacy ``definitions`` to ``$defs``. - - The promotion is contextual: ``definitions`` is renamed only when it - appears as a JSON Schema *meta-keyword* (sibling of ``properties`` / - ``$ref`` at a schema node), never when it appears as the *name of a - property* (i.e., as a key inside a ``properties`` dict). - - Without this gate, MCP servers that legitimately expose a tool - parameter named ``definitions`` (e.g. a CI/pipelines tool that uses - ``definitions`` for an array of pipeline-definition IDs) would have - that user-facing property name silently rewritten to ``$defs``. - Anthropic and OpenAI both reject ``$`` in property names - (``^[a-zA-Z0-9_.-]{1,64}$``), so the whole tool array gets a 400 and - every conversation breaks. - - The gate works by treating ``properties`` and ``patternProperties`` - specially during descent: we iterate the property-name -> schema map - directly, leaving the property names verbatim, then recurse into each - property's schema where ordinary JSON Schema semantics resume (so any - legitimately-nested ``definitions`` meta-keyword inside a property's - schema is still promoted). - """ - if isinstance(node, dict): - normalized = {} - for key, value in node.items(): - if key in ("properties", "patternProperties") and isinstance(value, dict): - # Keys of this dict are user-facing property names, not - # meta-keywords. Preserve them verbatim; recurse only into - # each property's schema, where ``definitions`` again has - # its JSON Schema meaning. - normalized[key] = { - prop_name: _rewrite_local_refs(prop_schema) - for prop_name, prop_schema in value.items() - } - else: - out_key = "$defs" if key == "definitions" else key - normalized[out_key] = _rewrite_local_refs(value) - ref = normalized.get("$ref") - if isinstance(ref, str) and ref.startswith("#/definitions/"): - normalized["$ref"] = "#/$defs/" + ref[len("#/definitions/"):] - return normalized - if isinstance(node, list): - return [_rewrite_local_refs(item) for item in node] - return node - - def _strip_nullable_union(node): - """Collapse JSON Schema nullable unions to provider-safe non-null schemas. - - Delegates to ``tools.schema_sanitizer.strip_nullable_unions`` so MCP - ingestion, the Anthropic guard, and the global sanitizer all share one - implementation. Keeps the ``nullable: true`` hint so runtime argument - coercion can still map a model-emitted ``"null"`` string to Python - ``None`` for this optional field. - """ - from tools.schema_sanitizer import strip_nullable_unions - - return strip_nullable_unions(node, keep_nullable_hint=True) - - def _collapse_const_unions(node): - """Collapse anyOf/oneOf unions of same-typed consts to property enums. - - Delegates to ``tools.schema_sanitizer.collapse_const_unions``. Runs - AFTER the nullable strip: single-non-null unions are already collapsed - by then, and unions of several const branches plus a null branch are - handled here (consts -> enum, null -> ``nullable: true`` hint). - Ported from block/goose tool_schema_normalize.rs (Apache-2.0). - """ - from tools.schema_sanitizer import collapse_const_unions - - return collapse_const_unions(node) - - def _repair_object_shape(node): - """Recursively repair object-shaped nodes: fill type, prune required.""" - if isinstance(node, list): - return [_repair_object_shape(item) for item in node] - if not isinstance(node, dict): - return node - - repaired = {k: _repair_object_shape(v) for k, v in node.items()} - - # Coerce missing / null type when the shape is clearly an object - # (has properties or required but no type). - if not repaired.get("type") and ( - "properties" in repaired or "required" in repaired - ): - repaired["type"] = "object" - - if repaired.get("type") == "object": - # Ensure properties exists so required can reference it safely - if "properties" not in repaired or not isinstance( - repaired.get("properties"), dict - ): - repaired["properties"] = {} if "properties" not in repaired else repaired["properties"] - if not isinstance(repaired.get("properties"), dict): - repaired["properties"] = {} - - # Prune required to only include names that exist in properties - required = repaired.get("required") - if isinstance(required, list): - props = repaired.get("properties") or {} - valid = [r for r in required if isinstance(r, str) and r in props] - if len(valid) != len(required): - if valid: - repaired["required"] = valid - else: - repaired.pop("required", None) - - return repaired - - normalized = _rewrite_local_refs(schema) - normalized = _strip_nullable_union(normalized) - normalized = _collapse_const_unions(normalized) - normalized = _repair_object_shape(normalized) - - # Ensure top-level is a well-formed object schema - if not isinstance(normalized, dict): - return {"type": "object", "properties": {}} - if normalized.get("type") == "object" and "properties" not in normalized: - normalized = {**normalized, "properties": {}} - - return normalized - - -def sanitize_mcp_name_component(value: str) -> str: - """Return an MCP name component safe for tool and prefix generation. - - Preserves Hermes's historical behavior of converting hyphens to - underscores, and also replaces any other character outside - ``[A-Za-z0-9_]`` with ``_`` so generated tool names are compatible with - provider validation rules. - """ - return re.sub(r"[^A-Za-z0-9_]", "_", str(value or "")) - - -# Native MCP tool-name prefix. Hermes uses the ``mcp____`` -# convention shared by Claude Code, Codex, and OpenCode (anomalyco/opencode -# #33533). The double-underscore delimiter disambiguates the server/tool -# boundary even when either component contains underscores, and matches the -# naming models are trained on. It also aligns native registration with the -# Anthropic-OAuth wire form (``_MCP_TOOL_PREFIX`` in anthropic_adapter.py), -# removing the single->double rewrite that path previously had to perform. -MCP_TOOL_NAME_PREFIX = "mcp__" -_MCP_NAME_DELIM = "__" - - -def mcp_prefixed_tool_name(server_name: str, tool_name: str) -> str: - """Build the registry/wire name for an MCP tool. - - Produces ``mcp____``. - """ - safe_server = sanitize_mcp_name_component(server_name) - safe_tool = sanitize_mcp_name_component(tool_name) - return f"{MCP_TOOL_NAME_PREFIX}{safe_server}{_MCP_NAME_DELIM}{safe_tool}" - - -def _convert_mcp_schema(server_name: str, mcp_tool) -> dict: - """Convert an MCP tool listing to the Hermes registry schema format. - - Args: - server_name: The logical server name for prefixing. - mcp_tool: An MCP ``Tool`` object with ``.name``, ``.description``, - and ``.input_schema`` (``.inputSchema`` before mcp 2.0). - - Returns: - A dict suitable for ``registry.register(schema=...)``. - """ - prefixed_name = mcp_prefixed_tool_name(server_name, mcp_tool.name) - return { - "name": prefixed_name, - "description": strip_unicode_tags( - mcp_tool.description or f"MCP tool {mcp_tool.name} from {server_name}" - ), - "parameters": _normalize_mcp_input_schema( - mcp_field(mcp_tool, "input_schema", "inputSchema") - ), - } - - -def _build_utility_schemas(server_name: str) -> List[dict]: - """Build schemas for the MCP utility tools (resources & prompts). - - Returns a list of (schema, handler_factory_name) tuples encoded as dicts - with keys: schema, handler_key. - """ - return [ - { - "schema": { - "name": mcp_prefixed_tool_name(server_name, "list_resources"), - "description": f"List available resources from MCP server '{server_name}'", - "parameters": { - "type": "object", - "properties": {}, - }, - }, - "handler_key": "list_resources", - }, - { - "schema": { - "name": mcp_prefixed_tool_name(server_name, "read_resource"), - "description": f"Read a resource by URI from MCP server '{server_name}'", - "parameters": { - "type": "object", - "properties": { - "uri": { - "type": "string", - "description": "URI of the resource to read", - }, - }, - "required": ["uri"], - }, - }, - "handler_key": "read_resource", - }, - { - "schema": { - "name": mcp_prefixed_tool_name(server_name, "list_prompts"), - "description": f"List available prompts from MCP server '{server_name}'", - "parameters": { - "type": "object", - "properties": {}, - }, - }, - "handler_key": "list_prompts", - }, - { - "schema": { - "name": mcp_prefixed_tool_name(server_name, "get_prompt"), - "description": f"Get a prompt by name from MCP server '{server_name}'", - "parameters": { - "type": "object", - "properties": { - "name": { - "type": "string", - "description": "Name of the prompt to retrieve", - }, - "arguments": { - "type": "object", - "description": "Optional arguments to pass to the prompt", - "properties": {}, - "additionalProperties": True, - }, - }, - "required": ["name"], - }, - }, - "handler_key": "get_prompt", - }, - ] - - -def _normalize_name_filter(value: Any, label: str) -> set[str]: - """Normalize include/exclude config to a set of tool-name patterns. - - Entries may be exact tool names or fnmatch-style globs - (``*_radar_*``, ``get_zones_*``). Matching happens in - :func:`matches_name_filter`. - """ - if value is None: - return set() - if isinstance(value, str): - return {value} - if isinstance(value, (list, tuple, set)): - return {str(item) for item in value} - logger.warning("MCP config %s must be a string or list of strings; ignoring %r", label, value) - return set() - - -def matches_name_filter(tool_name: str, patterns: set[str]) -> bool: - """True if ``tool_name`` matches any entry in ``patterns``. - - Exact names match literally; entries containing fnmatch metacharacters - (``*``, ``?``, ``[``) match as case-sensitive globs — the same pattern - semantics as ``approvals.deny``. Exact membership is checked first so - large literal lists stay O(1). - """ - if not patterns: - return False - if tool_name in patterns: - return True - return any( - fnmatch.fnmatchcase(tool_name, p) - for p in patterns - if "*" in p or "?" in p or "[" in p - ) - - -def _parse_boolish(value: Any, default: bool = True) -> bool: - """Parse a bool-like config value with safe fallback.""" - if value is None: - return default - if isinstance(value, bool): - return value - if isinstance(value, str): - lowered = value.strip().lower() - if lowered in {"true", "1", "yes", "on"}: - return True - if lowered in {"false", "0", "no", "off"}: - return False - logger.warning("MCP config expected a boolean-ish value, got %r; using default=%s", value, default) - return default - - -def _get_lifecycle_seconds(config: dict, key: str) -> Optional[float]: - """Return an optional positive lifecycle timeout from top-level/nested config.""" - raw = config.get(key) - lifecycle = config.get("lifecycle") - if raw is None and isinstance(lifecycle, dict): - raw = lifecycle.get(key) - if raw is None: - return None - try: - seconds = float(raw) - except (TypeError, ValueError): - logger.warning("MCP config %s must be a number of seconds; ignoring %r", key, raw) - return None - if seconds == 0: - return None - if seconds < 0: - logger.warning("MCP config %s must be positive; ignoring %r", key, raw) - return None - return seconds - - -_UTILITY_CAPABILITY_METHODS = { - "list_resources": "list_resources", - "read_resource": "read_resource", - "list_prompts": "list_prompts", - "get_prompt": "get_prompt", -} - -# Maps each utility handler to the MCP capability key that must be non-None -# on the server's ``initialize`` response for the handler to be registered. -# Source of truth: MCP spec — capabilities.resources / capabilities.prompts -# are present on the response only when the server actually implements -# those request families. Without this gate, tools-only servers (e.g. -# Context7 @upstash/context7-mcp, which advertises only ``tools``) had -# all four utility stubs registered and every model call to them came -# back with JSON-RPC ``-32601 Method not found``, which made the model -# conclude the server was broken even when the real tools worked. See -# #18051. -_UTILITY_CAPABILITY_ATTRS = { - "list_resources": "resources", - "read_resource": "resources", - "list_prompts": "prompts", - "get_prompt": "prompts", -} - - -def _track_mcp_tool_server(tool_name: str, server_name: str) -> None: - """Remember the exact raw MCP server that registered *tool_name*.""" - with _lock: - _mcp_tool_server_names[tool_name] = server_name - - -def _forget_mcp_tool_server(tool_name: str) -> None: - """Forget MCP server provenance for a deregistered tool.""" - with _lock: - _mcp_tool_server_names.pop(tool_name, None) - - -def _select_utility_schemas(server_name: str, server: MCPServerTask, config: dict) -> List[dict]: - """Select utility schemas based on config and server capabilities.""" - tools_filter = config.get("tools") or {} - resources_enabled = _parse_boolish(tools_filter.get("resources"), default=True) - prompts_enabled = _parse_boolish(tools_filter.get("prompts"), default=True) - - # ``initialize_result.capabilities`` is the source of truth: its sub-objects - # (``resources``, ``prompts``) are non-None iff the server advertises that - # request family. ``hasattr(server.session, ...)`` was the old gate but - # ClientSession always has the four method attributes defined on the class, - # so it never filtered anything. - advertised_caps = None - init_result = getattr(server, "initialize_result", None) - if init_result is not None: - advertised_caps = getattr(init_result, "capabilities", None) - - selected: List[dict] = [] - for entry in _build_utility_schemas(server_name): - handler_key = entry["handler_key"] - if handler_key in {"list_resources", "read_resource"} and not resources_enabled: - logger.debug("MCP server '%s': skipping utility '%s' (resources disabled)", server_name, handler_key) - continue - if handler_key in {"list_prompts", "get_prompt"} and not prompts_enabled: - logger.debug("MCP server '%s': skipping utility '%s' (prompts disabled)", server_name, handler_key) - continue - - # Preferred gate: check the server's advertised capabilities. Skip - # if the capability is explicitly not advertised. - if advertised_caps is not None: - cap_attr = _UTILITY_CAPABILITY_ATTRS[handler_key] - if getattr(advertised_caps, cap_attr, None) is None: - logger.debug( - "MCP server '%s': skipping utility '%s' " - "(server does not advertise '%s' capability)", - server_name, - handler_key, - cap_attr, - ) - continue - else: - # Legacy fallback for test fixtures or older code paths where - # initialize_result wasn't captured. Preserves the old behavior - # of registering every stub in that case rather than regressing - # any server that was working before this fix. - required_method = _UTILITY_CAPABILITY_METHODS[handler_key] - if not hasattr(server.session, required_method): - logger.debug( - "MCP server '%s': skipping utility '%s' (session lacks %s)", - server_name, - handler_key, - required_method, - ) - continue - selected.append(entry) - return selected - - -def _existing_tool_names() -> List[str]: - """Return tool names for all currently connected servers.""" - names: List[str] = [] - for _sname, server in _servers.items(): - if hasattr(server, "_registered_tool_names"): - names.extend(server._registered_tool_names) - continue - for mcp_tool in server._tools: - schema = _convert_mcp_schema(server.name, mcp_tool) - names.append(schema["name"]) - # Lazy servers registered from the schema cache have no MCPServerTask - # yet — their tools live in the registry only (#56832). - with _lock: - lazy_names = [ - n - for sname, tool_names in _lazy_server_tool_names.items() - if sname not in _servers - for n in tool_names - ] - names.extend(lazy_names) - return names - - -def _register_server_tools(name: str, server: MCPServerTask, config: dict) -> List[str]: - """Register tools from an already-connected server into the registry. - - Handles include/exclude filtering and utility tools. Toolset resolution - for ``mcp-{server}`` and raw server-name aliases is derived from the live - registry, rather than mutating ``toolsets.TOOLSETS`` at runtime. - - Lossy provider-safe name normalization can map distinct raw names to the - same registry name (for example ``read-file`` and ``read_file``). Such - collisions fail closed: every ambiguous entry is skipped rather than - selecting an arbitrary handler. - - Used by both initial discovery and dynamic refresh (list_changed). - - Returns: - List of registered prefixed tool names. - """ - from tools.registry import registry - - registered_names: List[str] = [] - toolset_name = f"mcp-{name}" - - # Selective tool loading: honour include/exclude lists from config. - # Rules (matching issue #690 spec, extended with glob support): - # tools.include — whitelist: only matching tool names are registered - # tools.exclude — blacklist: all tools EXCEPT matching ones are registered - # entries may be exact names or fnmatch globs (e.g. "*_radar_*") - # include takes precedence over exclude - # include: [] → register nothing (an explicit empty whitelist, as - # written by the install checklist's "uncheck everything" path) - # Neither set → register all tools (backward-compatible default) - tools_filter = config.get("tools") or {} - include_raw = tools_filter.get("include") - include_set = _normalize_name_filter( - include_raw, f"mcp_servers.{name}.tools.include" - ) - include_active = isinstance(include_raw, (str, list, tuple, set)) - exclude_set = _normalize_name_filter( - tools_filter.get("exclude"), f"mcp_servers.{name}.tools.exclude" - ) - - def _should_register(tool_name: str) -> bool: - if include_active: - return matches_name_filter(tool_name, include_set) - if exclude_set: - return not matches_name_filter(tool_name, exclude_set) - return True - - check_fn = _make_check_fn(name) - candidates: List[dict] = [] - - # Trust-tier metadata (security boundary): capture the server's - # configured trust tier and each tool's readOnlyHint annotation NOW, - # at discovery, so the call-time gate in _make_tool_handler classifies - # from data we control rather than re-reading server-supplied state. - _record_tool_trust_metadata(name, config, server._tools) - - for mcp_tool in server._tools: - if not _should_register(mcp_tool.name): - logger.debug( - "MCP server '%s': skipping tool '%s' (filtered by config)", - name, - mcp_tool.name, - ) - continue - - _scan_mcp_description(name, mcp_tool.name, mcp_tool.description or "") - schema = _convert_mcp_schema(name, mcp_tool) - candidates.append( - { - "registry_name": schema["name"], - "origin": f"tool {mcp_tool.name!r}", - "schema": schema, - "handler": _make_tool_handler( - name, mcp_tool.name, server.tool_timeout - ), - "check_fn": check_fn, - } - ) - - # Generated resource/prompt utility tools share the same namespace as raw - # MCP tools, so they must participate in the same collision preflight. - handler_factories = { - "list_resources": _make_list_resources_handler, - "read_resource": _make_read_resource_handler, - "list_prompts": _make_list_prompts_handler, - "get_prompt": _make_get_prompt_handler, - } - for entry in _select_utility_schemas(name, server, config): - schema = entry["schema"] - handler_key = entry["handler_key"] - candidates.append( - { - "registry_name": schema["name"], - "origin": f"generated utility {handler_key!r}", - "schema": schema, - "handler": handler_factories[handler_key]( - name, server.tool_timeout - ), - "check_fn": check_fn, - } - ) - - # Exact duplicate rows from a server are harmless but should not inflate - # counts. Distinct origins that collapse to one normalized name are unsafe. - unique_candidates: List[dict] = [] - seen_candidates: set[tuple[str, str]] = set() - origins_by_name: Dict[str, set[str]] = {} - for candidate in candidates: - key = (candidate["registry_name"], candidate["origin"]) - if key in seen_candidates: - logger.debug( - "MCP server '%s': duplicate registration candidate %s for '%s'; " - "keeping one", - name, - candidate["origin"], - candidate["registry_name"], - ) - continue - seen_candidates.add(key) - unique_candidates.append(candidate) - origins_by_name.setdefault(candidate["registry_name"], set()).add( - candidate["origin"] - ) - - # A generated resource/prompt utility that normalizes onto a server-native - # tool's name must not knock that native tool out of the registry: the - # native tool is the capability the user connected the server for, while the - # generated utility (read_resource/list_resources/list_prompts/get_prompt) - # is optional sugar that only matters when the server exposes no such tool - # of its own (#87112). Resolve that specific collision in favour of the - # native tool — keep it, drop the shadowed utility — and fall back to the - # conservative skip-everything only for genuinely ambiguous collisions (two - # or more native tools normalizing to one name, which we cannot - # disambiguate). The four utility keys are distinct, so a colliding set - # holds at most one utility origin. - ambiguous_names: Dict[str, List[str]] = {} - shadowed_utilities: set[tuple[str, str]] = set() - for registry_name, origins in origins_by_name.items(): - if len(origins) <= 1: - continue - utility_origins = sorted( - o for o in origins if o.startswith("generated utility ") - ) - native_origins = sorted(origins - set(utility_origins)) - if len(native_origins) == 1 and utility_origins: - for util_origin in utility_origins: - shadowed_utilities.add((registry_name, util_origin)) - logger.info( - "MCP server '%s': generated utility %s normalizes onto " - "server-native %s — keeping the native tool and dropping the " - "utility (the utility only applies when the server has no such " - "tool of its own)", - name, - ", ".join(utility_origins), - native_origins[0], - ) - continue - ambiguous_names[registry_name] = sorted(origins) - - for registry_name, origins in sorted(ambiguous_names.items()): - logger.error( - "MCP server '%s': name normalization collision for '%s' from %s; " - "skipping every colliding entry instead of choosing an arbitrary " - "handler", - name, - registry_name, - ", ".join(origins), - ) - - for candidate in unique_candidates: - registry_name = candidate["registry_name"] - if registry_name in ambiguous_names: - continue - if (registry_name, candidate["origin"]) in shadowed_utilities: - continue - - existing_toolset = registry.get_toolset_for_tool(registry_name) - if existing_toolset and existing_toolset != toolset_name: - if existing_toolset.startswith("mcp-"): - logger.error( - "MCP server '%s': %s normalizes to '%s', already owned by " - "MCP toolset '%s' — skipping to preserve the existing owner", - name, - candidate["origin"], - registry_name, - existing_toolset, - ) - else: - logger.warning( - "MCP server '%s': %s (→ '%s') collides with built-in tool " - "in toolset '%s' — skipping to preserve built-in", - name, - candidate["origin"], - registry_name, - existing_toolset, - ) - continue - - registry.register( - name=registry_name, - toolset=toolset_name, - schema=candidate["schema"], - handler=candidate["handler"], - check_fn=candidate["check_fn"], - is_async=False, - description=candidate["schema"]["description"], - scope=_server_registry_scope(name), - ) - - # The pre-check above is advisory only. Multiple servers connect in - # parallel, so ToolRegistry.register() is the atomic ownership gate. - if registry.get_toolset_for_tool(registry_name) != toolset_name: - logger.error( - "MCP server '%s': registration of %s as '%s' was rejected by " - "the registry; skipping provenance/count updates", - name, - candidate["origin"], - registry_name, - ) - continue - - _track_mcp_tool_server(registry_name, name) - registered_names.append(registry_name) - - if registered_names: - registry.register_toolset_alias(name, toolset_name) - # Write-through (#56832): refresh the on-disk schema cache after a - # live connect so the next startup can lazily register this server - # without spawning it. Cache failures never break registration. - try: - from tools.mcp_schema_cache import config_fingerprint, write_cache_entry - - tools_payload: List[dict] = [] - for mcp_tool in server._tools: - if not _should_register(mcp_tool.name): - continue - schema_obj = getattr(mcp_tool, "inputSchema", None) - tools_payload.append({ - "name": mcp_tool.name, - "description": mcp_tool.description or "", - "inputSchema": schema_obj if isinstance(schema_obj, dict) else {}, - # Persist the trust-relevant annotation so the lazy - # (cache-registered) path gates identically on next - # startup without spawning the server. - "annotations": { - "readOnlyHint": _annotation_read_only_hint(mcp_tool), - }, - }) - utility_payload = [ - {"schema": entry["schema"], "handler_key": entry["handler_key"]} - for entry in _select_utility_schemas(name, server, config) - ] - write_cache_entry( - name, - config_fingerprint(config), - tools=tools_payload, - utility_tools=utility_payload, - ttl_ms=(getattr(server, "_list_cache_meta", None) or {}).get("ttl_ms"), - cache_scope=(getattr(server, "_list_cache_meta", None) or {}).get("cache_scope"), - ) - except Exception as exc: - logger.debug("MCP schema cache write failed for '%s': %s", name, exc) - - return registered_names - - -class _CachedMCPTool: - """Minimal stand-in for MCP Tool objects loaded from the schema cache.""" - - __slots__ = ("name", "description", "inputSchema") - - def __init__(self, name: str, description: str, inputSchema: dict): - self.name = name - self.description = description - self.inputSchema = inputSchema or {} - - -def _register_from_cache_sync(name: str, config: dict, entry: dict) -> List[str]: - """Register a server's tools from a cached manifest, no child process. - - Lazy startup (#56832, design by Vansh5632): tools appear in the registry - immediately; the first real call routes through - ``_get_connected_server_for_call`` → ``_ensure_lazy_server_connected``. - """ - from tools.registry import registry - from tools.mcp_schema_cache import ( - config_fingerprint, - tools_from_cache_entry, - utility_tools_from_cache_entry, - ) - - registered_names: List[str] = [] - toolset_name = f"mcp-{name}" - fingerprint = config_fingerprint(config) - tool_timeout = _resolve_tool_timeout(config) - tools_filter = config.get("tools") or {} - include_raw = tools_filter.get("include") - include_set = _normalize_name_filter( - include_raw, f"mcp_servers.{name}.tools.include" - ) - # include: [] is an explicit empty whitelist (register nothing) — see the - # live discovery path above for the full filter rules. - include_active = isinstance(include_raw, (str, list, tuple, set)) - exclude_set = _normalize_name_filter( - tools_filter.get("exclude"), f"mcp_servers.{name}.tools.exclude" - ) - - def _should_register(tool_name: str) -> bool: - if include_active: - return matches_name_filter(tool_name, include_set) - if exclude_set: - return not matches_name_filter(tool_name, exclude_set) - return True - - check_fn = _make_check_fn(name) - # Trust-tier metadata for the lazy path: the cached manifest carries - # each tool's readOnlyHint (written by the live discovery path), and - # trust comes from operator config. Recording it before registration - # keeps the call-time gate identical whether the server was spawned - # live or registered from cache. Missing "annotations" in older cache - # files fails closed to write-capable. - cached_tool_objs = [ - SimpleNamespace( - name=raw.get("name"), - annotations=raw.get("annotations") - if isinstance(raw.get("annotations"), dict) else None, - ) - for raw in tools_from_cache_entry(entry) - if isinstance(raw, dict) and raw.get("name") - ] - _record_tool_trust_metadata(name, config, cached_tool_objs) - for raw in tools_from_cache_entry(entry): - if not isinstance(raw, dict): - continue - raw_name = raw.get("name") - if not raw_name or not _should_register(raw_name): - continue - raw_schema = raw.get("inputSchema") - mcp_tool = _CachedMCPTool( - raw_name, - raw.get("description") or "", - raw_schema if isinstance(raw_schema, dict) else {}, - ) - # Defense-in-depth: the cache file is user-writable JSON, so run the - # same injection scan the eager discovery path applies. - _scan_mcp_description(name, mcp_tool.name, mcp_tool.description or "") - schema = _convert_mcp_schema(name, mcp_tool) - registry_name = schema["name"] - existing_toolset = registry.get_toolset_for_tool(registry_name) - if existing_toolset and existing_toolset != toolset_name: - logger.warning( - "MCP server '%s' (lazy): cached tool '%s' collides with " - "toolset '%s' — skipping", - name, registry_name, existing_toolset, - ) - continue - registry.register( - name=registry_name, - toolset=toolset_name, - schema=schema, - handler=_make_tool_handler(name, raw_name, tool_timeout), - check_fn=check_fn, - is_async=False, - description=schema["description"], - scope=_mcp_registry_scope(), - ) - if registry.get_toolset_for_tool(registry_name) != toolset_name: - continue - _track_mcp_tool_server(registry_name, name) - registered_names.append(registry_name) - - handler_factories = { - "list_resources": _make_list_resources_handler, - "read_resource": _make_read_resource_handler, - "list_prompts": _make_list_prompts_handler, - "get_prompt": _make_get_prompt_handler, - } - for raw in utility_tools_from_cache_entry(entry): - if not isinstance(raw, dict): - continue - schema = raw.get("schema") - handler_key = raw.get("handler_key") - if not isinstance(schema, dict) or handler_key not in handler_factories: - continue - util_name = schema.get("name") or "" - if not util_name: - continue - existing_toolset = registry.get_toolset_for_tool(util_name) - if existing_toolset and existing_toolset != toolset_name: - continue - registry.register( - name=util_name, - toolset=toolset_name, - schema=schema, - handler=handler_factories[handler_key](name, tool_timeout), - check_fn=check_fn, - is_async=False, - description=schema.get("description") or "", - scope=_mcp_registry_scope(), - ) - if registry.get_toolset_for_tool(util_name) != toolset_name: - continue - _track_mcp_tool_server(util_name, name) - registered_names.append(util_name) - - if registered_names: - registry.register_toolset_alias(name, toolset_name) - with _lock: - _lazy_server_configs[name] = dict(config) - _lazy_server_fingerprints[name] = fingerprint - _lazy_server_tool_names[name] = list(registered_names) - logger.info( - "MCP server '%s' (lazy): registered %d tool(s) from schema cache", - name, len(registered_names), - ) - return registered_names - async def _discover_and_register_server(name: str, config: dict) -> List[str]: - """Connect to a single MCP server, discover tools, and register them. - - Returns list of registered tool names. - """ + """Connect one server, register its tools; return the registered names.""" connect_timeout = config.get("connect_timeout", _DEFAULT_CONNECT_TIMEOUT) - # List-based claim (not a ``nonlocal`` rebind): the claim callback runs - # inside ``_connect_server`` while this frame is suspended, and appending - # keeps type narrowing intact for the module's other ``server`` locals. + # The claim callback runs inside _connect_server while this frame is + # suspended; a list append avoids a nonlocal rebind. claimed: List[MCPServerTask] = [] def _claim_server(created: MCPServerTask) -> None: @@ -7803,8 +1881,8 @@ async def _discover_and_register_server(name: str, config: dict) -> List[str]: and not task.done() and not task_cancelling ): - # Recoverable park: the run task deliberately stays alive to - # self-probe, so adopt it into the registry for shutdown/revival. + # Recoverable park: the run task stays alive to self-probe, so + # adopt it for shutdown/revival. with _lock: _servers[name] = server _server_scope_keys[name] = _mcp_registry_scope() @@ -7837,16 +1915,11 @@ async def _discover_and_register_server(name: str, config: dict) -> List[str]: # --------------------------------------------------------------------------- def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: - """Connect to explicit MCP servers and register their tools. + """Connect the given ``{name: config}`` servers and register their tools. - Idempotent for already-connected server names. Servers with - ``enabled: false`` are skipped without disconnecting existing sessions. - - Args: - servers: Mapping of ``{server_name: server_config}``. - - Returns: - List of all currently registered MCP tool names. + Idempotent for connected names; ``enabled: false`` servers are skipped + without disconnecting existing sessions. Returns every registered MCP + tool name. """ if not _ensure_mcp_sdk(): logger.debug("MCP SDK not available -- skipping explicit MCP registration") @@ -7857,10 +1930,8 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: logger.debug("No explicit MCP servers provided") return [] - # Only attempt servers that aren't already connected (or currently - # connecting) and are enabled. Checking ``_server_connecting`` prevents - # duplicate subprocess spawns when ``discover_mcp_tools()`` is called - # from multiple entry-points before the first batch finishes (#58862). + # Candidates: enabled, not connected, not connecting (dedups concurrent + # discovery entry points), not lazily registered, not in backoff. with _lock: connecting = set(_server_connecting) new_servers = { @@ -7868,23 +1939,12 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: for k, v in servers.items() if k not in _servers and k not in connecting - # Servers already lazily registered from the schema cache are - # not re-registered; they connect on first tool use (#56832). and k not in _lazy_server_configs and _parse_boolish(v.get("enabled", True), default=True) - # Skip a server still serving its post-failure backoff. Without - # this, a server that fails to connect (and is therefore never - # recorded in ``_servers``) would be re-spawned on every worker - # session's discovery pass -- the #50394 restart storm. The - # cooldown is cleared automatically on the next successful - # connect or by a manual /mcp refresh. and not _connect_cooldown_active(k) } - # Cached entries with no live session are parked or mid-reconnect. - # Their tools are deregistered, so nothing else can reach - # _signal_reconnect — without this nudge a new session silently - # waits up to _PARKED_RETRY_INTERVAL for the next self-probe - # (#50170). Wake them now so their tools come back promptly. + # Known servers without a live session are parked or mid-reconnect; + # their tools are deregistered so nothing else can nudge them. stale_cached = [ _servers[k] for k in servers @@ -7906,11 +1966,8 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: if not new_servers: return _existing_tool_names() - # Lazy startup (#56832): servers gated with ``lazy: true`` whose config - # fingerprint matches a valid on-disk schema-cache entry register their - # tools from cache WITHOUT spawning/connecting. A missing or stale cache - # entry falls back to the normal eager connect below (which write-through - # refreshes the cache for next time). + # ``lazy: true`` servers with a valid schema-cache entry register from + # cache without connecting; a missing/stale entry falls back to eager. eager_servers: Dict[str, dict] = dict(new_servers) lazy_registered = 0 lazy_server_count = 0 @@ -7951,18 +2008,12 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: ) return _existing_tool_names() - # Start the background event loop for MCP connections _ensure_mcp_loop() - async def _discover_one(name: str, cfg: dict) -> List[str]: - """Connect to a single server and return its registered tool names.""" - return await _discover_and_register_server(name, cfg) - async def _discover_all(): server_names = list(new_servers.keys()) - # Connect to all servers in PARALLEL results = await asyncio.gather( - *(_discover_one(name, cfg) for name, cfg in new_servers.items()), + *(_discover_and_register_server(name, cfg) for name, cfg in new_servers.items()), return_exceptions=True, ) for name, result in zip(server_names, results): @@ -7972,10 +2023,6 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: with _lock: _server_connecting.discard(name) _server_connect_errors[name] = message - # Arm the per-server backoff so the next discovery pass - # doesn't immediately re-spawn this failing server - # (#50394). Isolated to this server -- healthy servers - # in the same batch are unaffected. _record_connect_failure(name) logger.warning( "Failed to connect to MCP server '%s'%s: %s", @@ -7989,12 +2036,8 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: _server_connect_errors.pop(name, None) _clear_connect_failure(name) - # Per-server timeouts are handled inside _discover_and_register_server. - # The outer timeout is generous: 120s total for parallel discovery. - # - # Temporarily clear the interrupt flag on the current thread so that MCP - # discovery is never cancelled by a stale interrupt from a prior agent - # session (executor threads get reused and may carry old interrupt state). + # Clear a stale interrupt flag (executor threads are reused) so a prior + # session's interrupt cannot cancel this discovery pass. from tools.interrupt import is_interrupted as _is_interrupted, set_interrupt as _set_interrupt _was_interrupted = _is_interrupted() if _was_interrupted: @@ -8002,10 +2045,8 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: try: _run_on_mcp_loop(_discover_all, timeout=120) except (TimeoutError, InterruptedError) as _e: - # When the outer timeout fires or the user interrupts, - # _discover_all's gather may not have finished, leaving - # entries stranded in _server_connecting. Those stale - # entries would block future reconnection attempts (#58862). + # Entries stranded in _server_connecting would block future + # reconnect attempts. with _lock: stale = [n for n in new_servers if n in _server_connecting] if stale: @@ -8027,7 +2068,6 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: if _was_interrupted: _set_interrupt(True) - # Log a summary so ACP callers get visibility into what was registered. with _lock: connected = [ n @@ -8051,32 +2091,24 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: def discover_mcp_tools() -> List[str]: - """Entry point: load config, connect to MCP servers, register tools. + """Entry point: load config, connect servers, register tools. - Called from ``model_tools`` after ``discover_builtin_tools()``. Safe to call even when - the ``mcp`` package is not installed (returns empty list). - - Idempotent for already-connected servers. If some servers failed on a - previous call, only the missing ones are retried. - - Returns: - List of all registered MCP tool names. + Safe without the ``mcp`` package (returns []). Idempotent: only servers + missing from a previous call are retried. Returns all MCP tool names. """ servers = _load_mcp_config() if not servers: logger.debug("No MCP servers configured") return [] - # SDK import is deferred to HERE so a config with zero MCP servers (the - # default) never pays the ~260ms `mcp` import on CLI startup. + # SDK import deferred to here so a config without servers never pays it. if not _ensure_mcp_sdk(): logger.debug("MCP SDK not available -- skipping MCP tool discovery") return [] - # Cross-process discovery guard (#62771). A lock loser waits for - # the holder, then performs its own process-local discovery. If locking is - # unavailable or the bounded wait expires, preserve the previous - # fail-soft behavior by running discovery unguarded. + # Cross-process guard: a lock loser waits for the holder, then runs its + # own discovery; if locking is unavailable or the wait expires, run + # unguarded (fail-soft). cookie = _try_acquire_mcp_discovery_lock() if cookie is None: logger.debug( @@ -8137,15 +2169,10 @@ def discover_mcp_tools() -> List[str]: cookie.release() def is_mcp_tool_parallel_safe(tool_name: str) -> bool: - """Check if an MCP tool belongs to a server that supports parallel tool calls. + """True when the tool's server opted into ``supports_parallel_tool_calls``. - MCP tool names follow the pattern ``mcp__{server}__{tool}``, but that - string shape is ambiguous when server names contain underscores. Use the - exact server provenance captured at registration time rather than prefix - matching, then check whether that server's config includes - ``supports_parallel_tool_calls: true``. - - Returns False for non-MCP tools or tools from servers without the flag. + Uses the provenance captured at registration, never the (ambiguous) + ``mcp__{server}__{tool}`` string shape. """ if not tool_name.startswith(MCP_TOOL_NAME_PREFIX): return False @@ -8155,96 +2182,62 @@ def is_mcp_tool_parallel_safe(tool_name: str) -> bool: def get_mcp_status() -> List[dict]: - """Return status of all configured MCP servers for banner display. + """Status of every configured server for banner/TUI display. - Returns a list of dicts with keys: name, transport, tools, connected, - disabled, and status. Includes connected servers, disabled servers, - in-flight connection attempts, recorded failures, and servers that are - configured but have not been started in this process yet. + Each dict has name, transport, tools, connected, disabled, status (one of + connected / disabled / connecting / failed / configured) and, for failed, + error. ``enabled: false`` is reported as disabled, not failed. """ - result: List[dict] = [] - - # Get configured servers from config configured = _load_mcp_config() if not configured: - return result + return [] with _lock: active_servers = dict(_servers) connecting = set(_server_connecting) connect_errors = dict(_server_connect_errors) + def _entry(name: str, transport: str, status: str, **extra) -> dict: + return { + "name": name, + "transport": transport, + "tools": 0, + "connected": False, + "disabled": status == "disabled", + "status": status, + **extra, + } + + result: List[dict] = [] for name, cfg in configured.items(): transport = cfg.get("transport", "http") if "url" in cfg else "stdio" enabled = _parse_boolish(cfg.get("enabled", True), default=True) server = active_servers.get(name) if server and server.session is not None: - entry = { - "name": name, - "transport": transport, - "tools": len(server._registered_tool_names) if hasattr(server, "_registered_tool_names") else len(server._tools), - "connected": True, - "disabled": False, - "status": "connected", - } + entry = _entry(name, transport, "connected", connected=True) + entry["tools"] = ( + len(server._registered_tool_names) + if hasattr(server, "_registered_tool_names") + else len(server._tools) + ) if server._sampling: entry["sampling"] = dict(server._sampling.metrics) - result.append(entry) elif not enabled: - # A server with enabled: false is intentionally not connected — it is - # disabled, not failed. Surface that distinction so consumers (banner, - # TUI) can render "disabled" rather than an alarming "failed". - result.append({ - "name": name, - "transport": transport, - "tools": 0, - "connected": False, - "disabled": True, - "status": "disabled", - }) + entry = _entry(name, transport, "disabled") elif name in connecting: - result.append({ - "name": name, - "transport": transport, - "tools": 0, - "connected": False, - "disabled": False, - "status": "connecting", - }) + entry = _entry(name, transport, "connecting") elif name in connect_errors: - result.append({ - "name": name, - "transport": transport, - "tools": 0, - "connected": False, - "disabled": False, - "status": "failed", - "error": connect_errors[name], - }) + entry = _entry(name, transport, "failed", error=connect_errors[name]) else: - result.append({ - "name": name, - "transport": transport, - "tools": 0, - "connected": False, - "disabled": False, - "status": "configured", - }) + entry = _entry(name, transport, "configured") + result.append(entry) return result def probe_mcp_server_tools() -> Dict[str, List[tuple]]: - """Temporarily connect to configured MCP servers and list their tools. - - Designed for ``hermes tools`` interactive configuration — connects to each - enabled server, grabs tool names and descriptions, then disconnects. - Does NOT register tools in the Hermes registry. - - Returns: - Dict mapping server name to list of (tool_name, description) tuples. - Servers that fail to connect are omitted from the result. - """ + """Connect to each enabled server, list ``(tool_name, description)`` and + disconnect, without registering anything. Failed servers are omitted.""" if not _ensure_mcp_sdk(): return {} @@ -8284,7 +2277,6 @@ def probe_mcp_server_tools() -> Dict[str, List[tuple]]: tools.append((t.name, desc)) result[name] = tools - # Shut down all probed connections await asyncio.gather( *(s.shutdown() for s in probed_servers), return_exceptions=True, @@ -8300,667 +2292,23 @@ def probe_mcp_server_tools() -> Dict[str, List[tuple]]: return result -# Serializes in-place mutation of an agent's tool snapshot. The reload RPC, -# the gateway reload, and the late-binding refresh thread all swap -# ``agent.tools`` / ``agent.valid_tool_names`` after the agent was built; the -# agent's run loop reads those during tool iteration, so a concurrent write -# mid-read could otherwise expose a half-updated list. -_agent_tools_lock = threading.Lock() - - def has_registered_mcp_tools() -> bool: - """True if any MCP server has actually registered tools into the registry. + """True if any MCP server has registered tools (cheap; no registry walk). - Cheap — checks the global MCP-tool→server name map under ``_lock``, no - registry walk. Used by the per-turn refresh hook so a session with no MCP - tools (the common case, and also a connected-but-zero-tool/prompt-only - server) skips the ``get_tool_definitions`` rebuild entirely. Checks - registered TOOLS, not connected servers, so a server that registers no tools - doesn't keep the hook firing every turn. + Checks registered TOOLS, not connected servers, so the per-turn refresh + hook stays idle for zero-tool servers. """ with _lock: return bool(_mcp_tool_server_names) def get_registered_mcp_server_names() -> set: - """Return the set of MCP server names that have actually registered at - least one tool into the registry (post-connection, post check_fn/include- - exclude filtering) -- i.e. the real, availability-filtered signal, not - just what's present in config.yaml under ``mcp_servers``. - - Used by capability-aware prompt building (e.g. gateway/session.py's - Slack platform note) to detect an MCP server that provides a given - platform's capability regardless of what its config key is named. - """ + """Server names that registered at least one tool (the live, filtered + signal — not merely what config.yaml lists).""" with _lock: return set(_mcp_tool_server_names.values()) - -def refresh_agent_mcp_tools( - agent, - *, - enabled_override=None, - disabled_override=None, - quiet_mode: bool = True, - content_aware: bool = False, - preserve_prefix: bool = False, -) -> set: - """Re-derive an already-built agent's tool snapshot from the live registry. - - The agent snapshots ``agent.tools`` once at build time and never re-reads - the registry (see ``run_agent`` / ``agent_init``). When MCP servers connect - *after* that snapshot — a slow HTTP/OAuth server that misses the bounded - startup wait, or a ``/reload-mcp`` — their tools are invisible until the - snapshot is rebuilt. This is the single shared rebuild used by every such - caller (the TUI ``reload.mcp`` RPC, the gateway reload, the late-binding - refresh thread, and the per-turn between-turns refresh) so they can't drift - apart again. - - The rebuild respects the agent's own ``enabled_toolsets`` / - ``disabled_toolsets`` (the same filtering it was built with) and diffs by - tool **name** (not count — a count compare misses an equal-size add/remove - swap). - - Crucially it is **additive-preserving**: ``get_tool_definitions`` returns - only the registry-derived tools, but ``agent_init`` appends two further - families directly onto ``agent.tools`` *after* that — external - memory-provider tools (mem0/honcho/…) and context-engine tools - (``lcm_*``). A naive ``agent.tools = get_tool_definitions(...)`` would - silently DELETE those. So after rebuilding the registry set we re-run the - same post-build injectors ``agent_init`` used, reconstructing the full - surface. The new ``(tools, valid_tool_names)`` pair is published together - under ``_agent_tools_lock`` so a concurrent reader never sees a - cross-attribute half-swap. - - ``preserve_prefix`` is for the callers that rebuild inside a live - conversation (the between-turns prologue). There the tool array is a - cached request prefix: every provider that renders ``tools`` ahead of the - messages re-prefills the entire history behind any byte that moves. A - plain rebuild moves two kinds of bytes — it drops a tool whose ``check_fn`` - merely flapped (a headless browser probe, an expired credential, a docker - blip), and it splices a late-landing tool into sorted position, which can - be index 0. With ``preserve_prefix`` the live order is authoritative: - existing tools keep their slot (schemas still refresh), a tool that is - still *registered* but momentarily unavailable is carried forward, a tool - that genuinely left the registry is still dropped, and new tools are - appended at the tail so the prefix only ever grows. Carrying an - unavailable tool forward changes nothing about dispatch — ``check_fn`` - gates exposure at snapshot time, never invocation, and every handler - already owns its own unavailability error. - - Returns the set of newly-added tool names (empty when nothing changed), so - callers can decide whether to notify the user / re-emit session info. The - caller owns the prompt-cache contract: this helper does NOT check turn state, - because each caller has a different policy (``/reload-mcp`` rebuilds after - explicit user consent; the late-binding and between-turns paths only rebuild - at a turn boundary, before that turn's ``tools=`` prefix is assembled). - """ - from model_tools import get_tool_definitions - from tools.registry import registry - - # Explicit reloads (/reload-mcp) pass freshly-resolved toolsets so a server - # the user just ENABLED in config is picked up; the agent's stored selection - # is then updated to match. The automatic paths (between-turns, late-binding) - # pass nothing and reuse the agent's build-time selection unchanged. - if enabled_override is not None or disabled_override is not None: - enabled = enabled_override if enabled_override is not None else getattr(agent, "enabled_toolsets", None) - disabled = disabled_override if disabled_override is not None else getattr(agent, "disabled_toolsets", None) - agent.enabled_toolsets = enabled - agent.disabled_toolsets = disabled - else: - enabled = getattr(agent, "enabled_toolsets", None) - disabled = getattr(agent, "disabled_toolsets", None) - - # Capture the registry generation this rebuild is derived from BEFORE the - # (potentially slow) get_tool_definitions call. Used at publish time to - # reject a stale write: if two callers race (e.g. the late-refresh daemon - # and the between-turns prologue around turn 1), a slower caller that - # computed an OLDER set must not clobber a newer set another caller already - # published. ``registry._generation`` bumps on every (de)register. - snapshot_generation = registry._generation - - # Registry-derived tools (built-ins + MCP), filtered to the agent's toolsets. - # Computed OUTSIDE the lock (get_tool_definitions can be slow); the diff and - # publish below happen together in ONE critical section so two concurrent - # callers can't torn-publish or compute overlapping ``added`` sets. - new_defs = list( - get_tool_definitions( - enabled_toolsets=enabled, - disabled_toolsets=disabled, - quiet_mode=quiet_mode, - ) - or [] - ) - new_names = {t["function"]["name"] for t in new_defs} - - # Re-append the post-build injected families that get_tool_definitions does - # NOT reproduce, so a refresh never strips them (memory-provider + context- - # engine tools). Staged entirely on LOCALS — the live ``agent.tools`` / - # ``valid_tool_names`` / ``_context_engine_tool_names`` are never touched - # until the single atomic publish below, so a concurrent reader - # (``build_api_kwargs``) can't see a partial rebuild or a cross-attribute - # half-swap. ``staged_engine_names`` are the context-engine routing names - # this rebuild actually appended (matching agent_init's dedup-aware add). - staged_engine_names = _reinject_post_build_tools(agent, new_defs, new_names) - - # Snapshot registry membership OUTSIDE ``_agent_tools_lock`` — it is the - # only input ``preserve_prefix`` needs beyond the two tool lists, and - # taking ``registry._lock`` under the tools lock would be the first place - # in the process to nest those two. - registered_names: set = set() - if preserve_prefix: - try: - registered_names = {entry.name for entry in registry.get_all_entries()} - except Exception: # noqa: BLE001 - # Fail open to the plain rebuild rather than pinning a stale list. - preserve_prefix = False - - # Single atomic read-diff-publish so the returned ``added`` is consistent - # with what was actually published, even under concurrent callers, and a - # stale (older-generation) rebuild can't overwrite a newer published one. - with _agent_tools_lock: - # Defensive: the published generation should be an int, but tolerate an - # agent that never set it (or set a non-int, e.g. a test mock) rather - # than throwing TypeError on the comparison and silently failing the - # whole refresh. - published_gen_raw = getattr(agent, "_tool_snapshot_generation", -1) - published_gen = published_gen_raw if isinstance(published_gen_raw, int) else -1 - if snapshot_generation < published_gen: - # A newer snapshot already won; our set is stale — drop it. - return set() - current_defs = list(getattr(agent, "tools", None) or []) - current = {t["function"]["name"] for t in current_defs} - if preserve_prefix: - new_defs, new_names = _merge_preserving_prefix( - current_defs, new_defs, registered_names, - ) - if new_names == current: - # Same NAME set. For MCP-reload callers that is "no change" — - # leave the live snapshot untouched (no churn). Content-aware - # callers (the compaction boundary) also diff the serialized - # bytes: dynamic schemas (image_generate capabilities, - # delegate_task limits, execute_code stubs) change CONTENT - # under stable names when config changes between compactions. - content_changed = False - if content_aware: - try: - _stable = json.dumps( - (getattr(agent, "tools", None) or []), - sort_keys=True, separators=(",", ":"), default=str, - ) - _new = json.dumps( - new_defs, sort_keys=True, separators=(",", ":"), - default=str, - ) - content_changed = _stable != _new - except Exception: # noqa: BLE001 - content_changed = False - if not content_changed: - # Record the generation so an in-flight older caller can't - # clobber. - agent._tool_snapshot_generation = max(published_gen, snapshot_generation) - return set() - agent.tools = new_defs - agent.valid_tool_names = new_names - # Publish context-engine routing names atomically with the snapshot. - engine_names = getattr(agent, "_context_engine_tool_names", None) - if isinstance(engine_names, set): - engine_names.clear() - engine_names.update(staged_engine_names) - agent._tool_snapshot_generation = max(published_gen, snapshot_generation) - added = new_names - current - # Every published snapshot re-pins the session's tool order so a later - # rebuild-for-existing-session (gateway agent-cache eviction) restores - # exactly these names — see ``restore_agent_tool_prefix``. - persist_agent_tool_names(agent) - return added - - -def reprobe_tool_availability() -> None: - """Explicit ``/reload-mcp`` hatch out of the tools[] freeze. - - Availability-gated tools (``check_fn``: Docker, HASS_TOKEN, OAuth…) are - frozen for the life of a session; a credential or daemon that appears - mid-session is only picked up when the user consciously asks. Drop the - ``check_fn`` verdict cache AND the ``get_tool_definitions`` memo (keyed on - registry generation, so it would otherwise replay the stale verdicts). - """ - from model_tools import _clear_tool_defs_cache - from tools.registry import invalidate_check_fn_cache - - invalidate_check_fn_cache() - _clear_tool_defs_cache() - - -def persist_agent_tool_names(agent) -> None: - """Best-effort: write ``agent.tools`` names to the session row (freeze pin).""" - db = getattr(agent, "_session_db", None) - session_id = getattr(agent, "session_id", None) - if not db or not session_id: - return - try: - db.update_session_tool_names( - session_id, - [t["function"]["name"] for t in (getattr(agent, "tools", None) or [])], - ) - except Exception: # noqa: BLE001 - logger.debug("tool_names persist skipped", exc_info=True) - - -def restore_agent_tool_prefix(agent, saved_names: list) -> bool: - """Fold a freshly built agent's ``tools`` onto the session's saved order. - - Closes the second door on the tools[] freeze: the gateway rebuilds a NEW - ``AIAgent`` for an existing session after agent-cache eviction, and - ``agent_init`` re-derives ``agent.tools`` from live ``check_fn`` probes - with no predecessor to preserve. The saved name list stands in for that - predecessor: a saved tool that is still registered but failed its probe - this time is carried forward from the registry's schema, a deregistered - one is dropped, and genuinely new tools append at the tail — the same - ``_merge_preserving_prefix`` rule the between-turns refresh uses. - Returns True when the snapshot was changed. - """ - if not saved_names: - return False - from tools.registry import registry - - fresh_defs = list(getattr(agent, "tools", None) or []) - fresh = {t["function"]["name"]: t for t in fresh_defs} - saved_defs = [] - for name in saved_names: - entry_def = fresh.get(name) - if entry_def is None: - entry = registry.get_entry(name) - if entry is None: - continue - entry_def = {"type": "function", "function": {**entry.schema, "name": entry.name}} - saved_defs.append(entry_def) - registered_names = {entry.name for entry in registry.get_all_entries()} - merged, merged_names = _merge_preserving_prefix(saved_defs, fresh_defs, registered_names) - with _agent_tools_lock: - if merged == fresh_defs: - return False - agent.tools = merged - agent.valid_tool_names = merged_names - if [t["function"]["name"] for t in merged] != list(saved_names): - persist_agent_tool_names(agent) - return True - - -def _merge_preserving_prefix( - current_defs: list, new_defs: list, registered_names: set, -) -> tuple[list, set]: - """Fold a fresh tool snapshot into a live one without moving existing bytes. - - The live tool array is a cached request prefix, so the merge is ordered by - ``current_defs``, not by the fresh list: - - * a name in both keeps its slot and takes the fresh schema (dynamic - overrides — delegate_task limits, execute_code stubs — still land); - * a name only in the live list is carried forward when it is still - registered (its ``check_fn`` flapped) and dropped when it is not (the - MCP server or plugin genuinely went away); - * a name only in the fresh list is appended at the tail, so a late-landing - MCP tool extends the prefix instead of splicing into sorted position. - """ - fresh = {} - for entry in new_defs: - name = (entry.get("function") or {}).get("name", "") - if name: - fresh[name] = entry - - merged = [] - for entry in current_defs: - name = (entry.get("function") or {}).get("name", "") - replacement = fresh.pop(name, None) - if replacement is not None: - merged.append(replacement) - elif name and name in registered_names: - merged.append(entry) - merged.extend(fresh.values()) - return merged, {(t.get("function") or {}).get("name", "") for t in merged} - - -def _reinject_post_build_tools(agent, tools_list: list, name_set: set) -> set: - """Append memory-provider and context-engine tools onto staged locals. - - Mirrors the post-``get_tool_definitions`` injection in ``agent_init`` so a - snapshot rebuild reconstructs the FULL tool surface, not just the - registry-derived subset. Operates ONLY on the caller's staged ``tools_list`` - / ``name_set`` (never the live agent attributes) so the rebuild stays atomic. - Idempotent (skips names already present) and fail-soft. - - Returns the set of context-engine routing names actually appended by THIS - rebuild — matching ``agent_init``'s dedup behavior (a name already provided - by a registry/plugin tool is NOT claimed for context-engine routing). The - caller publishes this into ``agent._context_engine_tool_names`` atomically - with the snapshot. - """ - def _add(schema: dict) -> bool: - name = schema.get("name", "") - if not name or name in name_set: - return False - tools_list.append({"type": "function", "function": schema}) - name_set.add(name) - return True - - # Memory-provider tools (mem0/honcho/byterover/supermemory/…). - try: - memory_manager = getattr(agent, "_memory_manager", None) - get_mem_schemas = getattr(memory_manager, "get_all_tool_schemas", None) if memory_manager else None - if callable(get_mem_schemas): - # Honor the same toolset gate inject_memory_provider_tools uses. - from agent.memory_manager import memory_provider_tools_enabled - if memory_provider_tools_enabled( - getattr(agent, "enabled_toolsets", None), - getattr(agent, "disabled_toolsets", None), - memory_tool_present="memory" in name_set, - ): - for schema in get_mem_schemas(): - if isinstance(schema, dict): - _add(schema) - except Exception: - logger.debug("Memory-provider tool re-injection skipped", exc_info=True) - - # Context-engine tools (lcm_grep/lcm_describe/…) — the `context_engine` - # toolset is intentionally empty, so these only exist via this append. - # Honor the same enabled_toolsets gate agent_init uses (#5544): without it a - # restricted-toolset platform (e.g. platform_toolsets: telegram: []) would - # re-leak lcm_* tools the build deliberately excluded, and pay the local- - # model latency penalty. - staged_engine_names: set = set() - try: - enabled = getattr(agent, "enabled_toolsets", None) - context_engine_allowed = enabled is None or "context_engine" in enabled - compressor = getattr(agent, "context_compressor", None) - get_schemas = getattr(compressor, "get_tool_schemas", None) if compressor else None - if context_engine_allowed and callable(get_schemas): - for schema in get_schemas(): - if not isinstance(schema, dict): - continue - name = schema.get("name", "") - # Only claim the routing name when WE appended the schema, so a - # name already owned by a registry/plugin tool keeps its own - # dispatch (matches agent_init.py's `continue`-before-claim). - if _add(schema) and name: - staged_engine_names.add(name) - except Exception: - logger.debug("Context-engine tool re-injection skipped", exc_info=True) - - return staged_engine_names - - -def shutdown_mcp_servers(*, scope: Optional[str] = None): - """Close MCP server connections and stop the background loop. - - Each server Task is signalled to exit its ``async with`` block so that - the anyio cancel-scope cleanup happens in the same Task that opened it. - All servers are shut down in parallel via ``asyncio.gather``. - - ``scope`` (a registry scope key) restricts teardown to the servers one - multiplexed profile owns — its ``/reload-mcp`` must not kill the other - profiles' connections — and leaves the shared loop running when anything - else is still connected. Without it every server goes, as before. - """ - with _lock: - selected = [ - name for name in _servers - if scope is None or _server_scope_keys.get(name) == scope - ] - servers_snapshot = [_servers[name] for name in selected] - - # Fast path: nothing to shut down. The connect-cooldown maps can still - # be populated here — a server that failed to connect is never recorded - # in ``_servers`` (that is the very premise of the #50394 cooldown), so - # "no live servers" is the MOST likely state in which stale backoff - # entries exist. Clear them so a post-shutdown restart re-attempts every - # configured server immediately. - if not servers_snapshot: - with _lock: - _server_connect_retry_after.clear() - _server_connect_failures.clear() - _stop_mcp_loop(only_if_idle=scope is not None) - return - - async def _shutdown(): - results = await asyncio.gather( - *(server.shutdown() for server in servers_snapshot), - return_exceptions=True, - ) - for server, result in zip(servers_snapshot, results): - if isinstance(result, Exception): - logger.debug( - "Error closing MCP server '%s': %s", server.name, result, - ) - with _lock: - for name in selected: - _servers.pop(name, None) - _server_scope_keys.pop(name, None) - # Drop connect-retry cooldowns too: a full shutdown/restart - # should re-attempt every server immediately, not honour a - # stale per-server backoff from before the restart (#50394). - _server_connect_retry_after.clear() - _server_connect_failures.clear() - - with _lock: - loop = _mcp_loop - if loop is not None and loop.is_running(): - from agent.async_utils import safe_schedule_threadsafe - future = safe_schedule_threadsafe( - _shutdown(), loop, - logger=logger, - log_message="MCP shutdown: failed to schedule", - ) - if future is not None: - try: - future.result(timeout=15) - except BaseException as exc: - logger.debug("Error during MCP shutdown: %s", exc) - - # Unconditional final sweep: whether the async ``_shutdown`` ran, - # timed out, or was never scheduled (loop already stopped), a full - # shutdown must leave no stale connect-cooldown state behind — the - # next start should re-attempt every server immediately (#50394). - with _lock: - _server_connect_retry_after.clear() - _server_connect_failures.clear() - - _stop_mcp_loop(only_if_idle=scope is not None) - - -def _kill_orphaned_mcp_children( - include_active: bool = False, - server_name: Optional[str] = None, -) -> None: - """Best-effort graceful shutdown of stdio MCP subprocesses to reap orphans. - - Orphans are PIDs that survived their session context exit (SDK teardown - did not terminate the process — common on Linux when stdio children escape - the parent cgroup on cancellation). By default only entries in - ``_orphan_stdio_pids`` are reaped so concurrent cron jobs and live user - sessions are not disrupted. - - Sends SIGTERM, waits 2 seconds, then escalates to SIGKILL for any - survivors, avoiding shared-resource collisions when multiple hermes - processes run on the same host (each has its own ``_stdio_pids`` dict). - - On POSIX, signals are sent via ``os.killpg`` to the spawn-time pgid when - one is tracked, so reparented grandchildren in the same process group - (e.g. ``claude mcp serve`` spawned by a stdio MCP wrapper that exited - first) are reaped alongside the direct child. Falls back to ``os.kill`` - on Windows and when no pgid is recorded. - - When ``server_name`` is set, only orphaned PIDs known to belong to that - MCP server are reaped. This lets stdio reconnects clean up their previous - transport without touching unrelated servers. - - With ``include_active=True`` also kills every PID in ``_stdio_pids`` — - used only at final shutdown, after the MCP event loop has stopped and no - sessions can still be in flight. - """ - import signal as _signal - - with _lock: - pids: Dict[int, str] = {} - for opid in _orphan_stdio_pids: - owner = _orphan_stdio_pid_servers.get(opid, "orphan") - if server_name is not None and owner != server_name: - continue - pids[opid] = owner - for opid in pids: - _orphan_stdio_pids.discard(opid) - _orphan_stdio_pid_servers.pop(opid, None) - if include_active: - active = dict(_stdio_pids) - if server_name is not None: - active = { - pid: owner - for pid, owner in active.items() - if owner == server_name - } - pids.update(active) - for pid in active: - _stdio_pids.pop(pid, None) - # Snapshot pgids for the pids we're about to kill, then drop the - # entries so a future spawn can't collide with stale state. - pgids: Dict[int, int] = {pid: _stdio_pgids[pid] for pid in pids if pid in _stdio_pgids} - for pid in pgids: - _stdio_pgids.pop(pid, None) - - # Fast path: no tracked stdio PIDs to reap. Skip the SIGTERM/sleep/SIGKILL - # dance entirely — otherwise every MCP-free shutdown pays a 2s sleep tax. - if not pids: - return - - # Pre-compute the gateway's own pgid so _send_signal can avoid killing it. - try: - _my_pgid = os.getpgrp() - except (AttributeError, OSError): - _my_pgid = None # Windows or restricted environment - - def _send_signal(pid: int, sig: int, server_name: str) -> None: - """SIGTERM/SIGKILL via pgroup on POSIX, fall back to pid signal.""" - pgid = pgids.get(pid) - killpg = getattr(os, "killpg", None) - if pgid is not None and killpg is not None: - if _my_pgid is not None and pgid == _my_pgid: - # The MCP child shares the gateway's own process group. - # Using killpg would deliver the signal to the gateway as - # well, crashing it (see #47134). Fall through to the - # per-pid kill() path instead. Warn because per-pid kill - # cannot reach grandchildren in this shared group — if the - # direct child has already exited, they may leak (inherent: - # group-killing them would also kill the gateway). - logger.warning( - "MCP server '%s' pgid %d matches gateway pgid; skipping " - "killpg to avoid self-kill and using per-pid kill — any " - "grandchildren in this group may not be reaped", - server_name, pgid, - ) - else: - try: - killpg(pgid, sig) - return - except (ProcessLookupError, PermissionError, OSError) as exc: - # Pgroup gone (all members exited) or refused — fall back to - # the per-pid path so we still try the direct child if alive. - logger.debug( - "killpg(%d, %d) failed for MCP server '%s': %s; falling back to kill(pid)", - pgid, sig, server_name, exc, - ) - try: - os.kill(pid, sig) - except (ProcessLookupError, PermissionError, OSError): - pass - - # Phase 1: SIGTERM (graceful) - for pid, server_name in pids.items(): - _send_signal(pid, _signal.SIGTERM, server_name) - logger.debug("Sent SIGTERM to orphaned MCP process %d (%s)", pid, server_name) - - # Phase 2: Wait for graceful exit - time.sleep(2) - - # Phase 3: SIGKILL any survivors - _sigkill = getattr(_signal, "SIGKILL", _signal.SIGTERM) - # ``os.kill(pid, 0)`` is NOT a no-op on Windows. Use the cross-platform - # existence check before escalating to SIGKILL. - from gateway.status import _pid_exists - for pid, server_name in pids.items(): - if not _pid_exists(pid): - continue # Good — exited after SIGTERM - _send_signal(pid, _sigkill, server_name) - logger.warning( - "Force-killed MCP process %d (%s) after SIGTERM timeout", - pid, server_name, - ) - - -def _stop_mcp_loop_if_idle() -> bool: - """Stop the MCP loop only when no registered server still owns it. - - Probe paths create temporary MCPServerTask instances that are not placed in - ``_servers``. They should clean up an otherwise-idle loop, but must not - tear down the process-global loop when live agent tools are registered on - it. Otherwise a dashboard/CLI probe can make later MCP tool calls fail - with ``MCP event loop is not running``. - """ - return _stop_mcp_loop(only_if_idle=True) - - -async def _drain_mcp_loop_tasks( - *, - timeout: float = _MCP_LOOP_DRAIN_TIMEOUT, -) -> None: - """Cancel every task still pending on the MCP loop and reap it. - - Cancelling is not enough on its own: ``Task.cancel()`` only schedules the - throw, so tasks need a cancellation cycle before the loop goes away. Wait - for them here — on their owning loop — but keep the final drain bounded so - a task that suppresses cancellation cannot hang process exit indefinitely. - """ - current = asyncio.current_task() - pending = [t for t in asyncio.all_tasks() if t is not current and not t.done()] - if not pending: - return - logger.debug("Draining %d pending task(s) from the MCP loop", len(pending)) - for task in pending: - task.cancel() - - done, still_pending = await asyncio.wait(pending, timeout=timeout) - for task in done: - if task.cancelled(): - continue - try: - task.exception() - except asyncio.CancelledError: - pass - except Exception as exc: - logger.debug("Pending MCP loop task ended during shutdown: %s", exc) - - if still_pending: - logger.warning( - "%d MCP loop task(s) still pending after %.1fs drain", - len(still_pending), timeout, - ) - - -async def _drain_and_stop_mcp_loop() -> None: - """Drain pending tasks, then stop the loop from its owning thread. - - Keeping both operations in one loop-owned sequence matters when the caller - times out waiting for a blocked loop. Queuing ``loop.stop`` separately from - the caller can overtake the scheduled drain before it receives a loop cycle, - leaving the drain coroutine itself pending when the loop is closed. - """ - loop = asyncio.get_running_loop() - try: - await _drain_mcp_loop_tasks(timeout=_MCP_LOOP_DRAIN_TIMEOUT) - finally: - loop.call_soon(loop.stop) - - def _stop_mcp_loop(*, only_if_idle: bool = False) -> bool: """Stop the background event loop and join its thread.""" global _mcp_loop, _mcp_thread @@ -8973,11 +2321,9 @@ def _stop_mcp_loop(*, only_if_idle: bool = False) -> bool: _mcp_loop = None _mcp_thread = None if loop is not None: - # Drain before stopping: closing the loop with tasks still suspended - # leaves their coroutines for the GC, whose finalizer then resumes them - # to run cleanup against a loop that is already closed -> "Event loop - # is closed" (#60197). ``shutdown_mcp_servers`` only reaps servers held - # in ``_servers``, so anything else left on this loop 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 @@ -9017,8 +2363,6 @@ def _stop_mcp_loop(*, only_if_idle: bool = False) -> bool: loop.close() except Exception as exc: logger.warning("Unable to close MCP event loop cleanly: %s", exc) - # After closing the loop, any stdio subprocesses that survived the - # graceful shutdown are now orphaned — include active PIDs too - # since the loop is gone and no session can still be in flight. + # The loop is gone, so no session can be in flight: reap active too. _kill_orphaned_mcp_children(include_active=True) return True diff --git a/tools/mcp_tool_agent.py b/tools/mcp_tool_agent.py new file mode 100644 index 0000000000..e7a509c7f3 --- /dev/null +++ b/tools/mcp_tool_agent.py @@ -0,0 +1,291 @@ +"""Live-agent tool-list maintenance after MCP (re)discovery: refreshing an +AIAgent's tools/tool names, preserving the cached tools[] prefix across rebuilds, +and re-injecting post-build tools.""" + +import logging +import json +import threading +from tools.mcp_tool_common import _core + +logger = logging.getLogger("tools.mcp_tool") + + +# Serializes in-place swaps of ``agent.tools`` / ``agent.valid_tool_names`` by +# the reload RPC, gateway reload and late-binding refresh thread; the run loop +# reads them during tool iteration and must never see a half-updated pair. +_agent_tools_lock = threading.Lock() + + +def refresh_agent_mcp_tools( + agent, + *, + enabled_override=None, + disabled_override=None, + quiet_mode: bool = True, + content_aware: bool = False, + preserve_prefix: bool = False, +) -> set: + """Re-derive an already-built agent's tool snapshot from the live registry. + + The agent snapshots ``agent.tools`` once at build time; servers that connect + later (slow OAuth server, ``/reload-mcp``) are invisible until rebuilt. This + is the single shared rebuild for the TUI RPC, gateway reload, late-binding + thread and between-turns refresh. It respects the agent's own toolset + filter, diffs by tool NAME (a count compare misses an equal-size swap), and + is additive-preserving: memory-provider and context-engine (``lcm_*``) tools + that ``agent_init`` appends after ``get_tool_definitions`` are re-injected, + since a naive rebuild would silently delete them. ``(tools, valid_tool_names)`` + are published together under ``_agent_tools_lock``. + + ``preserve_prefix`` is for rebuilds inside a live conversation, where the + tool array is a cached request prefix and any moved byte re-prefills the + whole history: existing tools keep their slot (schemas still refresh), a + still-registered tool whose ``check_fn`` merely flapped is carried forward, + a tool that left the registry is dropped, and new tools append at the tail. + Carrying an unavailable tool forward is safe: ``check_fn`` gates exposure, + never invocation, and every handler owns its own unavailability error. + + Returns the newly-added tool names (empty when unchanged). The caller owns + the prompt-cache contract (turn-boundary policy differs per caller). + """ + from model_tools import get_tool_definitions + from tools.registry import registry + + # Explicit reloads pass freshly-resolved toolsets (so a server just ENABLED + # in config is picked up) and the agent's selection is updated to match; + # automatic paths pass nothing and reuse the build-time selection. + if enabled_override is not None or disabled_override is not None: + enabled = enabled_override if enabled_override is not None else getattr(agent, "enabled_toolsets", None) + disabled = disabled_override if disabled_override is not None else getattr(agent, "disabled_toolsets", None) + agent.enabled_toolsets = enabled + agent.disabled_toolsets = disabled + else: + enabled = getattr(agent, "enabled_toolsets", None) + disabled = getattr(agent, "disabled_toolsets", None) + + # Capture the registry generation BEFORE the slow get_tool_definitions call; + # at publish time a slower caller holding an OLDER set must not clobber a + # newer set another caller already published. + snapshot_generation = registry._generation + + # Computed OUTSIDE the lock (can be slow); diff + publish happen together in + # one critical section so concurrent callers can't torn-publish. + new_defs = list( + get_tool_definitions( + enabled_toolsets=enabled, + disabled_toolsets=disabled, + quiet_mode=quiet_mode, + ) + or [] + ) + new_names = {t["function"]["name"] for t in new_defs} + + # Re-append the post-build families on LOCALS only; live agent attributes + # are untouched until the single atomic publish below. + staged_engine_names = _core._reinject_post_build_tools(agent, new_defs, new_names) + + # Registry membership is read OUTSIDE ``_agent_tools_lock``: taking + # ``registry._lock`` under the tools lock would be the first nesting of the two. + registered_names: set = set() + if preserve_prefix: + try: + registered_names = {entry.name for entry in registry.get_all_entries()} + except Exception: # noqa: BLE001 + preserve_prefix = False # fail open to the plain rebuild + + # Single atomic read-diff-publish so ``added`` matches what was published + # and a stale (older-generation) rebuild can't overwrite a newer one. + with _agent_tools_lock: + # Tolerate an agent that never set the generation (or a non-int mock) + # rather than failing the whole refresh on the comparison. + published_gen_raw = getattr(agent, "_tool_snapshot_generation", -1) + published_gen = published_gen_raw if isinstance(published_gen_raw, int) else -1 + if snapshot_generation < published_gen: + return set() # a newer snapshot already won + current_defs = list(getattr(agent, "tools", None) or []) + current = {t["function"]["name"] for t in current_defs} + if preserve_prefix: + new_defs, new_names = _merge_preserving_prefix( + current_defs, new_defs, registered_names, + ) + if new_names == current: + # Same NAME set: no change for MCP-reload callers. Content-aware + # callers (compaction boundary) also diff serialized bytes, since + # dynamic schemas change CONTENT under stable names. + content_changed = False + if content_aware: + try: + _stable = json.dumps( + (getattr(agent, "tools", None) or []), + sort_keys=True, separators=(",", ":"), default=str, + ) + _new = json.dumps( + new_defs, sort_keys=True, separators=(",", ":"), + default=str, + ) + content_changed = _stable != _new + except Exception: # noqa: BLE001 + content_changed = False + if not content_changed: + # Record the generation so an in-flight older caller can't clobber. + agent._tool_snapshot_generation = max(published_gen, snapshot_generation) + return set() + agent.tools = new_defs + agent.valid_tool_names = new_names + # Publish context-engine routing names atomically with the snapshot. + engine_names = getattr(agent, "_context_engine_tool_names", None) + if isinstance(engine_names, set): + engine_names.clear() + engine_names.update(staged_engine_names) + agent._tool_snapshot_generation = max(published_gen, snapshot_generation) + added = new_names - current + # Re-pin the session's tool order so a rebuild-for-existing-session + # (gateway agent-cache eviction) restores exactly these names. + persist_agent_tool_names(agent) + return added + + +def reprobe_tool_availability() -> None: + """Explicit ``/reload-mcp`` hatch out of the tools[] freeze: drop the + ``check_fn`` verdict cache AND the ``get_tool_definitions`` memo (keyed on + registry generation, so it would otherwise replay the stale verdicts).""" + from model_tools import _clear_tool_defs_cache + from tools.registry import invalidate_check_fn_cache + + invalidate_check_fn_cache() + _clear_tool_defs_cache() + + +def persist_agent_tool_names(agent) -> None: + """Best-effort: write ``agent.tools`` names to the session row (freeze pin).""" + db = getattr(agent, "_session_db", None) + session_id = getattr(agent, "session_id", None) + if not db or not session_id: + return + try: + db.update_session_tool_names( + session_id, + [t["function"]["name"] for t in (getattr(agent, "tools", None) or [])], + ) + except Exception: # noqa: BLE001 + logger.debug("tool_names persist skipped", exc_info=True) + + +def restore_agent_tool_prefix(agent, saved_names: list) -> bool: + """Fold a freshly built agent's ``tools`` onto the session's saved order. + + The gateway rebuilds a NEW AIAgent for an existing session after agent-cache + eviction, with no predecessor to preserve; the saved name list stands in: + a saved tool still registered but failing its probe is carried forward from + the registry schema, a deregistered one is dropped, new tools append at the + tail (same rule as ``_merge_preserving_prefix``). Returns True if changed. + """ + if not saved_names: + return False + from tools.registry import registry + + fresh_defs = list(getattr(agent, "tools", None) or []) + fresh = {t["function"]["name"]: t for t in fresh_defs} + saved_defs = [] + for name in saved_names: + entry_def = fresh.get(name) + if entry_def is None: + entry = registry.get_entry(name) + if entry is None: + continue + entry_def = {"type": "function", "function": {**entry.schema, "name": entry.name}} + saved_defs.append(entry_def) + registered_names = {entry.name for entry in registry.get_all_entries()} + merged, merged_names = _merge_preserving_prefix(saved_defs, fresh_defs, registered_names) + with _agent_tools_lock: + if merged == fresh_defs: + return False + agent.tools = merged + agent.valid_tool_names = merged_names + if [t["function"]["name"] for t in merged] != list(saved_names): + persist_agent_tool_names(agent) + return True + + +def _merge_preserving_prefix( + current_defs: list, new_defs: list, registered_names: set, +) -> tuple[list, set]: + """Fold a fresh tool snapshot into a live one without moving existing bytes. + + Ordered by ``current_defs`` (the cached request prefix): a name in both + keeps its slot but takes the fresh schema; a name only in the live list is + kept if still registered (``check_fn`` flapped) and dropped if not; a name + only in the fresh list is appended at the tail. + """ + fresh = {} + for entry in new_defs: + name = (entry.get("function") or {}).get("name", "") + if name: + fresh[name] = entry + + merged = [] + for entry in current_defs: + name = (entry.get("function") or {}).get("name", "") + replacement = fresh.pop(name, None) + if replacement is not None: + merged.append(replacement) + elif name and name in registered_names: + merged.append(entry) + merged.extend(fresh.values()) + return merged, {(t.get("function") or {}).get("name", "") for t in merged} + + +def _reinject_post_build_tools(agent, tools_list: list, name_set: set) -> set: + """Append memory-provider and context-engine tools onto the caller's staged + ``tools_list`` / ``name_set`` (never the live agent attributes), mirroring + ``agent_init``'s post-build injection. Idempotent and fail-soft. + + Returns the context-engine routing names THIS rebuild appended: a name + already owned by a registry/plugin tool is not claimed, matching agent_init. + """ + def _add(schema: dict) -> bool: + name = schema.get("name", "") + if not name or name in name_set: + return False + tools_list.append({"type": "function", "function": schema}) + name_set.add(name) + return True + + try: + memory_manager = getattr(agent, "_memory_manager", None) + get_mem_schemas = getattr(memory_manager, "get_all_tool_schemas", None) if memory_manager else None + if callable(get_mem_schemas): + # Same toolset gate inject_memory_provider_tools uses. + from agent.memory_manager import memory_provider_tools_enabled + if memory_provider_tools_enabled( + getattr(agent, "enabled_toolsets", None), + getattr(agent, "disabled_toolsets", None), + memory_tool_present="memory" in name_set, + ): + for schema in get_mem_schemas(): + if isinstance(schema, dict): + _add(schema) + except Exception: + logger.debug("Memory-provider tool re-injection skipped", exc_info=True) + + # The `context_engine` toolset is intentionally empty, so lcm_* tools exist + # only via this append. Honor the enabled_toolsets gate agent_init uses, or a + # restricted-toolset platform would re-leak tools the build excluded. + staged_engine_names: set = set() + try: + enabled = getattr(agent, "enabled_toolsets", None) + context_engine_allowed = enabled is None or "context_engine" in enabled + compressor = getattr(agent, "context_compressor", None) + get_schemas = getattr(compressor, "get_tool_schemas", None) if compressor else None + if context_engine_allowed and callable(get_schemas): + for schema in get_schemas(): + if not isinstance(schema, dict): + continue + name = schema.get("name", "") + # Claim the routing name only when WE appended the schema. + if _add(schema) and name: + staged_engine_names.add(name) + except Exception: + logger.debug("Context-engine tool re-injection skipped", exc_info=True) + + return staged_engine_names diff --git a/tools/mcp_tool_common.py b/tools/mcp_tool_common.py new file mode 100644 index 0000000000..d0bc4f3438 --- /dev/null +++ b/tools/mcp_tool_common.py @@ -0,0 +1,181 @@ +"""Small pure helpers shared by the tools.mcp_tool_* modules: SDK 1.x/2.x field +access, error-text sanitising, numeric/bool coercion, timeouts and jitter. No +origin state.""" + +import logging +import math +import os +import random +import re +from typing import Any, Optional + +logger = logging.getLogger("tools.mcp_tool") + + +class _OriginProxy: + """Attribute proxy for ``tools.mcp_tool`` resolved at access time. + + The split modules read origin state (``_servers``, ``_lock``, SDK symbols, + patchable helpers) through this so ``mock.patch("tools.mcp_tool.X")`` and + origin-side rebinds stay effective, and so no split module needs the origin + imported first (the origin imports them while it is still initialising). + """ + + __slots__ = () + + def __getattr__(self, name: str): + from tools import mcp_tool + + return getattr(mcp_tool, name) + + +_core = _OriginProxy() + + +_MISSING = object() + + +def mcp_field(obj, snake: str, camel: str, default=None): + """Read an MCP model field across the 1.x -> 2.x rename to snake_case. + + Pydantic aliases don't apply to attribute access, so ``getattr(result, + "isError", False)`` silently returns the default on 2.x — failed calls read + as successful, schemas as empty. Trying both spellings stays correct on + either SDK generation (``mcp`` is an optional extra at the user's version). + """ + value = getattr(obj, snake, _MISSING) + if value is not _MISSING: + return value + value = getattr(obj, camel, _MISSING) + return default if value is _MISSING else value + + +_DEFAULT_TOOL_TIMEOUT = 300 # seconds for tool calls + + +def _resolve_tool_timeout(config: dict) -> float: + """Per-server tool-call timeout. Precedence: ``mcp_servers..timeout`` + > ``timeouts.mcp.tool_call`` > the 300s default; values are platform-clamped + by ``resolve_timeout``.""" + per_server = config.get("timeout") + if per_server is not None: + return per_server + try: + from agent.deadline import resolve_timeout + + resolved = resolve_timeout("mcp.tool_call", default=_DEFAULT_TOOL_TIMEOUT) + if resolved is not None: + return resolved + except Exception: + logger.debug("mcp.tool_call timeout resolution failed", exc_info=True) + return _DEFAULT_TOOL_TIMEOUT + + +# Jitter on reconnect backoff so servers that lost the same backend don't +# retry in lockstep (thundering herd, synchronized log bursts). +_BACKOFF_JITTER = 0.2 # +/-20% + + +def _jittered(seconds: float) -> float: + """``seconds`` with +/-20% uniform jitter, floored at 0.""" + return max(0.0, seconds * random.uniform(1.0 - _BACKOFF_JITTER, + 1.0 + _BACKOFF_JITTER)) + + +# Credential patterns to strip from error messages. +_CREDENTIAL_PATTERN = re.compile( + r"(?:" + r"ghp_[A-Za-z0-9_]{1,255}" # GitHub PAT + r"|sk-[A-Za-z0-9_]{1,255}" # OpenAI-style key + r"|Bearer\s+\S+" # Bearer token + r"|token=[^\s&,;\"']{1,255}" # token=... + r"|key=[^\s&,;\"']{1,255}" # key=... + r"|API_KEY=[^\s&,;\"']{1,255}" # API_KEY=... + r"|password=[^\s&,;\"']{1,255}" # password=... + r"|secret=[^\s&,;\"']{1,255}" # secret=... + r")", + re.IGNORECASE, +) + + +def _env_ref_name(ref: str) -> str: + """Bare env-var name from a ``${...}`` body; strips a Cursor-style ``env:`` prefix.""" + ref = ref.strip() + if ref.startswith("env:"): + ref = ref[len("env:"):].strip() + return ref + + +def _sanitize_error(text: str) -> str: + """Replace credential-like patterns with [REDACTED] before text reaches the LLM.""" + return _CREDENTIAL_PATTERN.sub("[REDACTED]", text) + + +def _exc_str(exc: BaseException) -> str: + """Non-empty string for *exc*: some exceptions (``anyio.ClosedResourceError``) + carry no message, so fall back to ``repr`` to keep diagnostics.""" + text = str(exc).strip() + return text or repr(exc) + + +def _prepend_path(env: dict, directory: str) -> dict: + """Prepend *directory* to env PATH if it is not already present.""" + updated = dict(env or {}) + if not directory: + return updated + + existing = updated.get("PATH", "") + parts = [part for part in existing.split(os.pathsep) if part] + if directory not in parts: + parts = [directory, *parts] + updated["PATH"] = os.pathsep.join(parts) if parts else directory + return updated + + +def _safe_numeric(value, default, coerce=int, minimum=1): + """Coerce a config value (YAML strings included) to a number, clamped to + *minimum*; *default* on failure or non-finite floats.""" + try: + result = coerce(value) + if isinstance(result, float) and not math.isfinite(result): + return default + return max(result, minimum) + except (TypeError, ValueError, OverflowError): + return default + + +def _parse_boolish(value: Any, default: bool = True) -> bool: + """Parse a bool-like config value with safe fallback.""" + if value is None: + return default + if isinstance(value, bool): + return value + if isinstance(value, str): + lowered = value.strip().lower() + if lowered in {"true", "1", "yes", "on"}: + return True + if lowered in {"false", "0", "no", "off"}: + return False + logger.warning("MCP config expected a boolean-ish value, got %r; using default=%s", value, default) + return default + + +def _get_lifecycle_seconds(config: dict, key: str) -> Optional[float]: + """Return an optional positive lifecycle timeout from top-level/nested config.""" + raw = config.get(key) + lifecycle = config.get("lifecycle") + if raw is None and isinstance(lifecycle, dict): + raw = lifecycle.get(key) + if raw is None: + return None + try: + seconds = float(raw) + except (TypeError, ValueError): + logger.warning("MCP config %s must be a number of seconds; ignoring %r", key, raw) + return None + if seconds == 0: + return None + if seconds < 0: + logger.warning("MCP config %s must be positive; ignoring %r", key, raw) + return None + return seconds diff --git a/tools/mcp_tool_config.py b/tools/mcp_tool_config.py new file mode 100644 index 0000000000..80fd0dff2b --- /dev/null +++ b/tools/mcp_tool_config.py @@ -0,0 +1,382 @@ +"""MCP server config loading and stdio launch environment: ${VAR}/Cursor-style +interpolation, hidden-whitespace and suspicious-entry filtering, the filtered +subprocess env, command resolution, watchdog wrapping and the shared stderr log.""" + +import logging +import os +import re +import shutil +import sys +import threading +from typing import Callable +from datetime import datetime +from typing import Any, Dict, List, Optional, Set, Tuple +from tools.mcp_tool_common import _env_ref_name, _prepend_path, _core + +logger = logging.getLogger("tools.mcp_tool") + + +_mcp_stderr_log_fh: Optional[Any] = None + + +_mcp_stderr_log_lock = threading.Lock() + + +def _get_mcp_stderr_log() -> Any: + """Shared append-mode handle for MCP subprocess stderr, opened once per process. + + Must expose a real fd (``fileno()``) because asyncio wires the child's stderr + directly to it. Falls back to ``/dev/null``, then the real stderr. + """ + global _mcp_stderr_log_fh + with _mcp_stderr_log_lock: + if _mcp_stderr_log_fh is not None: + return _mcp_stderr_log_fh + try: + from hermes_constants import get_hermes_home + log_dir = get_hermes_home() / "logs" + log_dir.mkdir(parents=True, exist_ok=True) + log_path = log_dir / "mcp-stderr.log" + # Line-buffered so output lands promptly; errors="replace" tolerates + # garbled binary from misbehaving servers. + fh = open(log_path, "a", encoding="utf-8", errors="replace", buffering=1) + fh.fileno() # confirm a real fd before committing + _mcp_stderr_log_fh = fh + except Exception as exc: # pragma: no cover — best-effort fallback + logger.debug("Failed to open MCP stderr log, using devnull: %s", exc) + try: + _mcp_stderr_log_fh = open(os.devnull, "w", encoding="utf-8") + except Exception: + _mcp_stderr_log_fh = sys.stderr + return _mcp_stderr_log_fh + + +def _write_stderr_log_header(server_name: str) -> None: + """Write a session marker so operators can find each server's output in the + shared log without per-line prefixes (which would need a pipe + reader thread).""" + fh = _core._get_mcp_stderr_log() + try: + ts = datetime.now().strftime("%Y-%m-%d %H:%M:%S") + fh.write(f"\n===== [{ts}] starting MCP server '{server_name}' =====\n") + fh.flush() + except Exception: + pass + + +# Env vars safe to pass to stdio subprocesses (no secrets). +_SAFE_ENV_KEYS = frozenset({ + "PATH", "HOME", "USER", "LANG", "LC_ALL", "TERM", "SHELL", "TMPDIR", +}) + + +_SAFE_ENV_KEYS_CASE_INSENSITIVE = frozenset({ + # Windows process/location vars needed by launcher-style tools (e.g. + # Docker Desktop's MCP plugin discovery); none carry secrets. + "ALLUSERSPROFILE", + "APPDATA", + "COMMONPROGRAMFILES", + "COMMONPROGRAMFILES(X86)", + "COMMONPROGRAMW6432", + "COMPUTERNAME", + "COMSPEC", + "HOMEDRIVE", + "HOMEPATH", + "LOCALAPPDATA", + "NUMBER_OF_PROCESSORS", + "OS", + "PATHEXT", + "PROCESSOR_ARCHITECTURE", + "PROGRAMDATA", + "PROGRAMFILES", + "PROGRAMFILES(X86)", + "PROGRAMW6432", + "PUBLIC", + "SYSTEMDRIVE", + "SYSTEMROOT", + "TEMP", + "TMP", + "USERDOMAIN", + "USERNAME", + "USERPROFILE", + "WINDIR", +}) + + +# ${VAR_NAME} interpolation; any non-} chars allowed so MY-VAR / my.var work. +_ENV_VAR_PATTERN = re.compile(r"\$\{([^}]+)\}") + + +def _workspace_folder() -> str: + """Absolute workspace root for ``${workspaceFolder}``: the session's + authoritative root (terminal cwd / task override / $TERMINAL_CWD), else cwd.""" + try: + from tools.file_tools import _authoritative_workspace_root + + root = _authoritative_workspace_root() + if root: + return root + except Exception: + pass + return os.getcwd() + + +def _context_var_value(ref: str) -> Optional[str]: + """Resolve Cursor's case-sensitive context vars (``userHome``, + ``workspaceFolder``, ``workspaceFolderBasename``, ``pathSeparator``/``/``). + Returns None for anything else so it falls through to env-var lookup.""" + if ref == "userHome": + return os.path.expanduser("~") + if ref == "workspaceFolder": + return _core._workspace_folder() + if ref == "workspaceFolderBasename": + root = _core._workspace_folder() + return os.path.basename(root.rstrip("/\\")) or root + if ref in ("pathSeparator", "/"): + return os.sep + return None + + +def _build_safe_env(user_env: Optional[dict]) -> dict: + """Filtered env for stdio subprocesses so API keys/tokens don't leak. + + Passes only the safe baseline keys, ``XDG_*``, vars injected by an external + secret source (users configured that backend precisely so subprocesses can + consume them), plus the server config's own ``env``. + """ + try: + from hermes_cli.env_loader import get_secret_source + except Exception: # pragma: no cover — early bootstrap/import fallback + get_secret_source = None + env = {} + for key, value in os.environ.items(): + if ( + key in _SAFE_ENV_KEYS + or key.upper() in _SAFE_ENV_KEYS_CASE_INSENSITIVE + or key.startswith("XDG_") + or (get_secret_source is not None and get_secret_source(key)) + ): + env[key] = value + if user_env: + env.update(user_env) + return env + + +def _resolve_stdio_command(command: str, env: dict) -> tuple[str, dict]: + """Resolve a stdio command against the exact subprocess env, mainly so bare + ``npx``/``npm``/``node`` work under a filtered PATH.""" + resolved_command = os.path.expanduser(str(command).strip()) + resolved_env = dict(env or {}) + + if os.sep not in resolved_command: + path_arg = resolved_env.get("PATH") + which_hit = shutil.which(resolved_command, path=path_arg) + if which_hit is None and sys.platform == "win32" and resolved_env: + # shutil.which(path=...) uses the PARENT's PATHEXT, not the config + # env's, so retry with the config's PATHEXT (any key casing) applied. + cfg_pathext = next( + (v for k, v in resolved_env.items() + if k.upper() == "PATHEXT" and isinstance(v, str) and v.strip()), + None, + ) + if cfg_pathext and cfg_pathext != os.environ.get("PATHEXT"): + _saved = os.environ.get("PATHEXT") + try: + os.environ["PATHEXT"] = cfg_pathext + which_hit = shutil.which(resolved_command, path=path_arg) + finally: + if _saved is None: + os.environ.pop("PATHEXT", None) + else: + os.environ["PATHEXT"] = _saved + if which_hit: + resolved_command = which_hit + elif resolved_command in {"npx", "npm", "node"}: + hermes_home = os.path.expanduser( + os.getenv( + "HERMES_HOME", os.path.join(os.path.expanduser("~"), ".hermes") + ) + ) + candidates = [ + os.path.join(hermes_home, "node", "bin", resolved_command), + os.path.join(os.path.expanduser("~"), ".local", "bin", resolved_command), + # Canonical Node location for from-source Linux builds, the + # Hermes Docker image and Intel Homebrew. Needed when a user's + # hand-authored env.PATH omits it: npx's shebang re-execs + # /usr/bin/env node, so a symlink workaround fails one layer deeper. + os.path.join(os.sep, "usr", "local", "bin", resolved_command), + ] + for candidate in candidates: + if os.path.isfile(candidate) and os.access(candidate, os.X_OK): + resolved_command = candidate + break + + command_dir = os.path.dirname(resolved_command) + if command_dir: + resolved_env = _prepend_path(resolved_env, command_dir) + + return resolved_command, resolved_env + + +def _wrap_command_with_watchdog(command: str, args: list) -> tuple[str, list]: + """Wrap a stdio command in the parent-death watchdog (POSIX only; the + watchdog polls ``getppid()`` against our PID). Unchanged on non-POSIX or + if the PID cannot be read — watchdog bookkeeping must never block a connection.""" + if os.name != "posix": + # Relies on process groups (getpgid/killpg), same scope as the + # killpg-based orphan cleanup. + return command, args + try: + my_pid = os.getpid() + except Exception: + return command, args + watchdog_args = [ + os.path.join(os.path.dirname(os.path.abspath(__file__)), "mcp_stdio_watchdog.py"), + "--ppid", str(my_pid), + "--", + command, + *args, + ] + return sys.executable, watchdog_args + + +def _interpolate_env_vars(value): + """Recursively resolve ``${VAR}`` / Cursor ``${env:VAR}`` placeholders plus the + Cursor context vars (see ``_context_var_value``). + + Env refs resolve from the active profile's secret scope when multiplexing + (so ``${API_KEY}`` picks up the routed profile's value, not another + profile's in ``os.environ``). Unset vars keep the literal placeholder. + """ + from agent.secret_scope import get_secret as _get_secret + + if isinstance(value, str): + def _replace(m): + ctx = _context_var_value(m.group(1).strip()) + if ctx is not None: + return ctx + name = _env_ref_name(m.group(1)) + return _get_secret(name, m.group(0)) or m.group(0) + return _ENV_VAR_PATTERN.sub(_replace, value) + if isinstance(value, dict): + return {k: _interpolate_env_vars(v) for k, v in value.items()} + if isinstance(value, list): + return [_interpolate_env_vars(v) for v in value] + return value + + +# (server_name, dotted key path) pairs already warned about; config loads +# happen on every discovery pass, so warn once per process. +_whitespace_warned: Set[Tuple[str, str]] = set() + + +def _warn_hidden_whitespace(server_name: str, config: dict) -> List[str]: + """Warn once per (server, key path) about string values with leading/trailing + whitespace — a pasted trailing newline or leading space causes opaque + auth/connect failures and is invisible in config.yaml. + + Advisory only: values are never mutated (whitespace could be intentional) + and never logged (often secrets). Returns the flagged key paths. + """ + flagged: List[str] = [] + + def _walk(value: Any, path: str) -> None: + if isinstance(value, str): + if value != value.strip(): + flagged.append(path) + elif isinstance(value, dict): + for k, v in value.items(): + _walk(v, f"{path}.{k}" if path else str(k)) + elif isinstance(value, list): + for i, v in enumerate(value): + _walk(v, f"{path}[{i}]") + + _walk(config, "") + for key_path in flagged: + dedupe_key = (server_name, key_path) + if dedupe_key in _whitespace_warned: + continue + _whitespace_warned.add(dedupe_key) + logger.warning( + "MCP server '%s': config value '%s' has hidden leading or " + "trailing whitespace — this often causes authentication or " + "connection failures. Check for stray spaces/newlines in " + "config.yaml (or the referenced env var).", + server_name, + key_path, + ) + return flagged + + +def _filter_suspicious_mcp_servers(servers: Dict[str, dict]) -> Dict[str, dict]: + """Drop exfiltration-shaped MCP configs before any stdio spawn path.""" + try: + from hermes_cli.mcp_security import validate_mcp_server_entry as _validate_mcp_server_entry + except Exception: + _validate_mcp_server_entry: Callable[[str, dict[str, Any]], list[str]] | None = None + + if _validate_mcp_server_entry is None: + return servers + + safe_servers = {} + for name, cfg in servers.items(): + if not isinstance(cfg, dict): + safe_servers[name] = cfg + continue + issues = _validate_mcp_server_entry(name, cfg) + if issues: + logger.warning( + "Skipping suspicious MCP server '%s': %s", + name, + "; ".join(issues), + ) + continue + safe_servers[name] = cfg + return safe_servers + + +def _load_mcp_config() -> Dict[str, dict]: + """Read ``mcp_servers`` from config.yaml as ``{name: config}`` (empty on error + or in safe mode). Entries carry ``command``/``args``/``env`` (stdio) or + ``url``/``headers`` (HTTP) plus optional timeout/auth keys; ``${VAR}`` + placeholders are interpolated after ``.env`` is loaded.""" + try: + from hermes_cli.config import load_config + from utils import env_var_enabled as _env_enabled + + if _env_enabled("HERMES_SAFE_MODE"): + return {} + config = load_config() + servers = config.get("mcp_servers") + if not isinstance(servers, dict): + servers = {} + # Ensure .env vars are available for interpolation + try: + from hermes_cli.env_loader import load_hermes_dotenv + load_hermes_dotenv() + except Exception: + pass + safe_servers: Dict[str, dict] = {} + for name, cfg in _core._filter_suspicious_mcp_servers(servers).items(): + interpolated = _interpolate_env_vars(cfg) + if isinstance(interpolated, dict): + _warn_hidden_whitespace(name, interpolated) + safe_servers[name] = interpolated + try: + from hermes_cli.plugins import discover_plugins, get_plugin_manager + + discover_plugins() + portable = get_plugin_manager().get_portable_mcp_servers() + for name, cfg in _core._filter_suspicious_mcp_servers(portable).items(): + if name in safe_servers: + logger.warning( + "Portable MCP server '%s' conflicts with native config; skipping", + name, + ) + continue + safe_servers[name] = dict(cfg) + except Exception: + logger.debug("Failed to load portable MCP servers", exc_info=True) + return safe_servers + except Exception as exc: + logger.debug("Failed to load MCP config: %s", exc) + return {} diff --git a/tools/mcp_tool_content.py b/tools/mcp_tool_content.py new file mode 100644 index 0000000000..c1ebd534ba --- /dev/null +++ b/tools/mcp_tool_content.py @@ -0,0 +1,254 @@ +"""Rendering of MCP tool-result content blocks into model-facing text: size +capping, _meta filtering, image/audio caching to MEDIA tags, resource links and +embedded resources.""" + +import logging +from tools.ansi_strip import strip_unicode_tags +from tools.mcp_tool_common import mcp_field, _core +from tools.mcp_tool_schema import mcp_prefixed_tool_name + +logger = logging.getLogger("tools.mcp_tool") + + +# Hard allocation ceiling for one MCP text payload (chars): the first line of +# defense against a multi-megabyte flood being JSON-encoded and handed +# downstream. Deliberately far ABOVE the budget layer's 50K spillover threshold +# so ordinary large results reach spillover intact; only pathological floods +# are lossy-truncated here. +_MCP_HARD_RESULT_CAP_CHARS = 2_000_000 + + +def _truncate_mcp_text_result(text: str, max_chars: int = _MCP_HARD_RESULT_CAP_CHARS) -> str: + """Pass text at or under ``max_chars`` unchanged; otherwise keep a 40% head / + 60% tail split with an omission notice between.""" + if len(text) <= max_chars: + return text + head_chars = int(max_chars * 0.4) + tail_chars = max_chars - head_chars + omitted = len(text) - head_chars - tail_chars + return ( + text[:head_chars] + + f"\n\n... [MCP RESULT TRUNCATED - {omitted:,} chars omitted " + f"out of {len(text):,} total] ...\n\n" + + text[-tail_chars:] + ) + + +def _is_reserved_mcp_meta_key(key: str) -> bool: + """True if an MCP ``_meta`` key uses a protocol-reserved prefix: a + ``modelcontextprotocol`` or ``mcp`` label followed by at least one more + label. A trailing one (``com.example.mcp/...``) is a vendor namespace.""" + slash = key.find("/") + if slash <= 0: + return False + labels = key[:slash].split(".") + return any( + label in ("modelcontextprotocol", "mcp") and i < len(labels) - 1 + for i, label in enumerate(labels) + ) + + +def _strip_reserved_meta_keys(meta) -> "Optional[Dict[str, Any]]": + """Drop protocol-reserved keys from ``_meta``; None if nothing model-facing + remains or the input wasn't a mapping.""" + if not isinstance(meta, dict): + return None + out = {k: v for k, v in meta.items() + if isinstance(k, str) and not _is_reserved_mcp_meta_key(k)} + return out or None + + +def _mcp_image_extension_for_mime_type(mime_type: str) -> str: + """File extension for an MCP image MIME type (``.png`` fallback).""" + import mimetypes + normalized = (mime_type or "").split(";", 1)[0].strip().lower() + if normalized in {"image/jpeg", "image/jpg"}: + return ".jpg" + return mimetypes.guess_extension(normalized) or ".png" + + +def _cache_mcp_image_block(block) -> str: + """Cache an ``ImageContent`` block and return a ``MEDIA:`` tag. + + Returns "" (logging, not raising) when the block isn't an image, the base64 + is malformed, or the cache rejects the bytes: one bad block must not kill + the tool result, and the caller falls through to any text blocks. + """ + import base64 + + data = getattr(block, "data", None) + mime_type = mcp_field(block, "mime_type", "mimeType") + normalized_mime = str(mime_type or "").split(";", 1)[0].strip().lower() + if data is None or not normalized_mime.startswith("image/"): + return "" + + try: + raw_bytes = base64.b64decode(data) + except (TypeError, ValueError) as exc: + logger.warning("MCP image block decode failed (%s): %s", normalized_mime, exc) + return "" + + try: + from gateway.platforms.base import cache_image_from_bytes + + image_path = cache_image_from_bytes( + raw_bytes, + ext=_mcp_image_extension_for_mime_type(normalized_mime), + ) + except ImportError: + # gateway.platforms.base unavailable (e.g. cron without gateway deps): + # drop silently, callers get any text blocks that parsed. + logger.debug("MCP image caching skipped — gateway.platforms.base unavailable") + return "" + except Exception as exc: + logger.warning("MCP image block cache failed: %s", exc) + return "" + + return f"MEDIA:{image_path}" + + +# Hard cap on decoded resource bytes from one block, so a misbehaving server +# can't fill the cache disk. +_MCP_RESOURCE_MAX_BYTES = 50 * 1024 * 1024 + + +# Base64 expands ~4/3; reject oversized payloads BEFORE decoding so a multi-GB +# blob string is never transiently doubled in memory. +_MCP_RESOURCE_MAX_B64_CHARS = _MCP_RESOURCE_MAX_BYTES * 4 // 3 + 4 + + +def _mcp_resource_filename(uri: str, mime_type: str) -> str: + """Safe display filename from the URI's last path segment, used only as a + name hint: ``cache_document_from_bytes`` re-sanitizes and prefixes it, so + remote path components can't steer the cache location.""" + import mimetypes + import re as _re + from pathlib import Path + from urllib.parse import urlparse, unquote + + name = "" + if uri: + try: + name = Path(unquote(urlparse(str(uri)).path or "")).name + except (ValueError, TypeError): + name = "" + # Strip control chars (hostile URIs could inject newlines/ANSI into the + # filename and transcript marker) and cap length, preserving the extension. + name = _re.sub(r"[\x00-\x1f\x7f]", "", name).strip() + if len(name) > 150: + stem, dot, ext = name.rpartition(".") + if dot and 0 < len(ext) <= 12: + name = stem[: 150 - len(ext) - 1] + "." + ext + else: + name = name[:150] + if not name or name in {".", ".."}: + normalized = (mime_type or "").split(";", 1)[0].strip().lower() + ext = mimetypes.guess_extension(normalized) or ".bin" + name = f"resource{ext}" + return name + + +def _cache_mcp_audio_block(block) -> str: + """Cache an ``AudioContent`` block and return a ``MEDIA:`` tag; "" when not + audio or on any failure (same fail-open contract as the image path).""" + import base64 + + data = getattr(block, "data", None) + mime_type = str(mcp_field(block, "mime_type", "mimeType") or "").split(";", 1)[0].strip().lower() + if data is None or not mime_type.startswith("audio/"): + return "" + if len(data) > _core._MCP_RESOURCE_MAX_B64_CHARS: + return f"[MCP audio resource too large to cache: ~{len(data) * 3 // 4} bytes]" + try: + raw_bytes = base64.b64decode(data) + except (TypeError, ValueError) as exc: + logger.warning("MCP audio block decode failed (%s): %s", mime_type, exc) + return "" + if len(raw_bytes) > _core._MCP_RESOURCE_MAX_BYTES: + return f"[MCP audio resource too large to cache: {len(raw_bytes)} bytes]" + try: + from gateway.platforms.base import cache_audio_from_bytes + import mimetypes + + ext = ( + {"audio/wav": ".wav", "audio/x-wav": ".wav", "audio/wave": ".wav"}.get(mime_type) + or mimetypes.guess_extension(mime_type) + or ".ogg" + ) + audio_path = cache_audio_from_bytes(raw_bytes, ext=ext) + except ImportError: + logger.debug("MCP audio caching skipped — gateway.platforms.base unavailable") + return "" + except Exception as exc: + logger.warning("MCP audio block cache failed: %s", exc) + return "" + return f"MEDIA:{audio_path}" + + +def _render_mcp_resource_block(block, server_name: str = "") -> str: + """Render a ``ResourceLink`` or ``EmbeddedResource`` block as text. + + Embedded text → the text; embedded blob → decoded (size-capped) into the + document cache with a path marker; link → the URI plus a pointer at the + server's read_resource tool (no fetch here — links are only readable via + the originating session). "" for non-resource blocks; failures are + reported inline rather than silently dropped. + """ + block_type = getattr(block, "type", "") + + if block_type == "resource_link" or ( + hasattr(block, "uri") and not hasattr(block, "resource") and block_type != "text" + ): + uri = getattr(block, "uri", None) + if not uri: + return "" + name = getattr(block, "name", "") or "" + mime = mcp_field(block, "mime_type", "mimeType", "") or "" + details = f"uri={uri}" + if name: + details += f", name={name}" + if mime: + details += f", mimeType={mime}" + reader = ( + mcp_prefixed_tool_name(server_name, "read_resource") + if server_name + else "the MCP server's read_resource tool" + ) + return f"[MCP resource link: {details} — fetch it with {reader}]" + + resource = getattr(block, "resource", None) + if resource is None: + return "" + + text = getattr(resource, "text", None) + if text is not None: + return strip_unicode_tags(str(text)) + + blob = getattr(resource, "blob", None) + if blob is None: + return "" + + import base64 + + uri = str(getattr(resource, "uri", "") or "") + mime = str(mcp_field(resource, "mime_type", "mimeType", "") or "") + if len(blob) > _core._MCP_RESOURCE_MAX_B64_CHARS: + return f"[MCP embedded resource too large to cache: ~{len(blob) * 3 // 4} bytes, uri={uri}]" + try: + raw_bytes = base64.b64decode(blob) + except (TypeError, ValueError) as exc: + logger.warning("MCP embedded resource decode failed (%s): %s", mime or uri, exc) + return f"[MCP embedded resource could not be decoded: {mime or uri}]" + if len(raw_bytes) > _core._MCP_RESOURCE_MAX_BYTES: + return f"[MCP embedded resource too large to cache: {len(raw_bytes)} bytes, uri={uri}]" + try: + from gateway.platforms.base import cache_document_from_bytes + + path = cache_document_from_bytes(raw_bytes, _mcp_resource_filename(uri, mime)) + except ImportError: + logger.debug("MCP resource caching skipped — gateway.platforms.base unavailable") + return f"[MCP embedded resource received ({len(raw_bytes)} bytes, {mime or 'unknown type'}) but document cache unavailable in this process]" + except Exception as exc: + logger.warning("MCP embedded resource cache failed: %s", exc) + return f"[MCP embedded resource could not be cached: {mime or uri}]" + return f"[MCP resource saved to {path} ({mime or 'unknown type'}, {len(raw_bytes)} bytes) — read it with read_file or terminal tools]" diff --git a/tools/mcp_tool_errors.py b/tools/mcp_tool_errors.py new file mode 100644 index 0000000000..fe4d727a92 --- /dev/null +++ b/tools/mcp_tool_errors.py @@ -0,0 +1,541 @@ +"""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 logging +import asyncio +import errno +import os +import re +from typing import Any, List, Optional +from urllib.parse import urlparse +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 + + +def _handshake_rejected_as_modern(exc: BaseException) -> bool: + """True when a failed ``initialize`` signals a stateless-only (2026-07-28) server. + + Structural code check first, then substring fallback — never ``isinstance`` on + SDK exception types (they arrive wrapped in ExceptionGroups and drift across generations). + """ + err = getattr(exc, "error", None) + code = getattr(err, "code", None) or getattr(exc, "code", None) + if code in (_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, _core._JSONRPC_METHOD_NOT_FOUND): + return True + msg = str(exc).lower() + if not msg: + return False + return ( + "unsupported protocol version" in msg + or str(_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION) in msg + or _is_method_not_found_error(exc) + ) + + +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. Structural + ``MCPError.error.code`` check first, then substring fallback — including + "Unknown method: ", which some servers use; without it the + ping→list_tools keepalive fallback never latches and reconnect-loops. + """ + err = getattr(exc, "error", None) + code = getattr(err, "code", None) + if code == _core._JSONRPC_METHOD_NOT_FOUND: + return True + msg = str(exc).lower() + if not msg: + return False + return ( + str(_core._JSONRPC_METHOD_NOT_FOUND) in msg + or "method not found" in msg + or "unknown method" in msg + or "not found: ping" in msg + ) + + +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. + """ + + +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. + """ + + +def _unwrap_exception_group(exc: BaseException) -> BaseException: + """Extract the root-cause leaf from anyio ``(Base)ExceptionGroup`` wrappers. + + Group ``str()`` is opaque ("unhandled errors in a TaskGroup"), 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. + """ + while isinstance(exc, BaseExceptionGroup) and exc.exceptions: + fatal, _rest = exc.split((KeyboardInterrupt, SystemExit)) + if fatal is not None: + leaf: BaseException = fatal + while isinstance(leaf, BaseExceptionGroup) and leaf.exceptions: + leaf = leaf.exceptions[0] + raise leaf + chosen = exc.exceptions[0] + for sub in exc.exceptions: + if not _contains_only_cancellation(sub): + chosen = sub + break + exc = chosen + return exc + + +def _contains_only_cancellation(exc: BaseException) -> bool: + """True if ``exc`` is (or a group containing only) CancelledError.""" + if isinstance(exc, BaseExceptionGroup): + return all(_contains_only_cancellation(sub) for sub in exc.exceptions) + return isinstance(exc, asyncio.CancelledError) + + +def _classify_mcp_failure(exc: BaseException) -> str: + """Classify a connection failure as ``'permanent'`` or ``'transient'``. + + Permanent (deterministic — ``run()`` parks immediately instead of burning the + retry ladder): auth 401/403, NonMcpEndpointError, InvalidMcpUrlError, missing + stdio command (FileNotFoundError / ENOENT). Everything else keeps backoff retry. + """ + root = _unwrap_exception_group(exc) + if _core._is_auth_error(root): + return "permanent" + if isinstance(root, (NonMcpEndpointError, InvalidMcpUrlError)): + return "permanent" + if isinstance(root, FileNotFoundError): + return "permanent" + if isinstance(root, OSError) and getattr(root, "errno", None) == errno.ENOENT: + return "permanent" + # 401/403 HTTPStatusError that _is_auth_error's type-gate missed + # (auth types not importable in this environment). + status = getattr(getattr(root, "response", None), "status_code", None) + if status in (401, 403): + return "permanent" + return "transient" + + +def _validate_remote_mcp_url(server_name: str, url: Any) -> str: + """Return 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. + """ + if not isinstance(url, str): + raise InvalidMcpUrlError( + f"Invalid MCP URL for '{server_name}': expected a string, got " + f"{type(url).__name__}" + ) + stripped = url.strip() + if not stripped: + raise InvalidMcpUrlError( + f"Invalid MCP URL for '{server_name}': empty url" + ) + try: + parsed = urlparse(stripped) + except Exception as exc: # urlparse is very permissive — belt and braces + raise InvalidMcpUrlError( + f"Invalid MCP URL for '{server_name}': {stripped!r} ({exc})" + ) from exc + if parsed.scheme.lower() not in {"http", "https"}: + raise InvalidMcpUrlError( + f"Invalid MCP URL for '{server_name}': scheme must be http or " + f"https, got {parsed.scheme!r} ({stripped!r})" + ) + if not parsed.netloc: + raise InvalidMcpUrlError( + f"Invalid MCP URL for '{server_name}': missing host ({stripped!r})" + ) + # ``urlparse`` accepts ``http://:8080`` (empty host, explicit port) — reject it. + if not parsed.hostname: + raise InvalidMcpUrlError( + f"Invalid MCP URL for '{server_name}': missing hostname " + f"({stripped!r})" + ) + return stripped + + +def _resolve_client_cert(server_name: str, config: dict): + """Resolve ``client_cert`` / ``client_key`` into 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. + """ + raw_cert = config.get("client_cert") + raw_key = config.get("client_key") + + if raw_cert is None and raw_key is None: + return None + + def _expand(path: Any, label: str) -> str: + if not isinstance(path, str) or not path.strip(): + raise ValueError( + f"MCP server '{server_name}': {label} must be a non-empty " + f"string path (got {type(path).__name__})" + ) + expanded = os.path.expanduser(path.strip()) + if not os.path.isfile(expanded): + raise FileNotFoundError( + f"MCP server '{server_name}': {label} not found at " + f"{expanded!r}" + ) + return expanded + + if isinstance(raw_cert, (list, tuple)): + if raw_key is not None: + raise ValueError( + f"MCP server '{server_name}': specify either client_cert as " + f"a list [cert, key] OR client_cert + client_key, not both" + ) + if len(raw_cert) == 2: + return (_expand(raw_cert[0], "client_cert[0]"), _expand(raw_cert[1], "client_cert[1]")) + if len(raw_cert) == 3: + cert_path = _expand(raw_cert[0], "client_cert[0]") + key_path = _expand(raw_cert[1], "client_cert[1]") + password = raw_cert[2] + if not isinstance(password, str): + raise ValueError( + f"MCP server '{server_name}': client_cert[2] (key " + f"passphrase) must be a string" + ) + return (cert_path, key_path, password) + raise ValueError( + f"MCP server '{server_name}': client_cert list form must have 2 " + f"or 3 elements (got {len(raw_cert)})" + ) + + cert_path = _expand(raw_cert, "client_cert") + if raw_key is not None: + return (cert_path, _expand(raw_key, "client_key")) + return cert_path # single combined PEM (cert + key) + + +def _resolve_identity_header(server_name: str, config: dict): + """Resolve the optional per-server ``identity_header`` config. + + Shape: ``{name: "X-User-Id", value_from: "static"|"profile", value: "..."}`` + (``value`` required for static). Returns ``(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. + """ + raw = config.get("identity_header") + if raw is None: + return None + if not isinstance(raw, dict): + logger.warning( + "MCP server '%s': identity_header must be a mapping with " + "'name' and 'value'/'value_from' keys (got %s) — ignoring", + server_name, type(raw).__name__, + ) + return None + name = raw.get("name") + if not isinstance(name, str) or not name.strip(): + logger.warning( + "MCP server '%s': identity_header requires a non-empty " + "'name' — ignoring", server_name, + ) + return None + value_from = (raw.get("value_from") or "static").strip().lower() + if value_from == "static": + value = raw.get("value") + if not isinstance(value, str) or not value.strip(): + logger.warning( + "MCP server '%s': identity_header with value_from: static " + "requires a non-empty string 'value' — ignoring", + server_name, + ) + return None + return (name.strip(), value) + if value_from == "profile": + from hermes_cli.profiles import get_active_profile_name + return (name.strip(), get_active_profile_name()) + logger.warning( + "MCP server '%s': identity_header value_from must be 'static' or " + "'profile' (got %r) — ignoring", server_name, value_from, + ) + return None + + +def _apply_identity_header(server_name: str, config: dict, headers: dict) -> dict: + """Merge the resolved identity header into ``headers`` in place. + + An explicit per-server ``headers`` entry with the same name (any casing) + wins — the identity header never silently overrides user config. + """ + resolved = _resolve_identity_header(server_name, config) + if resolved is None: + return headers + name, value = resolved + if any(key.lower() == name.lower() for key in headers): + logger.debug( + "MCP server '%s': identity_header '%s' already set via explicit " + "headers config — keeping the explicit value", server_name, name, + ) + return headers + headers[name] = value + return headers + + +def _make_redirect_header_stripper( + original_url, + *, + strict: bool = False, + configured_header_names: "set[str] | frozenset[str]" = frozenset(), +): + """Build an httpx response hook that guards 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. + """ + + async def _strip_on_cross_origin_redirect(response): + if response.is_redirect and response.next_request: + target = response.next_request.url + if (target.scheme, target.host, target.port) != ( + original_url.scheme, original_url.host, original_url.port, + ): + response.next_request.headers.pop("authorization", None) + response.next_request.headers.pop("Authorization", None) + if strict: + for _name in configured_header_names: + while _name in response.next_request.headers: + del response.next_request.headers[_name] + + return _strip_on_cross_origin_redirect + + +def _format_connect_error(exc: BaseException) -> str: + """Render nested MCP connection errors into an actionable short message.""" + + def _find_missing(current: BaseException) -> Optional[str]: + nested = getattr(current, "exceptions", None) + if nested: + for child in nested: + missing = _find_missing(child) + if missing: + return missing + return None + if isinstance(current, FileNotFoundError): + if getattr(current, "filename", None): + return str(current.filename) + match = re.search(r"No such file or directory: '([^']+)'", str(current)) + if match: + return match.group(1) + for attr in ("__cause__", "__context__"): + nested_exc = getattr(current, attr, None) + if isinstance(nested_exc, BaseException): + missing = _find_missing(nested_exc) + if missing: + return missing + return None + + def _flatten_messages(current: BaseException) -> List[str]: + nested = getattr(current, "exceptions", None) + if nested: + flattened: List[str] = [] + for child in nested: + flattened.extend(_flatten_messages(child)) + return flattened + messages = [] + text = str(current).strip() + if text: + messages.append(text) + for attr in ("__cause__", "__context__"): + nested_exc = getattr(current, attr, None) + if isinstance(nested_exc, BaseException): + messages.extend(_flatten_messages(nested_exc)) + return messages or [current.__class__.__name__] + + missing = _find_missing(exc) + if missing: + message = f"missing executable '{missing}'" + if os.path.basename(missing) in {"npx", "npm", "node"}: + message += ( + " (ensure Node.js is installed and PATH includes its bin directory, " + "or set mcp_servers..command to an absolute path and include " + "that directory in mcp_servers..env.PATH)" + ) + return _sanitize_error(message) + + deduped: List[str] = [] + for item in _flatten_messages(exc): + if item not in deduped: + deduped.append(item) + return _sanitize_error("; ".join(deduped[:3])) + + +# Lazily-built caches so this module imports even without the SDK OAuth module. +_AUTH_ERROR_TYPES: tuple = () +_HTTP_STATUS_ERROR_TYPES: Optional[tuple] = None + + +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, so both go in the tuple. + """ + global _HTTP_STATUS_ERROR_TYPES + if _HTTP_STATUS_ERROR_TYPES is not None: + return _HTTP_STATUS_ERROR_TYPES + found: list = [] + sdk_mod = _core.sdk_httpx() + if sdk_mod is not None: + found.append(sdk_mod.HTTPStatusError) + try: + import httpx + if httpx.HTTPStatusError not in found: + found.append(httpx.HTTPStatusError) + except ImportError: + pass + _HTTP_STATUS_ERROR_TYPES = tuple(found) + return _HTTP_STATUS_ERROR_TYPES + + +def _get_auth_error_types() -> tuple: + """Cached tuple of 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`. + """ + global _AUTH_ERROR_TYPES + if _AUTH_ERROR_TYPES: + return _AUTH_ERROR_TYPES + types: list = [] + try: + from mcp.client.auth import OAuthFlowError, OAuthTokenError + types.extend([OAuthFlowError, OAuthTokenError]) + except ImportError: + pass + try: + from mcp.client.auth import UnauthorizedError # type: ignore # older SDKs + types.append(UnauthorizedError) + except ImportError: + pass + try: + from tools.mcp_oauth import OAuthNonInteractiveError + types.append(OAuthNonInteractiveError) + except ImportError: + pass + types.extend(_http_status_error_types()) + _AUTH_ERROR_TYPES = tuple(types) + return _AUTH_ERROR_TYPES + + +def _is_auth_error(exc: BaseException) -> bool: + """True if ``exc`` indicates an MCP OAuth failure. + + ``HTTPStatusError`` counts only with status 401; other HTTP errors fall + through to the generic error path. + """ + types = _get_auth_error_types() + if not types or not isinstance(exc, types): + return False + status_error_types = _http_status_error_types() + if status_error_types and isinstance(exc, status_error_types): + return getattr(exc.response, "status_code", None) == 401 + return True + + +# Lower-cased substrings meaning the server-side transport session expired / +# was GC'd. The OAuth token is still valid — only the transport needs rebuilding. +_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. +_EXC_TRAVERSAL_MAX_NODES = 10_000 + + +def _is_session_expired_error(exc: BaseException) -> bool: + """True if ``exc`` looks like an MCP transport session expiry. + + Streamable-HTTP servers GC session state (idle TTL, restart, pod rotation) + while the OAuth token stays valid, so unlike :func:`_is_auth_error` the fix + is a transport reconnect (``_reconnect_event``), not an OAuth refresh. + """ + # AnyIO stream exceptions are often message-less (``str(ClosedResourceError()) == ""``), + # so type checks are needed in addition to marker matching. + try: + from anyio import BrokenResourceError, ClosedResourceError, EndOfStream + + transport_error_types = ( + BrokenResourceError, + ClosedResourceError, + EndOfStream, + ) + except ImportError: # pragma: no cover - AnyIO is supplied by the MCP SDK + transport_error_types = () + + # 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 + budget = _EXC_TRAVERSAL_MAX_NODES + while stack and budget > 0: + current = stack.pop() + if current is None: + continue + identity = id(current) + if identity in seen: + continue + seen.add(identity) + budget -= 1 + + if isinstance(current, InterruptedError): + return False + if isinstance(current, transport_error_types): + transport_error_found = True + + # 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 msg and any(marker in msg for marker in _SESSION_EXPIRED_MARKERS): + transport_error_found = True + + stack.extend(getattr(current, "exceptions", ())) + stack.append(getattr(current, "__cause__", None)) + stack.append(getattr(current, "__context__", None)) + + return transport_error_found diff --git a/tools/mcp_tool_handlers.py b/tools/mcp_tool_handlers.py new file mode 100644 index 0000000000..0fc23c2a46 --- /dev/null +++ b/tools/mcp_tool_handlers.py @@ -0,0 +1,785 @@ +"""Registry-facing sync handlers for MCP tools and utility tools (resources/prompts), plus the per-call recovery ladder: trust gating, circuit breaker, auth (401) refresh, session-expired reconnect and dead-stdio respawn retry. Split from tools/mcp_tool.py.""" + +import logging +import asyncio +import contextvars +import inspect +import json +import time +from contextlib import asynccontextmanager +from types import SimpleNamespace +from typing import Any, Dict, List, Optional +from tools.registry import tool_error +from tools.ansi_strip import strip_unicode_tags +from tools.mcp_tool_common import _exc_str, _sanitize_error, mcp_field, _core +from tools.mcp_tool_content import _MCP_HARD_RESULT_CAP_CHARS, _cache_mcp_audio_block, _cache_mcp_image_block, _render_mcp_resource_block, _strip_reserved_meta_keys, _truncate_mcp_text_result +from tools.mcp_tool_errors import _is_session_expired_error + +logger = logging.getLogger("tools.mcp_tool") + + +def _trust_gate_check(server_name: str, tool_name: str) -> Optional[str]: + """Approval gate for write-capable tools on ``trust: untrusted`` servers. + + Returns None to proceed, or a ``tool_error`` string when blocked. + Fail-closed: approval-system errors block the call. + """ + trust = _core._server_trust_levels.get(server_name, _core._TRUST_FULL) + if trust != _core._TRUST_UNTRUSTED: + return None + if _core._tool_read_only_hints.get(server_name, {}).get(tool_name) is True: + return None + + # Lazy import: tools.approval routes the prompt to whichever surface owns + # the session (CLI, TUI, Telegram, ...) and normalizes the answer. + try: + from tools.approval import request_elicitation_consent + + answer = request_elicitation_consent( + ( + f"MCP tool '{tool_name}' on UNTRUSTED server " + f"'{server_name}' wants to run. This tool is write-capable " + f"(no readOnlyHint=true annotation) and may modify external " + f"state." + ), + ( + f"Server '{server_name}' is configured 'trust: untrusted'. " + f"Approve to run '{tool_name}' once, or deny to block it." + ), + surface=f"mcp-trust/{server_name}", + ) + except Exception as exc: + logger.error( + "MCP trust gate: approval check failed for %s.%s: %s", + server_name, tool_name, exc, exc_info=True, + ) + return tool_error( + f"MCP tool '{tool_name}' on untrusted server '{server_name}' " + f"was blocked: the approval system was unavailable " + f"(fail-closed)." + ) + + if answer == "accept": + return None + logger.info( + "MCP trust gate: user %s '%s' on untrusted server '%s'", + "cancelled" if answer == "cancel" else "denied", + tool_name, server_name, + ) + return tool_error( + f"The user did not approve running write-capable MCP tool " + f"'{tool_name}' on untrusted server '{server_name}'. The command " + f"was NOT run. Do not retry without explicit user direction." + ) + + +def _result_is_error(result) -> bool: + """True only for a JSON payload carrying an ``error`` key (non-JSON = success).""" + try: + return "error" in json.loads(result) + except (json.JSONDecodeError, TypeError): + return False + + +def _retry_once(server_name: str, retry_call, op_description: str, what: str): + """Re-run ``retry_call`` after a recovery step. + + Returns the result (and closes the circuit breaker) when it is not an error + payload; None when the retry raised or errored, so the caller falls through. + """ + try: + result = retry_call() + except Exception as retry_exc: + logger.warning( + "MCP %s/%s retry after %s failed: %s", + server_name, op_description, what, retry_exc, + ) + return None + if _result_is_error(result): + return None + _core._reset_server_error(server_name) + return result + + +def _handle_auth_error_and_retry( + server_name: str, + exc: BaseException, + retry_call, + op_description: str, +): + """Attempt OAuth recovery and one retry; return None to fall through. + + Non-auth exceptions return None. Otherwise: ask ``MCPOAuthManager.handle_401`` + whether recovery is viable; if so, signal ``_reconnect_event`` so the server + task rebuilds the session with fresh credentials, wait for ready, retry once. + Any failure returns the structured ``needs_reauth`` error so the model stops + trying to refresh manually. + """ + if not _core._is_auth_error(exc): + return None + + from tools.mcp_oauth_manager import get_manager + manager = get_manager() + + async def _recover(): + return await manager.handle_401(server_name, None) + + try: + recovered = _core._run_on_mcp_loop(_recover, timeout=10) + except Exception as rec_exc: + logger.warning( + "MCP OAuth '%s': recovery attempt failed: %s", + server_name, rec_exc, + ) + recovered = False + + if recovered: + with _core._lock: + srv = _core._servers.get(server_name) + reconnected = False + if srv is not None and hasattr(srv, "_reconnect_event"): + reconnected = _core._signal_reconnect_and_wait( + server_name, + srv, + op_description=f"{op_description} after OAuth recovery", + timeout=15, + ) + + # OAuth recovery + reconnect is independent evidence the server is viable, + # so close the breaker here, not only on retry success — otherwise a failing + # retry would leave it pinned open forever. A broken server re-trips it via + # _bump_server_error on the retry. + if reconnected: + _core._reset_server_error(server_name) + + result = _retry_once(server_name, retry_call, op_description, "auth recovery") + if result is not None: + return result + + # No recovery, or retry failed: structured needs_reauth error + breaker strike. + _core._bump_server_error(server_name) + return tool_error( + f"MCP server '{server_name}' requires re-authentication. " + f"Run `hermes mcp login {server_name}` (or delete the tokens " + f"file under ~/.hermes/mcp-tokens/ and restart). Do NOT retry " + f"this tool — ask the user to re-authenticate.", + needs_reauth=True, + server=server_name, + ) + + +def _handle_session_expired_and_retry( + server_name: str, + exc: BaseException, + retry_call, + op_description: str, +): + """Trigger a transport reconnect and retry once on session expiry. + + Unlike :func:`_handle_auth_error_and_retry` this skips ``handle_401`` — the + token is still valid, only the server-side session is stale. Returns None to + fall through (not session-expired, no server record / loop, reconnect did not + ready in time, or the retry also failed). + """ + if not _is_session_expired_error(exc): + return None + + with _core._lock: + srv = _core._servers.get(server_name) + if srv is None or not hasattr(srv, "_reconnect_event"): + return None + + loop = _core._mcp_loop + if loop is None or not loop.is_running(): + return None + + logger.info( + "MCP server '%s': %s failed with session-expired error (%s); " + "signalling transport reconnect and retrying once.", + server_name, op_description, exc, + ) + + if not _core._signal_reconnect_and_wait( + server_name, + srv, + op_description=op_description, + timeout=15, + ): + logger.warning( + "MCP server '%s': reconnect did not ready within 15s after " + "session-expired error; falling through to error response.", + server_name, + ) + return None + + return _retry_once(server_name, retry_call, op_description, "session reconnect") + + +class _StdioChildExited(RuntimeError): + """A server's stdio subprocess was gone when (or while) a call ran. + + Deliberately NOT a TimeoutError: nothing timed out — the child was already + dead (typically a gateway restart killed it under a live agent session). + Handled by :func:`_handle_stdio_child_exited_and_retry`. + """ + + +def _handle_stdio_child_exited_and_retry( + server_name: str, + exc: Exception, + retry_call, + op_description: str, +): + """Respawn a dead stdio child and retry the call once; None if not our error. + + Cannot hot-cycle respawns: this never spawns anything — it sets + ``_reconnect_event`` once and waits for the server task to publish a fresh + session, so spawn frequency stays governed by ``run()``'s rapid-drop budget. + Single-shot: a child that dies again immediately reports and stops. + """ + if not isinstance(exc, _StdioChildExited): + return None + + with _core._lock: + srv = _core._servers.get(server_name) + + reconnected = False + if srv is not None and hasattr(srv, "_reconnect_event"): + logger.info( + "MCP server '%s': %s found the stdio subprocess dead (%s); " + "respawning and retrying once.", + server_name, op_description, exc, + ) + loop = _core._mcp_loop + if loop is not None and loop.is_running(): + reconnected = _core._signal_reconnect_and_wait( + server_name, + srv, + op_description=op_description, + timeout=_core._STDIO_RESPAWN_WAIT_SEC, + ) + else: + # No MCP loop to wait on (non-async adapters, tests) — still request + # the respawn so the next call lands on a live transport. + _core._signal_reconnect(srv) + + if reconnected: + try: + result = retry_call() + except _StdioChildExited as retry_exc: + # Died again right after respawn: broken server, not a restart + # artifact. Stop here — run()'s budget takes it to the park. + logger.warning( + "MCP server '%s': %s stdio subprocess exited again right " + "after respawn (%s); not retrying further.", + server_name, op_description, retry_exc, + ) + _core._bump_server_error(server_name) + return tool_error( + f"MCP server '{server_name}' respawned its stdio subprocess " + f"and it exited again immediately. The server is not " + f"starting cleanly — do NOT retry this tool; ask the user to " + f"check the server's command and its stderr log." + ) + except Exception as retry_exc: + logger.warning( + "MCP %s/%s retry after stdio respawn failed: %s", + server_name, op_description, retry_exc, + ) + _core._bump_server_error(server_name) + return tool_error(_sanitize_error( + f"MCP call failed after respawning the stdio subprocess for " + f"'{server_name}': {type(retry_exc).__name__}: " + f"{_exc_str(retry_exc)}" + )) + if _result_is_error(result): + _core._bump_server_error(server_name) + else: + _core._reset_server_error(server_name) + return result + + _core._bump_server_error(server_name) + return tool_error( + f"MCP server '{server_name}' stdio subprocess had exited (this is " + f"not a timeout — the call never reached the server). A respawn was " + f"requested but no fresh session came back within " + f"{_core._STDIO_RESPAWN_WAIT_SEC:.0f}s. Wait a few seconds before retrying; " + f"if it keeps failing the server is not starting and needs the user." + ) + + +def _interrupted_call_result() -> str: + """Standardized JSON error for a user-interrupted MCP tool call.""" + return tool_error("MCP call interrupted: user sent a new message") + + +def _mark_server_call_started(server: Any) -> None: + """Record a user-visible MCP operation when the server supports it.""" + mark_tool_call = getattr(server, "mark_tool_call", None) + if callable(mark_tool_call): + mark_tool_call() + + +@asynccontextmanager +async def _track_inflight_rpc(server: Any, server_name: str, op: str): + """Register the running RPC on the server so teardown can fail it fast. + + If a deliberate reconnect/shutdown teardown cancels the task + (``_fail_inflight_calls`` sets ``_reconnecting`` first) the cancel becomes a + clean retryable RuntimeError; external cancels (caller timeout, user + interrupt) propagate unchanged. Test doubles without ``_inflight_tasks`` + simply skip tracking. + """ + inflight = getattr(server, "_inflight_tasks", None) + task = asyncio.current_task() + if task is not None and inflight is not None: + inflight.add(task) + try: + yield + except asyncio.CancelledError: + if getattr(server, "_reconnecting", False): + raise RuntimeError( + f"MCP {op} on '{server_name}' was aborted by a reconnect " + f"teardown; retry the request on the rebuilt session" + ) from None + raise + finally: + if task is not None and inflight is not None: + inflight.discard(task) + + +def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float): + """Return a sync registry handler (``handler(args_dict, **kwargs) -> str``) + that calls an MCP tool via the background loop.""" + + def _handler(args: dict, **kwargs) -> str: + # Security boundary: untrusted-server write tools need approval before + # ANY transport work, including the lazy first-use spawn below. + gate_error = _trust_gate_check(server_name, tool_name) + if gate_error is not None: + return gate_error + + # Circuit breaker. After the cooldown the breaker is half-open: the next + # call goes through as a probe; success resets it, failure re-bumps (which + # re-stamps the open-time and re-arms the cooldown). + if _core._server_error_counts.get(server_name, 0) >= _core._CIRCUIT_BREAKER_THRESHOLD: + opened_at = _core._server_breaker_opened_at.get(server_name, 0.0) + age = time.monotonic() - opened_at + if age < _core._CIRCUIT_BREAKER_COOLDOWN_SEC: + remaining = max(1, int(_core._CIRCUIT_BREAKER_COOLDOWN_SEC - age)) + return tool_error( + f"MCP server '{server_name}' is unreachable after " + f"{_core._server_error_counts[server_name]} consecutive " + f"failures. Auto-retry available in ~{remaining}s. " + f"Do NOT retry this tool yet — use alternative " + f"approaches or ask the user to check the MCP server." + ) + + server = _core._get_connected_server_for_call(server_name) + if not server: + _core._bump_server_error(server_name) + return tool_error(f"MCP server '{server_name}' is not connected") + + # No session: a reconnect may be completing (fresh session swaps in + # asynchronously), so wait briefly before charging a breaker strike. + if not server.session and not _core._wait_for_server_session_ready( + server, timeout=min(5.0, float(tool_timeout or 5.0)), + ): + # Still down — reconnecting or parked (e.g. dead stdio child). + # Probing a dead transport would re-arm the breaker forever, so + # ask the server task to rebuild (respawns stdio) and return a + # clean "reconnecting" error; the breaker resets once the fresh + # session initializes. + _core._bump_server_error(server_name) + if _core._signal_reconnect(server): + return tool_error( + f"MCP server '{server_name}' transport is down; " + f"reconnect requested. Do NOT retry this tool " + f"immediately — give it a few seconds to come back." + ) + return tool_error(f"MCP server '{server_name}' is not connected") + + async def _call(): + _mark_server_call_started(server) + async with server._rpc_lock, _track_inflight_rpc( + server, server_name, f"tools/call {tool_name}" + ): + # Snapshot contextvars so an elicitation callback (fired on the + # MCP recv loop, which doesn't inherit them) can replay them + # for gateway platform / session routing. + server._pending_call_context = contextvars.copy_context() + try: + # Fast-fail: an already-dead stdio child must not hold this + # slot for the full tool timeout. callable() + real-bool check + # because MagicMock attributes return truthy Mocks. + _stdio_dead = getattr(server, "_stdio_children_dead", None) + if ( + callable(_stdio_dead) + and isinstance(_stdio_dead_result := _stdio_dead(), bool) + and _stdio_dead_result + ): + # server.session is stale so the transport-down path above + # never fired; hand this to the respawn-and-retry path. + raise _StdioChildExited( + f"MCP stdio subprocess for '{server_name}' had " + f"already exited when the call was dispatched" + ) + _call_coro = server.session.call_tool(tool_name, arguments=args) + _watch_children = getattr(server, "_watch_stdio_children", None) + _watch_ok = ( + _watch_children is not None + and inspect.iscoroutinefunction(_watch_children) + and asyncio.iscoroutine(_call_coro) + ) + if not _watch_ok: + # Stubbed sessions return a non-awaitable, or there is no + # child-watcher to race: plain await. + result = ( + await _call_coro + if asyncio.iscoroutine(_call_coro) + else _call_coro + ) + else: + # Race the RPC against the stdio-children watcher so a + # mid-call death fails immediately. + rpc_task = asyncio.ensure_future(_call_coro) + watch_task = asyncio.ensure_future(_watch_children()) + try: + done, _pending = await asyncio.wait( + {rpc_task, watch_task}, + return_when=asyncio.FIRST_COMPLETED, + ) + if watch_task in done and not rpc_task.done(): + rpc_task.cancel() + # Nothing clears server.session on a mid-call + # death; the respawn-and-retry path owns the + # reconnect signal. + raise _StdioChildExited( + f"MCP stdio subprocess for " + f"'{server_name}' exited mid-call" + ) + result = await rpc_task + finally: + watch_task.cancel() + if not rpc_task.done(): + rpc_task.cancel() + await asyncio.gather( + rpc_task, watch_task, return_exceptions=True + ) + finally: + server._pending_call_context = None + # Round-trip completed: transport is healthy even if the tool + # returned isError. Clear the rapid-drop budget. + _mark_proven = getattr(server, "_mark_session_proven", None) + if _mark_proven is not None: + _mark_proven() + # CallToolResult: .content (blocks) and .is_error (.isError before mcp 2.0). + if mcp_field(result, "is_error", "isError", False): + error_text = "" + for block in (result.content or []): + if getattr(block, "text", None): + error_text += block.text + continue + # EmbeddedResource error payloads carry text under .resource.text. + res_text = getattr(getattr(block, "resource", None), "text", None) + if res_text: + error_text += str(res_text) + return tool_error(_sanitize_error( + _truncate_mcp_text_result( + error_text or "MCP tool returned an error" + ) + )) + + # Text blocks pass through; image/audio blocks are cached via the + # gateway image-cache so they flow out as MEDIA: tags; resource blocks + # (PDFs, docs, ...) are materialized rather than silently dropped. + parts: List[str] = [] + for block in (result.content or []): + if hasattr(block, "text") and block.text: + parts.append(strip_unicode_tags(block.text)) + continue + image_tag = _cache_mcp_image_block(block) + if image_tag: + parts.append(image_tag) + continue + audio_tag = _cache_mcp_audio_block(block) + if audio_tag: + parts.append(audio_tag) + continue + resource_text = _render_mcp_resource_block(block, server_name) + if resource_text: + parts.append(resource_text) + continue + # Benign empty renders log at debug; warn only for unknown shapes. + block_type = getattr(block, "type", None) or type(block).__name__ + if block_type in {"text", "resource", "audio", "image"}: + logger.debug( + "MCP %s: content block type %r rendered empty", + server_name, block_type, + ) + else: + logger.warning( + "MCP %s: dropping unsupported content block type %r", + server_name, block_type, + ) + text_result = "\n".join(parts) if parts else "" + + # Hard-cap pathological payloads; ordinary large results pass to spillover. + text_result = _truncate_mcp_text_result(text_result) + + # content is the primary (model-oriented) payload; structuredContent + # supplements it. Server-level `_meta` is surfaced too, minus + # protocol-reserved keys (`modelcontextprotocol`/`mcp` label followed + # by another label); vendor-namespaced keys pass through. + structured = mcp_field(result, "structured_content", "structuredContent") + # Cap structuredContent too (multi-MB JSON flood); over the hard cap it + # degrades to the head+tail truncated string. + if structured is not None: + try: + _structured_json = json.dumps(structured, ensure_ascii=False, default=str) + except (TypeError, ValueError): + _structured_json = None + if _structured_json is not None and len(_structured_json) > _MCP_HARD_RESULT_CAP_CHARS: + structured = _truncate_mcp_text_result(_structured_json) + meta = _strip_reserved_meta_keys(mcp_field(result, "meta", "meta")) + if structured is not None or meta is not None: + payload: Dict[str, Any] = {} + if text_result: + payload["result"] = text_result + if structured is not None: + if text_result: + payload["structuredContent"] = structured + else: + payload["result"] = structured + if meta is not None: + payload["_meta"] = meta + if "result" not in payload: + payload["result"] = text_result + try: + return json.dumps(payload, ensure_ascii=False) + except (TypeError, ValueError): + # Non-serializable metadata: drop the extras, keep the call. + return json.dumps({"result": text_result}, ensure_ascii=False) + return json.dumps({"result": text_result}, ensure_ascii=False) + + def _call_once(): + return _core._run_on_mcp_loop(_call, timeout=tool_timeout) + + try: + result = _call_once() + # An error payload from the tool itself still counts as a strike. + if _result_is_error(result): + _core._bump_server_error(server_name) + else: + _core._reset_server_error(server_name) + return result + except InterruptedError: + return _interrupted_call_result() + except Exception as exc: + # Recovery ladder, in order: dead stdio child (respawn + retry), + # auth (OAuth recovery + retry), session expiry (reconnect + retry). + # Each returns None when the exception is not its kind. + for recover in ( + _handle_stdio_child_exited_and_retry, + _handle_auth_error_and_retry, + _handle_session_expired_and_retry, + ): + recovered = recover(server_name, exc, _call_once, f"tools/call {tool_name}") + if recovered is not None: + return recovered + + _core._bump_server_error(server_name) + logger.error( + "MCP tool %s/%s call failed: %s", + server_name, tool_name, exc, + ) + return tool_error(_sanitize_error( + f"MCP call failed: {type(exc).__name__}: {_exc_str(exc)}" + )) + + return _handler + + +def _make_utility_handler(server_name: str, tool_timeout: float, op: str, + log_label: str, build_call): + """Shared shape of the four utility handlers (resources/prompts). + + ``build_call(server, args)`` returns an error string (validation failed) or a + zero-arg coroutine function doing the RPC under ``_rpc_lock``. The wrapper owns + the connected check and the auth / session-expired recovery ladder. + """ + + def _handler(args: dict, **kwargs) -> str: + server = _core._get_connected_server_for_call(server_name) + if not server or not server.session: + return tool_error(f"MCP server '{server_name}' is not connected") + + call = build_call(server, args) + if isinstance(call, str): + return call + + def _call_once(): + return _core._run_on_mcp_loop(call, timeout=tool_timeout) + + try: + return _call_once() + except InterruptedError: + return _interrupted_call_result() + except Exception as exc: + for recover in (_handle_auth_error_and_retry, _handle_session_expired_and_retry): + recovered = recover(server_name, exc, _call_once, op) + if recovered is not None: + return recovered + logger.error("MCP %s/%s failed: %s", server_name, log_label, exc) + return tool_error(_sanitize_error( + f"MCP call failed: {type(exc).__name__}: {_exc_str(exc)}" + )) + + return _handler + + +def _make_list_resources_handler(server_name: str, tool_timeout: float): + """Return a sync handler that lists resources from an MCP server.""" + + def _build(server, args): + async def _call(): + _mark_server_call_started(server) + async with server._rpc_lock: + all_resources = await _core._paginate_full_list( + server.session.list_resources, "resources", server_name + ) + resources = [] + for r in all_resources: + entry = {} + if hasattr(r, "uri"): + entry["uri"] = str(r.uri) + if hasattr(r, "name"): + entry["name"] = r.name + if hasattr(r, "description") and r.description: + entry["description"] = r.description + # Key stays camelCase — this is the tool's own JSON output shape. + _mime = mcp_field(r, "mime_type", "mimeType") + if _mime: + entry["mimeType"] = _mime + resources.append(entry) + return json.dumps({"resources": resources}, ensure_ascii=False) + return _call + + return _make_utility_handler(server_name, tool_timeout, "resources/list", "list_resources", _build) + + +def _make_read_resource_handler(server_name: str, tool_timeout: float): + """Return a sync handler that reads a resource by URI from an MCP server.""" + + def _build(server, args): + uri = args.get("uri") + if not uri: + return tool_error("Missing required parameter 'uri'") + + async def _call(): + _mark_server_call_started(server) + async with server._rpc_lock: + result = await server.session.read_resource(uri) + parts: List[str] = [] + contents = result.contents if hasattr(result, "contents") else [] + for block in contents: + if getattr(block, "text", None) is not None: + parts.append(strip_unicode_tags(block.text)) + elif getattr(block, "blob", None) is not None: + # Materialize binary contents into the document cache + # (same contract as EmbeddedResource blocks in tool results). + rendered = _render_mcp_resource_block( + SimpleNamespace(type="resource", resource=block), + server_name, + ) + parts.append(rendered or f"[binary data, {len(block.blob)} bytes]") + return json.dumps({"result": "\n".join(parts) if parts else ""}, ensure_ascii=False) + return _call + + return _make_utility_handler(server_name, tool_timeout, "resources/read", "read_resource", _build) + + +def _make_list_prompts_handler(server_name: str, tool_timeout: float): + """Return a sync handler that lists prompts from an MCP server.""" + + def _build(server, args): + async def _call(): + _mark_server_call_started(server) + async with server._rpc_lock: + all_prompts = await _core._paginate_full_list( + server.session.list_prompts, "prompts", server_name + ) + prompts = [] + for p in all_prompts: + entry = {} + if hasattr(p, "name"): + entry["name"] = p.name + if hasattr(p, "description") and p.description: + entry["description"] = p.description + if hasattr(p, "arguments") and p.arguments: + entry["arguments"] = [ + { + "name": a.name, + **({"description": a.description} if hasattr(a, "description") and a.description else {}), + **({"required": a.required} if hasattr(a, "required") else {}), + } + for a in p.arguments + ] + prompts.append(entry) + return json.dumps({"prompts": prompts}, ensure_ascii=False) + return _call + + return _make_utility_handler(server_name, tool_timeout, "prompts/list", "list_prompts", _build) + + +def _make_get_prompt_handler(server_name: str, tool_timeout: float): + """Return a sync handler that gets a prompt by name from an MCP server.""" + + def _build(server, args): + name = args.get("name") + if not name: + return tool_error("Missing required parameter 'name'") + arguments = args.get("arguments", {}) + + async def _call(): + _mark_server_call_started(server) + async with server._rpc_lock: + result = await server.session.get_prompt(name, arguments=arguments) + messages = [] + for msg in (result.messages if hasattr(result, "messages") else []): + entry = {} + if hasattr(msg, "role"): + entry["role"] = msg.role + if hasattr(msg, "content"): + content = msg.content + if hasattr(content, "text"): + entry["content"] = strip_unicode_tags(content.text) + elif isinstance(content, str): + entry["content"] = strip_unicode_tags(content) + else: + entry["content"] = strip_unicode_tags(str(content)) + messages.append(entry) + resp = {"messages": messages} + if hasattr(result, "description") and result.description: + resp["description"] = result.description + return json.dumps(resp, ensure_ascii=False) + return _call + + return _make_utility_handler(server_name, tool_timeout, "prompts/get", "get_prompt", _build) + + +def _make_check_fn(server_name: str): + """Return a check function that verifies the MCP connection is alive.""" + + def _check() -> bool: + with _core._lock: + server = _core._servers.get(server_name) + if server is not None and ( + server.session is not None or server._is_recycled_stdio() + ): + return True + # Lazy (schema-cache registered) servers count as available: the + # first real call spawns/connects them. + return server_name in _core._lazy_server_configs + + return _check diff --git a/tools/mcp_tool_health.py b/tools/mcp_tool_health.py new file mode 100644 index 0000000000..eb9ebab18e --- /dev/null +++ b/tools/mcp_tool_health.py @@ -0,0 +1,392 @@ +"""Session health for MCPServerTask: dynamic tool refresh on list_changed notifications, server log forwarding, keepalive probes, suspect-mark / lazy-verify, in-flight call fail-fast, stdio child liveness and stdio idle/lifetime recycling. Split from tools/mcp_tool.py.""" + +import asyncio +import json +import logging +import time +from typing import Optional +from tools.mcp_tool_errors import _is_method_not_found_error, _unwrap_exception_group +from tools.mcp_tool_schema import mcp_prefixed_tool_name +from tools.mcp_tool_registration import _forget_mcp_tool_server +from tools.mcp_tool_common import _core + +logger = logging.getLogger("tools.mcp_tool") + + +class MCPServerHealthMixin: + """Methods of :class:`tools.mcp_tool.MCPServerTask` (mixed in; relies on its attributes).""" + + __slots__ = () + + def _is_http(self) -> bool: + """Check if this server uses HTTP transport.""" + return "url" in self._config + + def _is_recycled_stdio(self) -> bool: + """Return True when a stdio server was intentionally recycled.""" + return not self._is_http() and self._recycled_reason is not None + + def mark_tool_call(self) -> None: + """Record that a user-visible MCP operation is starting.""" + self._last_tool_call_at = time.monotonic() + + def _mark_lifecycle_started(self) -> None: + now = time.monotonic() + self._lifecycle_started_at = now + self._last_tool_call_at = now + self._recycled_reason = None + + def _stdio_recycle_reason(self, now: Optional[float] = None) -> Optional[str]: + """Return the stdio recycle reason if idle/age limits have elapsed.""" + if self._is_http() or self._rpc_lock.locked(): + return None + now = time.monotonic() if now is None else now + if ( + self._max_lifetime_seconds is not None + and now - self._lifecycle_started_at >= self._max_lifetime_seconds + ): + return "max_lifetime_seconds" + if ( + self._idle_timeout_seconds is not None + and now - self._last_tool_call_at >= self._idle_timeout_seconds + ): + return "idle_timeout_seconds" + return None + + def _next_stdio_recycle_deadline(self) -> Optional[float]: + """Return the next monotonic recycle deadline for stdio, if any.""" + if self._is_http() or self._rpc_lock.locked(): + return None + deadlines = [] + if self._max_lifetime_seconds is not None: + deadlines.append(self._lifecycle_started_at + self._max_lifetime_seconds) + if self._idle_timeout_seconds is not None: + deadlines.append(self._last_tool_call_at + self._idle_timeout_seconds) + return min(deadlines) if deadlines else None + + def _mark_stdio_recycled(self, reason: str) -> None: + """Mark a stdio session dormant before its transport finishes closing.""" + self._recycled_reason = reason + self.session = None + + async def _refresh_tools_task(self): + """Run a dynamic tool refresh and log failures from background tasks.""" + try: + await self._refresh_tools() + except Exception: + logger.exception("MCP server '%s': dynamic tool refresh failed", self.name) + + def _schedule_tools_refresh(self) -> asyncio.Task: + """Schedule a background tool refresh and keep it strongly referenced.""" + task = asyncio.create_task(self._refresh_tools_task()) + self._pending_refresh_tasks.add(task) + task.add_done_callback(self._pending_refresh_tasks.discard) + return task + + def _make_logging_callback(self): + """Build a ``logging_callback`` that forwards server ``notifications/message`` + into Hermes logging tagged with the server name (the SDK default drops them).""" + async def _on_log(params): + try: + level = _core._MCP_LOG_LEVEL_MAP.get( + str(getattr(params, "level", "info")).lower(), logging.INFO, + ) + data = getattr(params, "data", None) + if not isinstance(data, str): + try: + data = json.dumps(data, ensure_ascii=False, default=str) + except (TypeError, ValueError): + data = str(data) + # Cap payloads so a chatty server can't flood agent.log. + if len(data) > 2000: + data = data[:2000] + "... [truncated]" + logger_name = getattr(params, "logger", None) + origin = f"{self.name}/{logger_name}" if logger_name else self.name + logger.log(level, "MCP server log [%s]: %s", origin, data) + except Exception: + logger.debug( + "Failed to handle MCP log notification from '%s'", + self.name, exc_info=True, + ) + return _on_log + + def _make_message_handler(self): + """Build a ``message_handler`` for ``ClientSession``: only + ``ToolListChangedNotification`` triggers a refresh; prompt/resource changes are logged.""" + async def _handler(message): + try: + if isinstance(message, Exception): + logger.debug("MCP message handler (%s): exception: %s", self.name, message) + return + if _core._MCP_NOTIFICATION_TYPES and isinstance(message, _core.ServerNotification): + # mcp 2.0 made ServerNotification a plain union (payload IS the + # message) instead of a RootModel (payload under ``.root``). + # ``isinstance`` accepts both; only the unwrap differs — without + # it ``.root`` raises into the catch-all and refreshes stop. + match getattr(message, "root", message): + case _core.ToolListChangedNotification(): + logger.info( + "MCP server '%s': received tools/list_changed notification", + self.name, + ) + # Refresh in a separate task: some servers emit + # list_changed right after initialize while another + # request is in flight, and refreshing synchronously + # inside the handler can wedge the stdio JSON-RPC stream. + self._schedule_tools_refresh() + # Yield one tick so short-lived notification contexts + # (and tests) can observe the scheduled refresh. + await asyncio.sleep(0) + case _core.PromptListChangedNotification(): + logger.debug("MCP server '%s': prompts/list_changed (ignored)", self.name) + case _core.ResourceListChangedNotification(): + logger.debug("MCP server '%s': resources/list_changed (ignored)", self.name) + case _: + pass + except Exception: + logger.exception("Error in MCP message handler for '%s'", self.name) + return _handler + + async def _refresh_tools(self): + """Re-fetch tools on ``tools/list_changed`` and update the registry. + + The lock serializes rapid-fire notifications. After the list_tools + ``await``, all mutations are synchronous — atomic on the event loop. + """ + from tools.registry import registry + + if not self._advertises_tools(): + # Shouldn't happen, but tools/list would raise MCPError(-32601). + return + + async with self._refresh_lock: + old_tool_names = set(self._registered_tool_names) + + # 1. Fetch the current tool list (follow nextCursor). + async with self._rpc_lock: + new_mcp_tools = await _core._paginate_full_list( + self.session.list_tools, "tools", self.name + ) + + # 2. Remove only stale names first — no nuke-and-repave: live agent + # turns may hold tool-call IDs pointing at existing handlers, and + # in-place replacement avoids transient "tool not connected" races. + toolset_name = f"mcp-{self.name}" + stale_tool_names = old_tool_names - { + mcp_prefixed_tool_name(self.name, tool.name) + for tool in new_mcp_tools + } + for tool_name in stale_tool_names: + # Never remove a colliding name currently owned by another server. + if registry.get_toolset_for_tool(tool_name) != toolset_name: + continue + registry.deregister(tool_name, scope=_core._server_registry_scope(self.name)) + _forget_mcp_tool_server(tool_name) + + # 3. Re-register; the helper may skip names ambiguous after normalization. + self._tools = new_mcp_tools + registered_names = _core._register_server_tools( + self.name, self, self._config + ) + + # A raw name can become ambiguous without changing its normalized + # name, so the pre-pass misses it: drop any old entry the final + # collision-checked registration no longer owns. + registered_name_set = set(registered_names) + for tool_name in old_tool_names - registered_name_set: + if registry.get_toolset_for_tool(tool_name) != toolset_name: + continue + registry.deregister(tool_name, scope=_core._server_registry_scope(self.name)) + _forget_mcp_tool_server(tool_name) + self._registered_tool_names = registered_names + + # 4. Log what changed (user-visible). + new_tool_names = set(self._registered_tool_names) + added = new_tool_names - old_tool_names + removed = old_tool_names - new_tool_names + changes = [] + if added: + changes.append(f"added: {', '.join(sorted(added))}") + if removed: + changes.append(f"removed: {', '.join(sorted(removed))}") + if changes: + logger.warning( + "MCP server '%s': tools changed dynamically — %s. " + "Verify these changes are expected.", + self.name, "; ".join(changes), + ) + else: + logger.info( + "MCP server '%s': dynamically refreshed %d tool(s) (no changes)", + self.name, len(self._registered_tool_names), + ) + + async def _keepalive_probe(self) -> None: + """Exercise the session; raise on a genuine connection failure. + + ``ping`` first (cheap, OPTIONAL utility). On -32601 latch + ``_ping_unsupported`` and fall back to ``list_tools`` when the server + advertises tools; otherwise the -32601 propagates (no liveness primitive + left). The latch resets on each fresh transport connection. + """ + if not self._ping_unsupported: + try: + await asyncio.wait_for(self.session.send_ping(), timeout=30.0) + return + except Exception as exc: + if _is_method_not_found_error(exc): + # Ping is definitively unsupported. + if not self._advertises_tools(): + raise + self._ping_unsupported = True + logger.info( + "MCP server '%s': does not implement the optional " + "'ping' utility (-32601); using 'list_tools' for " + "keepalive on this connection.", + self.name, + ) + elif isinstance(exc, (TimeoutError, asyncio.TimeoutError)) and self._advertises_tools(): + # A server that silently drops ping looks like a dead transport. + # Confirm with list_tools before declaring it dead; if that + # also fails, propagate the original failure. + try: + await asyncio.wait_for(self.session.list_tools(), timeout=30.0) + except Exception: + raise exc from None + # Transport alive; latch so later keepalives skip the 30s wait. + self._ping_unsupported = True + logger.info( + "MCP server '%s': ping timed out but list_tools " + "succeeded — server silently drops ping; using " + "'list_tools' for keepalive on this connection.", + self.name, + ) + return + else: + # Closed transport, expired session, etc. — real failure. + raise + + # Fallback probe for servers without ping support. + await asyncio.wait_for(self.session.list_tools(), timeout=30.0) + + def _mark_session_proven(self) -> None: + """Record that the session demonstrated real health (keepalive or tool-call success). + + Only then is the reconnect budget cleared: a handshake that drops moments + later must keep consuming ``_reconnect_retries`` so a flapping transport + still reaches the park instead of respawning forever. + """ + if not self._session_proven: + self._session_proven = True + self._reconnect_retries = 0 + if self._was_parked: + self._was_parked = False + logger.warning( + "MCP server '%s': revived — session healthy again after " + "parking (state: parked → connected)", + self.name, + ) + # A proven fresh transport clears the one-time permanent-failure + # grace and any race bookkeeping. + self._permanent_grace_used = False + self._teardown_race = False + + def mark_suspect(self, reason: str) -> None: + """Latch a suspicion (no I/O). The NEXT call verifies via + :meth:`ensure_healthy` and recycles the transport if the probe fails.""" + if self._suspect_reason is None and reason: + logger.warning( + "MCP server '%s': connection marked suspect (%s); next call " + "will health-check it", + self.name, reason, + ) + self._suspect_reason = reason or None + + async def ensure_healthy(self, timeout: float = 5.0) -> bool: + """Verify a suspect connection before reuse; recycle if dead. + + True when healthy (suspicion cleared). On failure requests a reconnect, + drops the stale session so the caller's no-session path takes over, and + returns False. Never raises. + """ + reason = self._suspect_reason + if not reason: + return True + if self.session is None: + # Nothing to verify — the reconnect path owns recovery now. + self._suspect_reason = None + self._reconnect_event.set() + return False + try: + await asyncio.wait_for(self._keepalive_probe(), timeout=timeout) + except Exception as exc: + root = _unwrap_exception_group(exc) + logger.warning( + "MCP server '%s': suspect connection (%s) failed health " + "check (%s: %s) — requesting reconnect (state: suspect → " + "degraded)", + self.name, reason, type(root).__name__, root, + ) + self._suspect_reason = None + self.mark_suspect(f"health check failed after {reason}") + self.session = None + self._ready.clear() + self._reconnect_event.set() + return False + logger.info( + "MCP server '%s': suspect connection passed health check " + "(%s) — clearing suspicion", + self.name, reason, + ) + self._suspect_reason = None + self._mark_session_proven() + return True + + def _fail_inflight_calls(self, reason: str) -> None: + """Cancel every in-flight RPC on this connection. + + Called from lifecycle exits BEFORE the transport unwinds: the SDK does + not always fail pending requests when streams close, so a call would + otherwise wait out the full tool timeout. Cancelling anything flags + ``_teardown_race`` so run() treats the next reconnect as recovery + rather than charging the rapid-drop budget. + """ + victims = [t for t in self._inflight_tasks if not t.done()] + if not victims: + return + self._reconnecting = True + self._teardown_race = True + self.mark_suspect(f"{reason} tore down {len(victims)} in-flight call(s)") + for task in victims: + task.cancel() + + def _stdio_children_dead(self) -> bool: + """True when every stdio child we spawned has exited. + + Best-effort: False (unknown → don't fail fast) for HTTP, no captured + PIDs, missing psutil, or a failed probe. + """ + pids = getattr(self, "_stdio_child_pids", None) + if not pids or self._is_http(): + return False + try: + import psutil + except ImportError: + return False + for pid in pids: + # pid_exists handles Windows without signal-permission noise. + try: + alive = psutil.pid_exists(pid) + except Exception: + return False + if alive: + return False + return True + + async def _watch_stdio_children(self) -> None: + """Poll child liveness while a stdio RPC is in flight; resolves when a + tracked child dies so the caller cancels the RPC instead of waiting out the timeout.""" + while True: + if self._stdio_children_dead(): + return + await asyncio.sleep(0.25) diff --git a/tools/mcp_tool_lifecycle.py b/tools/mcp_tool_lifecycle.py new file mode 100644 index 0000000000..6be36a732e --- /dev/null +++ b/tools/mcp_tool_lifecycle.py @@ -0,0 +1,342 @@ +"""MCP process lifecycle: stdio child PID tracking and orphan cleanup, graceful +server shutdown and draining of the background MCP loop.""" + +import logging +import asyncio +import os +import time +from typing import Dict, Optional +from tools.mcp_tool_common import _core + +logger = logging.getLogger("tools.mcp_tool") + + +# Live stdio MCP children (pid -> server_name), added after connection and +# removed on normal shutdown, so they can be force-killed if SDK teardown fails. +_stdio_pids: Dict[int, str] = {} + + +# PIDs that survived their session context exit (SDK teardown failed to kill +# them); detected in _run_stdio's finally, reaped by _kill_orphaned_mcp_children(). +# Kept separate from _stdio_pids so cleanup sweeps never race active sessions. +_orphan_stdio_pids: set = set() + + +_orphan_stdio_pid_servers: Dict[int, str] = {} + + +# pid -> pgid captured at spawn. The SDK spawns children with +# start_new_session=True (PGID == PID); grandchildren inherit that PGID and +# keep it after the direct child exits, so killpg still reaches them. Tracked +# separately from _stdio_pids so the PGID survives the child's removal. +# Empty on Windows (os.getpgid is POSIX-only). +_stdio_pgids: Dict[int, int] = {} + + +def _snapshot_child_pids() -> set: + """Current direct-child PIDs: /proc on Linux, else psutil, else empty set.""" + my_pid = os.getpid() + + # /proc//task//children is per-THREAD, and stdio_client() spawns + # from the MCP loop thread, so union every task's children — reading only + # the main thread's file returns an empty set on every Linux install. + try: + task_dir = f"/proc/{my_pid}/task" + tids = os.listdir(task_dir) + found: set = set() + for tid in tids: + try: + with open(f"{task_dir}/{tid}/children", encoding="utf-8") as f: + found.update(int(p) for p in f.read().split() if p.strip()) + except (FileNotFoundError, OSError, ValueError): + continue # thread exited between listdir and open + return found + except (FileNotFoundError, OSError, ValueError): + pass + + try: + import psutil + return {c.pid for c in psutil.Process(my_pid).children()} + except Exception: + pass + + return set() + + +# argv markers of non-MCP gateway children that can race into the snapshot +# delta during an MCP spawn (defense-in-depth; LSP/slash_worker already use +# start_new_session). Matched against argv[1:] because Python/Java children +# start with the interpreter path. +_NON_MCP_CHILD_CMDLINE_MARKERS: tuple[str, ...] = ( + "tui_gateway.slash_worker", + "tui_gateway.entry", + "-dorg.eclipse.equinox.launcher", # jdtls (legacy arg style) + "eclipse.jdt.ls", + "org.eclipse.equinox.launcher_", +) + + +def _filter_mcp_children(pids: set) -> set: + """Drop non-MCP children from a PID snapshot delta. + + Tracking a stray child in _stdio_pgids is catastrophic if it lacks + start_new_session: its pgid can be the TUI parent's, so the shutdown + killpg() would kill the TUI itself. + """ + if not pids: + return pids + try: + import psutil + except ImportError: + return pids # keep all PIDs (prior behavior) + filtered: set = set() + for pid in pids: + try: + argv = psutil.Process(pid).cmdline() + except (psutil.NoSuchProcess, psutil.AccessDenied, OSError): + # Raced away or zombie — cannot be our fresh server, unsafe to track. + continue + if any( + marker in arg + for arg in argv[1:] + for marker in _NON_MCP_CHILD_CMDLINE_MARKERS + ): + continue + filtered.add(pid) + return filtered + + +def shutdown_mcp_servers(*, scope: Optional[str] = None): + """Close MCP server connections (in parallel) and stop the background loop. + + Each server Task is signalled to exit its own ``async with`` so the anyio + cancel-scope cleanup runs in the Task that opened it. ``scope`` restricts + teardown to one multiplexed profile's servers (its ``/reload-mcp`` must not + kill other profiles') and leaves the shared loop running if anything else + is still connected. + """ + with _core._lock: + selected = [ + name for name in _core._servers + if scope is None or _core._server_scope_keys.get(name) == scope + ] + servers_snapshot = [_core._servers[name] for name in selected] + + # Fast path: nothing to shut down. Still clear the connect-cooldown maps — + # a server that failed to connect is never in ``_servers``, so this is the + # most likely state for stale backoff entries; a restart must retry at once. + if not servers_snapshot: + with _core._lock: + _core._server_connect_retry_after.clear() + _core._server_connect_failures.clear() + _core._stop_mcp_loop(only_if_idle=scope is not None) + return + + async def _shutdown(): + results = await asyncio.gather( + *(server.shutdown() for server in servers_snapshot), + return_exceptions=True, + ) + for server, result in zip(servers_snapshot, results): + if isinstance(result, Exception): + logger.debug( + "Error closing MCP server '%s': %s", server.name, result, + ) + with _core._lock: + for name in selected: + _core._servers.pop(name, None) + _core._server_scope_keys.pop(name, None) + # Drop connect-retry cooldowns too: a restart must re-attempt every + # server immediately, not honour a stale per-server backoff. + _core._server_connect_retry_after.clear() + _core._server_connect_failures.clear() + + with _core._lock: + loop = _core._mcp_loop + if loop is not None and loop.is_running(): + from agent.async_utils import safe_schedule_threadsafe + future = safe_schedule_threadsafe( + _shutdown(), loop, + logger=logger, + log_message="MCP shutdown: failed to schedule", + ) + if future is not None: + try: + future.result(timeout=15) + except BaseException as exc: + logger.debug("Error during MCP shutdown: %s", exc) + + # Unconditional final sweep: whether ``_shutdown`` ran, timed out, or was + # never scheduled, no stale connect-cooldown state may survive shutdown. + with _core._lock: + _core._server_connect_retry_after.clear() + _core._server_connect_failures.clear() + + _core._stop_mcp_loop(only_if_idle=scope is not None) + + +def _kill_orphaned_mcp_children( + include_active: bool = False, + server_name: Optional[str] = None, +) -> None: + """Best-effort reap of stdio MCP subprocesses: SIGTERM, wait 2s, SIGKILL survivors. + + By default only ``_orphan_stdio_pids`` (PIDs that outlived their session + context) are reaped so concurrent cron jobs / live sessions are untouched; + ``include_active=True`` also kills every ``_stdio_pids`` entry and is only + for final shutdown after the MCP loop has stopped. ``server_name`` limits + the sweep to one server (stdio reconnects cleaning up their old transport). + + On POSIX signals go via ``os.killpg`` to the spawn-time pgid when tracked, + so reparented grandchildren are reaped too; falls back to ``os.kill``. + """ + import signal as _signal + + with _core._lock: + pids: Dict[int, str] = {} + for opid in _orphan_stdio_pids: + owner = _orphan_stdio_pid_servers.get(opid, "orphan") + if server_name is not None and owner != server_name: + continue + pids[opid] = owner + for opid in pids: + _orphan_stdio_pids.discard(opid) + _orphan_stdio_pid_servers.pop(opid, None) + if include_active: + active = dict(_stdio_pids) + if server_name is not None: + active = { + pid: owner + for pid, owner in active.items() + if owner == server_name + } + pids.update(active) + for pid in active: + _stdio_pids.pop(pid, None) + # Snapshot pgids for the pids we're about to kill, then drop them so a + # future spawn can't collide with stale state. + pgids: Dict[int, int] = {pid: _stdio_pgids[pid] for pid in pids if pid in _stdio_pgids} + for pid in pgids: + _stdio_pgids.pop(pid, None) + + # Fast path: nothing to reap — skip the 2s sleep every MCP-free shutdown + # would otherwise pay. + if not pids: + return + + # Our own pgid, so _send_signal never killpg()s the gateway itself. + try: + _my_pgid = os.getpgrp() + except (AttributeError, OSError): + _my_pgid = None # Windows or restricted environment + + def _send_signal(pid: int, sig: int, server_name: str) -> None: + """SIGTERM/SIGKILL via pgroup on POSIX, fall back to pid signal.""" + pgid = pgids.get(pid) + killpg = getattr(os, "killpg", None) + if pgid is not None and killpg is not None: + if _my_pgid is not None and pgid == _my_pgid: + # Child shares the gateway's pgroup: killpg would kill the + # gateway too, so use per-pid kill. Warn because per-pid kill + # can't reach grandchildren in this group (inherent trade-off). + logger.warning( + "MCP server '%s' pgid %d matches gateway pgid; skipping " + "killpg to avoid self-kill and using per-pid kill — any " + "grandchildren in this group may not be reaped", + server_name, pgid, + ) + else: + try: + killpg(pgid, sig) + return + except (ProcessLookupError, PermissionError, OSError) as exc: + # Pgroup gone or refused — still try the direct child. + logger.debug( + "killpg(%d, %d) failed for MCP server '%s': %s; falling back to kill(pid)", + pgid, sig, server_name, exc, + ) + try: + os.kill(pid, sig) + except (ProcessLookupError, PermissionError, OSError): + pass + + for pid, server_name in pids.items(): + _send_signal(pid, _signal.SIGTERM, server_name) + logger.debug("Sent SIGTERM to orphaned MCP process %d (%s)", pid, server_name) + + time.sleep(2) + + _sigkill = getattr(_signal, "SIGKILL", _signal.SIGTERM) + # ``os.kill(pid, 0)`` is NOT a no-op on Windows; use the portable check. + from gateway.status import _pid_exists + for pid, server_name in pids.items(): + if not _pid_exists(pid): + continue # exited after SIGTERM + _send_signal(pid, _sigkill, server_name) + logger.warning( + "Force-killed MCP process %d (%s) after SIGTERM timeout", + pid, server_name, + ) + + +def _stop_mcp_loop_if_idle() -> bool: + """Stop the MCP loop only when no registered server still owns it. + + Probe paths create temporary MCPServerTasks not placed in ``_servers``; + they may clean up an idle loop but must not tear down the process-global + loop under live agent tools, or later calls fail with + ``MCP event loop is not running``. + """ + return _core._stop_mcp_loop(only_if_idle=True) + + +async def _drain_mcp_loop_tasks( + *, + timeout: Optional[float] = None, +) -> None: + """Cancel every task still pending on the MCP loop and reap it. + + ``Task.cancel()`` only schedules the throw, so tasks need a cancellation + cycle before the loop goes away; wait for them here, on their owning loop, + bounded so a task that suppresses cancellation cannot hang process exit. + """ + if timeout is None: + timeout = _core._MCP_LOOP_DRAIN_TIMEOUT + current = asyncio.current_task() + pending = [t for t in asyncio.all_tasks() if t is not current and not t.done()] + if not pending: + return + logger.debug("Draining %d pending task(s) from the MCP loop", len(pending)) + for task in pending: + task.cancel() + + done, still_pending = await asyncio.wait(pending, timeout=timeout) + for task in done: + if task.cancelled(): + continue + try: + task.exception() + except asyncio.CancelledError: + pass + except Exception as exc: + logger.debug("Pending MCP loop task ended during shutdown: %s", exc) + + if still_pending: + logger.warning( + "%d MCP loop task(s) still pending after %.1fs drain", + len(still_pending), timeout, + ) + + +async def _drain_and_stop_mcp_loop() -> None: + """Drain pending tasks, then stop the loop from its owning thread. + + Both must run as one loop-owned sequence: a ``loop.stop`` queued separately + by a timed-out caller can overtake the scheduled drain, leaving the drain + coroutine itself pending when the loop is closed. + """ + loop = asyncio.get_running_loop() + try: + await _drain_mcp_loop_tasks(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT) + finally: + loop.call_soon(loop.stop) diff --git a/tools/mcp_tool_registration.py b/tools/mcp_tool_registration.py new file mode 100644 index 0000000000..542ea40d79 --- /dev/null +++ b/tools/mcp_tool_registration.py @@ -0,0 +1,537 @@ +"""Registering a connected (or schema-cached) MCP server's tools into the tool +registry: include/exclude filtering, trust-tier metadata capture, utility-tool +selection, name-collision resolution and the schema-cache write-through.""" + +import logging +from types import SimpleNamespace +from typing import TYPE_CHECKING, Any, Callable, Dict, List +from tools.mcp_tool_common import _parse_boolish, _core +from tools.mcp_tool_handlers import _make_check_fn, _make_get_prompt_handler, _make_list_prompts_handler, _make_list_resources_handler, _make_read_resource_handler +from tools.mcp_tool_common import _resolve_tool_timeout +from tools.mcp_tool_schema import _UTILITY_CAPABILITY_ATTRS, _UTILITY_CAPABILITY_METHODS, _build_utility_schemas, _normalize_name_filter, matches_name_filter + +if TYPE_CHECKING: # pragma: no cover + from tools.mcp_tool import MCPServerTask + +logger = logging.getLogger("tools.mcp_tool") + + +def _normalize_server_trust(value: Any) -> str: + """Normalize a config ``trust`` value: None -> ``full`` (backward-compatible + default); any unrecognized string -> ``untrusted``, so a misspelled tier + fails closed rather than silently disabling gating.""" + if value is None: + return _core._TRUST_FULL + text = str(value).strip().lower() + if text == _core._TRUST_FULL: + return _core._TRUST_FULL + if text == _core._TRUST_UNTRUSTED: + return _core._TRUST_UNTRUSTED + logger.warning( + "MCP trust: unrecognized trust value %r — treating as 'untrusted' " + "(valid values: full, untrusted)", value, + ) + return _core._TRUST_UNTRUSTED + + +def _annotation_read_only_hint(mcp_tool: Any) -> bool: + """True only when annotations (SDK object or schema-cache dict) carry + ``readOnlyHint is True``; anything else is False — unknown metadata means + the tool must be treated as write-capable.""" + annotations = getattr(mcp_tool, "annotations", None) + if annotations is None: + return False + if isinstance(annotations, dict): + hint = annotations.get("readOnlyHint") + else: + hint = getattr(annotations, "readOnlyHint", None) + return hint is True + + +def _record_tool_trust_metadata( + server_name: str, config: dict, tools: List[Any] +) -> None: + """Capture per-server trust and per-tool readOnlyHint at discovery.""" + with _core._lock: + _core._server_trust_levels[server_name] = _normalize_server_trust( + (config or {}).get("trust") + ) + hints = _core._tool_read_only_hints.setdefault(server_name, {}) + for tool in tools: + name = getattr(tool, "name", None) + if name: + hints[name] = _annotation_read_only_hint(tool) + + +def _track_mcp_tool_server(tool_name: str, server_name: str) -> None: + """Remember the exact raw MCP server that registered *tool_name*.""" + with _core._lock: + _core._mcp_tool_server_names[tool_name] = server_name + + +def _forget_mcp_tool_server(tool_name: str) -> None: + """Forget MCP server provenance for a deregistered tool.""" + with _core._lock: + _core._mcp_tool_server_names.pop(tool_name, None) + + +def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: dict) -> List[dict]: + """Select utility schemas based on config and server capabilities.""" + tools_filter = config.get("tools") or {} + resources_enabled = _parse_boolish(tools_filter.get("resources"), default=True) + prompts_enabled = _parse_boolish(tools_filter.get("prompts"), default=True) + + # ``initialize_result.capabilities`` is the source of truth: its sub-objects + # are non-None iff the server advertises that request family. The old + # ``hasattr(server.session, ...)`` gate never filtered anything because + # ClientSession defines all four methods on the class. + advertised_caps = None + init_result = getattr(server, "initialize_result", None) + if init_result is not None: + advertised_caps = getattr(init_result, "capabilities", None) + + selected: List[dict] = [] + for entry in _build_utility_schemas(server_name): + handler_key = entry["handler_key"] + if handler_key in {"list_resources", "read_resource"} and not resources_enabled: + logger.debug("MCP server '%s': skipping utility '%s' (resources disabled)", server_name, handler_key) + continue + if handler_key in {"list_prompts", "get_prompt"} and not prompts_enabled: + logger.debug("MCP server '%s': skipping utility '%s' (prompts disabled)", server_name, handler_key) + continue + + if advertised_caps is not None: + cap_attr = _UTILITY_CAPABILITY_ATTRS[handler_key] + if getattr(advertised_caps, cap_attr, None) is None: + logger.debug( + "MCP server '%s': skipping utility '%s' " + "(server does not advertise '%s' capability)", + server_name, + handler_key, + cap_attr, + ) + continue + else: + # Legacy fallback when initialize_result wasn't captured (test + # fixtures, older paths): register every stub, as before. + required_method = _UTILITY_CAPABILITY_METHODS[handler_key] + if not hasattr(server.session, required_method): + logger.debug( + "MCP server '%s': skipping utility '%s' (session lacks %s)", + server_name, + handler_key, + required_method, + ) + continue + selected.append(entry) + return selected + + +def _existing_tool_names() -> List[str]: + """Return tool names for all currently connected servers.""" + names: List[str] = [] + for _sname, server in _core._servers.items(): + if hasattr(server, "_registered_tool_names"): + names.extend(server._registered_tool_names) + continue + for mcp_tool in server._tools: + schema = _core._convert_mcp_schema(server.name, mcp_tool) + names.append(schema["name"]) + # Lazy servers registered from the schema cache have no MCPServerTask yet — + # their tools live only in the registry. + with _core._lock: + lazy_names = [ + n + for sname, tool_names in _core._lazy_server_tool_names.items() + if sname not in _core._servers + for n in tool_names + ] + names.extend(lazy_names) + return names + + +# Utility tool key -> handler factory; each takes (server_name, tool_timeout). +_UTILITY_HANDLER_FACTORIES = { + "list_resources": _make_list_resources_handler, + "read_resource": _make_read_resource_handler, + "list_prompts": _make_list_prompts_handler, + "get_prompt": _make_get_prompt_handler, +} + + +def _make_tool_filter(name: str, config: dict) -> Callable[[str], bool]: + """Build the include/exclude predicate for a server's tool names. + + Rules: ``tools.include`` is a whitelist, ``tools.exclude`` a blacklist; + entries may be exact names or fnmatch globs; include wins over exclude; + ``include: []`` is an explicit empty whitelist (register nothing); neither + set registers everything. + """ + tools_filter = config.get("tools") or {} + include_raw = tools_filter.get("include") + include_set = _normalize_name_filter(include_raw, f"mcp_servers.{name}.tools.include") + include_active = isinstance(include_raw, (str, list, tuple, set)) + exclude_set = _normalize_name_filter( + tools_filter.get("exclude"), f"mcp_servers.{name}.tools.exclude" + ) + + def _should_register(tool_name: str) -> bool: + if include_active: + return matches_name_filter(tool_name, include_set) + if exclude_set: + return not matches_name_filter(tool_name, exclude_set) + return True + + return _should_register + + +def _resolve_name_collisions(name: str, candidates: List[dict]): + """Preflight registry-name collisions among one server's candidates. + + Returns ``(unique_candidates, ambiguous_names, shadowed_utilities)``. Exact + duplicate rows (same name + origin) are dropped silently; a generated + utility that normalizes onto a server-native tool's name is shadowed (the + native tool wins); any other multi-origin collision is ambiguous and every + colliding entry is skipped (fail closed). + """ + unique_candidates: List[dict] = [] + seen_candidates: set[tuple[str, str]] = set() + origins_by_name: Dict[str, set[str]] = {} + for candidate in candidates: + key = (candidate["registry_name"], candidate["origin"]) + if key in seen_candidates: + logger.debug( + "MCP server '%s': duplicate registration candidate %s for '%s'; " + "keeping one", + name, + candidate["origin"], + candidate["registry_name"], + ) + continue + seen_candidates.add(key) + unique_candidates.append(candidate) + origins_by_name.setdefault(candidate["registry_name"], set()).add( + candidate["origin"] + ) + + ambiguous_names: Dict[str, List[str]] = {} + shadowed_utilities: set[tuple[str, str]] = set() + for registry_name, origins in origins_by_name.items(): + if len(origins) <= 1: + continue + utility_origins = sorted( + o for o in origins if o.startswith("generated utility ") + ) + native_origins = sorted(origins - set(utility_origins)) + if len(native_origins) == 1 and utility_origins: + for util_origin in utility_origins: + shadowed_utilities.add((registry_name, util_origin)) + logger.info( + "MCP server '%s': generated utility %s normalizes onto " + "server-native %s — keeping the native tool and dropping the " + "utility (the utility only applies when the server has no such " + "tool of its own)", + name, + ", ".join(utility_origins), + native_origins[0], + ) + continue + ambiguous_names[registry_name] = sorted(origins) + + for registry_name, origins in sorted(ambiguous_names.items()): + logger.error( + "MCP server '%s': name normalization collision for '%s' from %s; " + "skipping every colliding entry instead of choosing an arbitrary " + "handler", + name, + registry_name, + ", ".join(origins), + ) + return unique_candidates, ambiguous_names, shadowed_utilities + + +def _write_schema_cache(name: str, server: "MCPServerTask", config: dict, should_register) -> None: + """Write-through: persist the manifest so the next startup can register this + server lazily without spawning it. Never raises.""" + try: + from tools.mcp_schema_cache import config_fingerprint, write_cache_entry + + tools_payload: List[dict] = [] + for mcp_tool in server._tools: + if not should_register(mcp_tool.name): + continue + schema_obj = getattr(mcp_tool, "inputSchema", None) + tools_payload.append({ + "name": mcp_tool.name, + "description": mcp_tool.description or "", + "inputSchema": schema_obj if isinstance(schema_obj, dict) else {}, + # Persisted so the lazy path trust-gates identically next startup. + "annotations": { + "readOnlyHint": _annotation_read_only_hint(mcp_tool), + }, + }) + utility_payload = [ + {"schema": entry["schema"], "handler_key": entry["handler_key"]} + for entry in _select_utility_schemas(name, server, config) + ] + cache_meta = getattr(server, "_list_cache_meta", None) or {} + write_cache_entry( + name, + config_fingerprint(config), + tools=tools_payload, + utility_tools=utility_payload, + ttl_ms=cache_meta.get("ttl_ms"), + cache_scope=cache_meta.get("cache_scope"), + ) + except Exception as exc: + logger.debug("MCP schema cache write failed for '%s': %s", name, exc) + + +def _register_server_tools(name: str, server: "MCPServerTask", config: dict) -> List[str]: + """Register an already-connected server's tools (plus utility tools) into + the registry; used by initial discovery and list_changed refresh. + + Toolset resolution for ``mcp-{server}`` / raw-name aliases derives from the + live registry rather than mutating ``toolsets.TOOLSETS``. Lossy name + normalization can map distinct raw names (``read-file``/``read_file``) to + one registry name; such collisions fail closed — every ambiguous entry is + skipped. Returns the registered prefixed names. + """ + from tools.registry import registry + + registered_names: List[str] = [] + toolset_name = f"mcp-{name}" + + _should_register = _make_tool_filter(name, config) + check_fn = _make_check_fn(name) + candidates: List[dict] = [] + + # Security boundary: capture trust tier and readOnlyHint NOW, at discovery, + # so the call-time gate classifies from data we control, not re-read + # server-supplied state. + _record_tool_trust_metadata(name, config, server._tools) + + for mcp_tool in server._tools: + if not _should_register(mcp_tool.name): + logger.debug( + "MCP server '%s': skipping tool '%s' (filtered by config)", + name, + mcp_tool.name, + ) + continue + + _core._scan_mcp_description(name, mcp_tool.name, mcp_tool.description or "") + schema = _core._convert_mcp_schema(name, mcp_tool) + candidates.append( + { + "registry_name": schema["name"], + "origin": f"tool {mcp_tool.name!r}", + "schema": schema, + "handler": _core._make_tool_handler( + name, mcp_tool.name, server.tool_timeout + ), + "check_fn": check_fn, + } + ) + + # Generated resource/prompt utility tools share the same namespace as raw + # MCP tools, so they must participate in the same collision preflight. + for entry in _select_utility_schemas(name, server, config): + schema = entry["schema"] + handler_key = entry["handler_key"] + candidates.append( + { + "registry_name": schema["name"], + "origin": f"generated utility {handler_key!r}", + "schema": schema, + "handler": _UTILITY_HANDLER_FACTORIES[handler_key]( + name, server.tool_timeout + ), + "check_fn": check_fn, + } + ) + + unique_candidates, ambiguous_names, shadowed_utilities = _resolve_name_collisions(name, candidates) + + for candidate in unique_candidates: + registry_name = candidate["registry_name"] + if registry_name in ambiguous_names: + continue + if (registry_name, candidate["origin"]) in shadowed_utilities: + continue + + existing_toolset = registry.get_toolset_for_tool(registry_name) + if existing_toolset and existing_toolset != toolset_name: + if existing_toolset.startswith("mcp-"): + logger.error( + "MCP server '%s': %s normalizes to '%s', already owned by " + "MCP toolset '%s' — skipping to preserve the existing owner", + name, + candidate["origin"], + registry_name, + existing_toolset, + ) + else: + logger.warning( + "MCP server '%s': %s (→ '%s') collides with built-in tool " + "in toolset '%s' — skipping to preserve built-in", + name, + candidate["origin"], + registry_name, + existing_toolset, + ) + continue + + registry.register( + name=registry_name, + toolset=toolset_name, + schema=candidate["schema"], + handler=candidate["handler"], + check_fn=candidate["check_fn"], + is_async=False, + description=candidate["schema"]["description"], + scope=_core._server_registry_scope(name), + ) + + # The pre-check above is advisory only. Multiple servers connect in + # parallel, so ToolRegistry.register() is the atomic ownership gate. + if registry.get_toolset_for_tool(registry_name) != toolset_name: + logger.error( + "MCP server '%s': registration of %s as '%s' was rejected by " + "the registry; skipping provenance/count updates", + name, + candidate["origin"], + registry_name, + ) + continue + + _core._track_mcp_tool_server(registry_name, name) + registered_names.append(registry_name) + + if registered_names: + registry.register_toolset_alias(name, toolset_name) + _write_schema_cache(name, server, config, _should_register) + + return registered_names + + +class _CachedMCPTool: + """Minimal stand-in for MCP Tool objects loaded from the schema cache.""" + + __slots__ = ("name", "description", "inputSchema") + + def __init__(self, name: str, description: str, inputSchema: dict): + self.name = name + self.description = description + self.inputSchema = inputSchema or {} + + +def _register_from_cache_sync(name: str, config: dict, entry: dict) -> List[str]: + """Lazy startup: register a server's tools from a cached manifest with no + child process. The first real call routes through + ``_get_connected_server_for_call`` -> ``_ensure_lazy_server_connected``.""" + from tools.registry import registry + from tools.mcp_schema_cache import ( + config_fingerprint, + tools_from_cache_entry, + utility_tools_from_cache_entry, + ) + + registered_names: List[str] = [] + toolset_name = f"mcp-{name}" + fingerprint = config_fingerprint(config) + tool_timeout = _resolve_tool_timeout(config) + _should_register = _make_tool_filter(name, config) + check_fn = _make_check_fn(name) + # Record trust metadata before registration so the call-time gate is + # identical whether the server was spawned live or registered from cache. + # Missing "annotations" in older cache files fails closed to write-capable. + cached_tool_objs = [ + SimpleNamespace( + name=raw.get("name"), + annotations=raw.get("annotations") + if isinstance(raw.get("annotations"), dict) else None, + ) + for raw in tools_from_cache_entry(entry) + if isinstance(raw, dict) and raw.get("name") + ] + _record_tool_trust_metadata(name, config, cached_tool_objs) + for raw in tools_from_cache_entry(entry): + if not isinstance(raw, dict): + continue + raw_name = raw.get("name") + if not raw_name or not _should_register(raw_name): + continue + raw_schema = raw.get("inputSchema") + mcp_tool = _CachedMCPTool( + raw_name, + raw.get("description") or "", + raw_schema if isinstance(raw_schema, dict) else {}, + ) + # Defense-in-depth: the cache file is user-writable JSON, so apply the + # same injection scan as eager discovery. + _core._scan_mcp_description(name, mcp_tool.name, mcp_tool.description or "") + schema = _core._convert_mcp_schema(name, mcp_tool) + registry_name = schema["name"] + existing_toolset = registry.get_toolset_for_tool(registry_name) + if existing_toolset and existing_toolset != toolset_name: + logger.warning( + "MCP server '%s' (lazy): cached tool '%s' collides with " + "toolset '%s' — skipping", + name, registry_name, existing_toolset, + ) + continue + registry.register( + name=registry_name, + toolset=toolset_name, + schema=schema, + handler=_core._make_tool_handler(name, raw_name, tool_timeout), + check_fn=check_fn, + is_async=False, + description=schema["description"], + scope=_core._mcp_registry_scope(), + ) + if registry.get_toolset_for_tool(registry_name) != toolset_name: + continue + _core._track_mcp_tool_server(registry_name, name) + registered_names.append(registry_name) + + for raw in utility_tools_from_cache_entry(entry): + if not isinstance(raw, dict): + continue + schema = raw.get("schema") + handler_key = raw.get("handler_key") + if not isinstance(schema, dict) or handler_key not in _UTILITY_HANDLER_FACTORIES: + continue + util_name = schema.get("name") or "" + if not util_name: + continue + existing_toolset = registry.get_toolset_for_tool(util_name) + if existing_toolset and existing_toolset != toolset_name: + continue + registry.register( + name=util_name, + toolset=toolset_name, + schema=schema, + handler=_UTILITY_HANDLER_FACTORIES[handler_key](name, tool_timeout), + check_fn=check_fn, + is_async=False, + description=schema.get("description") or "", + scope=_core._mcp_registry_scope(), + ) + if registry.get_toolset_for_tool(util_name) != toolset_name: + continue + _core._track_mcp_tool_server(util_name, name) + registered_names.append(util_name) + + if registered_names: + registry.register_toolset_alias(name, toolset_name) + with _core._lock: + _core._lazy_server_configs[name] = dict(config) + _core._lazy_server_fingerprints[name] = fingerprint + _core._lazy_server_tool_names[name] = list(registered_names) + logger.info( + "MCP server '%s' (lazy): registered %d tool(s) from schema cache", + name, len(registered_names), + ) + return registered_names diff --git a/tools/mcp_tool_sampling.py b/tools/mcp_tool_sampling.py new file mode 100644 index 0000000000..0ba22d55ca --- /dev/null +++ b/tools/mcp_tool_sampling.py @@ -0,0 +1,521 @@ +"""MCP client-side handlers for server-initiated requests: sampling (sampling/createMessage, text and tool-use results) and elicitation. Split from tools/mcp_tool.py.""" + +import asyncio +import json +import logging +import time +from typing import List, Optional +from tools.mcp_tool_common import _MISSING, _exc_str, _safe_numeric, _sanitize_error, mcp_field, _core +from tools.mcp_tool_schema import _normalize_mcp_input_schema + +logger = logging.getLogger("tools.mcp_tool") + + +class SamplingHandler: + """Handles sampling/createMessage requests for one MCP server. + + 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. + + Callable; 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``. + """ + + _STOP_REASON_MAP = {"stop": "endTurn", "length": "maxTokens", "tool_calls": "toolUse"} + + def __init__(self, server_name: str, config: dict): + self.server_name = server_name + self.max_rpm = _safe_numeric(config.get("max_rpm", 10), 10, int) + self.timeout = _safe_numeric(config.get("timeout", 30), 30, float) + self.max_tokens_cap = _safe_numeric(config.get("max_tokens_cap", 4096), 4096, int) + self.max_tool_rounds = _safe_numeric( + config.get("max_tool_rounds", 5), 5, int, minimum=0, + ) + self.model_override = config.get("model") + self.allowed_models = config.get("allowed_models", []) + + _log_levels = {"debug": logging.DEBUG, "info": logging.INFO, "warning": logging.WARNING} + self.audit_level = _log_levels.get( + str(config.get("log_level", "info")).lower(), logging.INFO, + ) + + self._rate_timestamps: List[float] = [] + self._tool_loop_count = 0 + self.metrics = {"requests": 0, "errors": 0, "tokens_used": 0, "tool_use_count": 0} + + def _check_rate_limit(self) -> bool: + """Sliding-window (60s) limiter; True if the request is allowed.""" + now = time.time() + window = now - 60 + self._rate_timestamps[:] = [t for t in self._rate_timestamps if t > window] + if len(self._rate_timestamps) >= self.max_rpm: + return False + self._rate_timestamps.append(now) + return True + + def _resolve_model(self, preferences) -> Optional[str]: + """Config override > server hint > None (use default).""" + if self.model_override: + return self.model_override + if preferences and hasattr(preferences, "hints") and preferences.hints: + for hint in preferences.hints: + if hasattr(hint, "name") and hint.name: + return hint.name + return None + + @staticmethod + def _extract_tool_result_text(block) -> str: + """Extract text from a ToolResultContent block.""" + if not hasattr(block, "content") or block.content is None: + return "" + items = block.content if isinstance(block.content, list) else [block.content] + return "\n".join(item.text for item in items if hasattr(item, "text")) + + def _convert_messages(self, params) -> List[dict]: + """Convert MCP SamplingMessages to OpenAI format. + + Uses ``msg.content_as_list`` when the SDK provides it; dispatches per + block by duck-typing. + """ + # A tool-use id is the discriminator for a tool *result* block; it must be + # read under both spellings (mcp_field) — on mcp 2.x a bare + # ``hasattr(b, "toolUseId")`` is False, silently dropping tool results. + def _tool_use_id(block): + return mcp_field(block, "tool_use_id", "toolUseId", _MISSING) + + def _is_tool_use(block): + return hasattr(block, "name") and hasattr(block, "input") + + messages: List[dict] = [] + for msg in params.messages: + 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] + 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) + ] + + for tr in tool_results: + messages.append({ + "role": "tool", + "tool_call_id": _tool_use_id(tr), + "content": self._extract_tool_result_text(tr), + }) + + if tool_uses: + tc_list = [] + for tu in tool_uses: + tc_list.append({ + "id": getattr(tu, "id", f"call_{len(tc_list)}"), + "type": "function", + "function": { + "name": tu.name, + "arguments": json.dumps(tu.input, ensure_ascii=False) if isinstance(tu.input, dict) else str(tu.input), + }, + }) + msg_dict: dict = {"role": msg.role, "tool_calls": tc_list} + text_parts = [b.text for b in content_blocks if hasattr(b, "text")] + if text_parts: + msg_dict["content"] = "\n".join(text_parts) + messages.append(msg_dict) + elif content_blocks: + # Pure text/image content. + if len(content_blocks) == 1 and hasattr(content_blocks[0], "text"): + messages.append({"role": msg.role, "content": content_blocks[0].text}) + else: + parts = [] + for block in content_blocks: + block_mime = mcp_field( + block, "mime_type", "mimeType", _MISSING + ) + if hasattr(block, "text"): + parts.append({"type": "text", "text": block.text}) + elif hasattr(block, "data") and block_mime is not _MISSING: + parts.append({ + "type": "image_url", + "image_url": {"url": f"data:{block_mime};base64,{block.data}"}, + }) + else: + logger.warning( + "Unsupported sampling content block type: %s (skipped)", + type(block).__name__, + ) + if parts: + messages.append({"role": msg.role, "content": parts}) + + return messages + + @staticmethod + def _error(message: str, code: int = -1): + """Return ErrorData (MCP spec) or raise as fallback.""" + if _core._MCP_SAMPLING_TYPES: + return _core.ErrorData(code=code, message=message) + raise Exception(message) + + def _build_tool_use_result(self, choice, response): + """Build a CreateMessageResultWithTools from an LLM tool_calls response.""" + self.metrics["tool_use_count"] += 1 + + # Tool-loop governance. + if self.max_tool_rounds == 0: + self._tool_loop_count = 0 + return self._error( + f"Tool loops disabled for server '{self.server_name}' (max_tool_rounds=0)" + ) + + self._tool_loop_count += 1 + if self._tool_loop_count > self.max_tool_rounds: + self._tool_loop_count = 0 + return self._error( + f"Tool loop limit exceeded for server '{self.server_name}' " + f"(max {self.max_tool_rounds} rounds)" + ) + + content_blocks = [] + for tc in choice.message.tool_calls: + args = tc.function.arguments + if isinstance(args, str): + try: + parsed = json.loads(args) + except (json.JSONDecodeError, ValueError): + logger.warning( + "MCP server '%s': malformed tool_calls arguments " + "from LLM (wrapping as raw): %.100s", + self.server_name, args, + ) + parsed = {"_raw": args} + else: + parsed = args if isinstance(args, dict) else {"_raw": str(args)} + + content_blocks.append(_core.ToolUseContent( + type="tool_use", + id=tc.id, + name=tc.function.name, + input=parsed, + )) + + logger.log( + self.audit_level, + "MCP server '%s' sampling response: model=%s, tokens=%s, tool_calls=%d", + self.server_name, response.model, + getattr(getattr(response, "usage", None), "total_tokens", "?"), + len(content_blocks), + ) + + return _core.CreateMessageResultWithTools( + 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).""" + self._tool_loop_count = 0 + response_text = choice.message.content or "" + + logger.log( + self.audit_level, + "MCP server '%s' sampling response: model=%s, tokens=%s", + self.server_name, response.model, + getattr(getattr(response, "usage", None), "total_tokens", "?"), + ) + + return _core.CreateMessageResult( + role="assistant", + content=_core.TextContent(type="text", text=_sanitize_error(response_text)), + model=response.model, + 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(), + ), + } + + async def __call__(self, context, params): + """SDK sampling callback (``SamplingFnT``). Returns CreateMessageResult, + CreateMessageResultWithTools, or ErrorData.""" + if not self._check_rate_limit(): + logger.warning( + "MCP server '%s' sampling rate limit exceeded (%d/min)", + self.server_name, self.max_rpm, + ) + self.metrics["errors"] += 1 + return self._error( + f"Sampling rate limit exceeded for server '{self.server_name}' " + f"({self.max_rpm} requests/minute)" + ) + + model = self._resolve_model( + mcp_field(params, "model_preferences", "modelPreferences") + ) + + from agent.auxiliary_client import call_llm + + 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, + ) + self.metrics["errors"] += 1 + return self._error( + f"Model '{resolved_model}' not allowed for server " + f"'{self.server_name}'. Allowed: {', '.join(self.allowed_models)}" + ) + + messages = self._convert_messages(params) + system_prompt = mcp_field(params, "system_prompt", "systemPrompt") + if system_prompt: + messages.insert(0, {"role": "system", "content": system_prompt}) + + max_tokens = min( + mcp_field(params, "max_tokens", "maxTokens", self.max_tokens_cap), + self.max_tokens_cap, + ) + call_temperature = None + if hasattr(params, "temperature") and params.temperature is not None: + call_temperature = params.temperature + + # Forward server-provided tools. + call_tools = None + server_tools = getattr(params, "tools", None) + if server_tools: + call_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 + ] + + 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), + ) + + # Offload the sync LLM call so the MCP loop is not blocked. + def _sync_call(): + return call_llm( + task="mcp", + model=resolved_model or None, + messages=messages, + temperature=call_temperature, + max_tokens=max_tokens, + tools=call_tools, + timeout=self.timeout, + ) + + try: + response = await asyncio.wait_for( + asyncio.to_thread(_sync_call), timeout=self.timeout, + ) + except asyncio.TimeoutError: + self.metrics["errors"] += 1 + return self._error( + f"Sampling LLM call timed out after {self.timeout}s " + f"for server '{self.server_name}'" + ) + except Exception as exc: + self.metrics["errors"] += 1 + return self._error( + f"Sampling LLM call failed: {_sanitize_error(_exc_str(exc))}" + ) + + # Empty choices happen on content filtering / provider errors. + if not getattr(response, "choices", None): + self.metrics["errors"] += 1 + return self._error( + f"LLM returned empty response (no choices) for server " + f"'{self.server_name}'" + ) + + choice = response.choices[0] + self.metrics["requests"] += 1 + total_tokens = getattr(getattr(response, "usage", None), "total_tokens", 0) + if isinstance(total_tokens, int): + self.metrics["tokens_used"] += total_tokens + + if ( + choice.finish_reason == "tool_calls" + and hasattr(choice.message, "tool_calls") + and choice.message.tool_calls + ): + return self._build_tool_use_result(choice, response) + + return self._build_text_result(choice, response) + + +def _format_elicitation_schema_summary(schema: dict, server_name: str) -> str: + """Render a flat-object requested_schema as a human-readable field list + (names, types, descriptions) so the user knows what they're approving.""" + 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(): + field_type = "" + field_desc = "" + if isinstance(field_spec, dict): + field_type = str(field_spec.get("type", "") or "") + field_desc = str(field_spec.get("description", "") or "") + suffix = f" ({field_type})" if field_type else "" + if field_desc: + lines.append(f" - {field_name}{suffix}: {field_desc}") + else: + lines.append(f" - {field_name}{suffix}") + return "\n".join(lines) + + +class ElicitationHandler: + """Handles ``elicitation/create`` requests for one MCP server. + + Callable; 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. + _OUTER_TIMEOUT_GRACE_SECONDS = 5 + + 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. + 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. + self.owner = owner + self.metrics = { + "requests": 0, + "accepted": 0, + "declined": 0, + "errors": 0, + } + + def session_kwargs(self) -> dict: + """Kwargs to pass to ClientSession for elicitation support.""" + return {"elicitation_callback": self} + + 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. + mode = getattr(params, "mode", "form") + if mode == "url": + logger.info( + "MCP server '%s' requested URL-mode elicitation; " + "declining (URL-mode elicitation not implemented)", + self.server_name, + ) + self.metrics["declined"] += 1 + return _core.ElicitResult(action="decline") + + 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. + 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: + from tools.approval import request_elicitation_consent + except Exception as exc: # pragma: no cover -- defensive + logger.error( + "MCP server '%s' elicitation: approval system unavailable: %s", + self.server_name, exc, + ) + self.metrics["errors"] += 1 + return _core.ElicitResult(action="decline") + + # Offload the sync consent flow to a thread — inline it would freeze the + # MCP loop and every other RPC on this session. The recv-loop task does + # NOT inherit the agent's contextvars, so replay the snapshot captured on + # owner._pending_call_context for gateway-platform detection. + captured = getattr(self.owner, "_pending_call_context", None) if self.owner else None + + def _invoke_consent() -> str: + if captured is None: + return request_elicitation_consent( + message, + description, + timeout_seconds=int(self.timeout), + surface=f"mcp-elicitation/{self.server_name}", + ) + # Context.run executes a context once — copy so multiple + # elicitations within one tool call work. + return captured.copy().run( + request_elicitation_consent, + message, + description, + timeout_seconds=int(self.timeout), + surface=f"mcp-elicitation/{self.server_name}", + ) + + try: + answer = await asyncio.wait_for( + 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), + ) + self.metrics["errors"] += 1 + return _core.ElicitResult(action="cancel") + except Exception as exc: + logger.error( + "MCP server '%s' elicitation failed: %s", + self.server_name, exc, exc_info=True, + ) + self.metrics["errors"] += 1 + return _core.ElicitResult(action="decline") + + if answer == "accept": + self.metrics["accepted"] += 1 + return _core.ElicitResult(action="accept", content={}) + if answer == "cancel": + self.metrics["errors"] += 1 + return _core.ElicitResult(action="cancel") + self.metrics["declined"] += 1 + return _core.ElicitResult(action="decline") diff --git a/tools/mcp_tool_schema.py b/tools/mcp_tool_schema.py new file mode 100644 index 0000000000..db8f43472b --- /dev/null +++ b/tools/mcp_tool_schema.py @@ -0,0 +1,307 @@ +"""MCP tool schema conversion and naming: JSON-schema normalisation for provider +compatibility, mcp__server__tool naming, utility-tool schemas, include/exclude +filters and description injection scanning.""" + +import logging +import fnmatch +import re +from typing import Any, List +from tools.ansi_strip import strip_unicode_tags +from tools.mcp_tool_common import mcp_field + +logger = logging.getLogger("tools.mcp_tool") + + +# Prompt-injection indicators in MCP tool descriptions. WARNING-level only: +# log but never block, since false positives would break legitimate servers. +_MCP_INJECTION_PATTERNS = [ + (re.compile(r"ignore\s+(all\s+)?previous\s+instructions", re.I), + "prompt override attempt ('ignore previous instructions')"), + (re.compile(r"you\s+are\s+now\s+a", re.I), + "identity override attempt ('you are now a...')"), + (re.compile(r"your\s+new\s+(task|role|instructions?)\s+(is|are)", re.I), + "task override attempt"), + (re.compile(r"system\s*:\s*", re.I), + "system prompt injection attempt"), + (re.compile(r"<\s*(system|human|assistant)\s*>", re.I), + "role tag injection attempt"), + (re.compile(r"do\s+not\s+(tell|inform|mention|reveal)", re.I), + "concealment instruction"), + (re.compile(r"(curl|wget|fetch)\s+https?://", re.I), + "network command in description"), + (re.compile(r"base64\.(b64decode|decodebytes)", re.I), + "base64 decode reference"), + (re.compile(r"exec\s*\(|eval\s*\(", re.I), + "code execution reference"), + (re.compile(r"import\s+(subprocess|os|shutil|socket)", re.I), + "dangerous import reference"), +] + + +def _scan_mcp_description(server_name: str, tool_name: str, description: str) -> List[str]: + """Scan a tool description for injection patterns; returns finding strings + (empty = clean) and logs a warning when any match.""" + findings = [] + if not description: + return findings + for pattern, reason in _MCP_INJECTION_PATTERNS: + if pattern.search(description): + findings.append(reason) + if findings: + logger.warning( + "MCP server '%s' tool '%s': suspicious description content — %s. " + "Description: %.200s", + server_name, tool_name, "; ".join(findings), + description, + ) + return findings + + +def _normalize_mcp_input_schema(schema: dict | None) -> dict: + """Normalize MCP input schemas so one form is valid on OpenAI, Anthropic, + Gemini and Moonshot. + + Repairs, applied recursively: ``definitions``/``#/definitions/`` refs -> + ``$defs`` (Moonshot rejects the draft-07 form); missing/null ``type`` on an + object-shaped node -> ``"object"``; an object without ``properties`` gets an + empty one so ``required`` can't dangle; ``required`` pruned to names present + in ``properties`` (Gemini 400s otherwise); nullable ``anyOf`` unions + collapsed to the non-null branch (Anthropic rejects nullable branches), + optionality living solely in the parent's ``required``. + """ + if not schema: + return {"type": "object", "properties": {}} + + def _rewrite_local_refs(node): + """Promote legacy ``definitions`` to ``$defs`` — but ONLY where it is a + JSON Schema meta-keyword, never as a property NAME inside ``properties``/ + ``patternProperties``. A tool parameter legitimately named ``definitions`` + rewritten to ``$defs`` would 400 the whole tool array (Anthropic/OpenAI + forbid ``$`` in property names). Property names are kept verbatim and + recursion resumes ordinary semantics inside each property's schema.""" + if isinstance(node, dict): + normalized = {} + for key, value in node.items(): + if key in ("properties", "patternProperties") and isinstance(value, dict): + normalized[key] = { + prop_name: _rewrite_local_refs(prop_schema) + for prop_name, prop_schema in value.items() + } + else: + out_key = "$defs" if key == "definitions" else key + normalized[out_key] = _rewrite_local_refs(value) + ref = normalized.get("$ref") + if isinstance(ref, str) and ref.startswith("#/definitions/"): + normalized["$ref"] = "#/$defs/" + ref[len("#/definitions/"):] + return normalized + if isinstance(node, list): + return [_rewrite_local_refs(item) for item in node] + return node + + def _strip_nullable_union(node): + """Shared implementation with the Anthropic guard and global sanitizer. + Keeps the ``nullable: true`` hint so runtime coercion can still map a + model-emitted ``"null"`` string to ``None``.""" + from tools.schema_sanitizer import strip_nullable_unions + + return strip_nullable_unions(node, keep_nullable_hint=True) + + def _collapse_const_unions(node): + """Collapse anyOf/oneOf unions of same-typed consts to enums. Must run + AFTER the nullable strip: consts -> enum, null branch -> ``nullable`` hint.""" + from tools.schema_sanitizer import collapse_const_unions + + return collapse_const_unions(node) + + def _repair_object_shape(node): + """Recursively fill missing object ``type``, ensure ``properties``, prune ``required``.""" + if isinstance(node, list): + return [_repair_object_shape(item) for item in node] + if not isinstance(node, dict): + return node + + repaired = {k: _repair_object_shape(v) for k, v in node.items()} + + if not repaired.get("type") and ( + "properties" in repaired or "required" in repaired + ): + repaired["type"] = "object" + + if repaired.get("type") == "object": + if not isinstance(repaired.get("properties"), dict): + repaired["properties"] = {} + + required = repaired.get("required") + if isinstance(required, list): + props = repaired.get("properties") or {} + valid = [r for r in required if isinstance(r, str) and r in props] + if len(valid) != len(required): + if valid: + repaired["required"] = valid + else: + repaired.pop("required", None) + + return repaired + + normalized = _rewrite_local_refs(schema) + normalized = _strip_nullable_union(normalized) + normalized = _collapse_const_unions(normalized) + normalized = _repair_object_shape(normalized) + + if not isinstance(normalized, dict): + return {"type": "object", "properties": {}} + if normalized.get("type") == "object" and "properties" not in normalized: + normalized = {**normalized, "properties": {}} + + return normalized + + +def sanitize_mcp_name_component(value: str) -> str: + """Replace every char outside ``[A-Za-z0-9_]`` with ``_`` (hyphens included, + the historical behavior) so generated names pass provider validation.""" + return re.sub(r"[^A-Za-z0-9_]", "_", str(value or "")) + + +# ``mcp____``: the convention shared by Claude Code, Codex and +# OpenCode. The double underscore disambiguates the server/tool boundary even +# when either contains underscores, and matches the Anthropic-OAuth wire form. +MCP_TOOL_NAME_PREFIX = "mcp__" + + +_MCP_NAME_DELIM = "__" + + +def mcp_prefixed_tool_name(server_name: str, tool_name: str) -> str: + """Registry/wire name: ``mcp____``.""" + safe_server = sanitize_mcp_name_component(server_name) + safe_tool = sanitize_mcp_name_component(tool_name) + return f"{MCP_TOOL_NAME_PREFIX}{safe_server}{_MCP_NAME_DELIM}{safe_tool}" + + +def _convert_mcp_schema(server_name: str, mcp_tool) -> dict: + """Convert an MCP ``Tool`` (``.input_schema``, or ``.inputSchema`` before + mcp 2.0) to a ``registry.register(schema=...)`` dict.""" + return { + "name": mcp_prefixed_tool_name(server_name, mcp_tool.name), + "description": strip_unicode_tags( + mcp_tool.description or f"MCP tool {mcp_tool.name} from {server_name}" + ), + "parameters": _normalize_mcp_input_schema( + mcp_field(mcp_tool, "input_schema", "inputSchema") + ), + } + + +def _build_utility_schemas(server_name: str) -> List[dict]: + """Schemas for the resource/prompt utility tools as ``{schema, handler_key}`` dicts.""" + return [ + { + "schema": { + "name": mcp_prefixed_tool_name(server_name, "list_resources"), + "description": f"List available resources from MCP server '{server_name}'", + "parameters": { + "type": "object", + "properties": {}, + }, + }, + "handler_key": "list_resources", + }, + { + "schema": { + "name": mcp_prefixed_tool_name(server_name, "read_resource"), + "description": f"Read a resource by URI from MCP server '{server_name}'", + "parameters": { + "type": "object", + "properties": { + "uri": { + "type": "string", + "description": "URI of the resource to read", + }, + }, + "required": ["uri"], + }, + }, + "handler_key": "read_resource", + }, + { + "schema": { + "name": mcp_prefixed_tool_name(server_name, "list_prompts"), + "description": f"List available prompts from MCP server '{server_name}'", + "parameters": { + "type": "object", + "properties": {}, + }, + }, + "handler_key": "list_prompts", + }, + { + "schema": { + "name": mcp_prefixed_tool_name(server_name, "get_prompt"), + "description": f"Get a prompt by name from MCP server '{server_name}'", + "parameters": { + "type": "object", + "properties": { + "name": { + "type": "string", + "description": "Name of the prompt to retrieve", + }, + "arguments": { + "type": "object", + "description": "Optional arguments to pass to the prompt", + "properties": {}, + "additionalProperties": True, + }, + }, + "required": ["name"], + }, + }, + "handler_key": "get_prompt", + }, + ] + + +def _normalize_name_filter(value: Any, label: str) -> set[str]: + """Normalize include/exclude config to a set of exact names or fnmatch globs.""" + if value is None: + return set() + if isinstance(value, str): + return {value} + if isinstance(value, (list, tuple, set)): + return {str(item) for item in value} + logger.warning("MCP config %s must be a string or list of strings; ignoring %r", label, value) + return set() + + +def matches_name_filter(tool_name: str, patterns: set[str]) -> bool: + """True if ``tool_name`` matches any entry: exact names literally, entries + with ``*``/``?``/``[`` as case-sensitive globs (same semantics as + ``approvals.deny``). Exact membership is checked first so big lists stay O(1).""" + if not patterns: + return False + if tool_name in patterns: + return True + return any( + fnmatch.fnmatchcase(tool_name, p) + for p in patterns + if "*" in p or "?" in p or "[" in p + ) + + +_UTILITY_CAPABILITY_METHODS = { + "list_resources": "list_resources", + "read_resource": "read_resource", + "list_prompts": "list_prompts", + "get_prompt": "get_prompt", +} + + +# Utility handler -> capability key that must be non-None on the server's +# ``initialize`` response for the handler to be registered. Without this gate a +# tools-only server got all four stubs and every call returned JSON-RPC -32601, +# making the model conclude the server was broken. +_UTILITY_CAPABILITY_ATTRS = { + "list_resources": "resources", + "read_resource": "resources", + "list_prompts": "prompts", + "get_prompt": "prompts", +} diff --git a/tools/mcp_tool_transport.py b/tools/mcp_tool_transport.py new file mode 100644 index 0000000000..8a53d549f2 --- /dev/null +++ b/tools/mcp_tool_transport.py @@ -0,0 +1,676 @@ +"""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 +import os +from typing import Dict, Optional +from tools.mcp_tool_config import _wrap_command_with_watchdog +from tools.mcp_tool_errors import NonMcpEndpointError, _apply_identity_header, _handshake_rejected_as_modern, _make_redirect_header_stripper, _resolve_client_cert +from tools.mcp_tool_lifecycle import _filter_mcp_children, _orphan_stdio_pid_servers, _orphan_stdio_pids, _stdio_pgids, _stdio_pids +from tools.mcp_tool_common import _core + +logger = logging.getLogger("tools.mcp_tool") + + +class MCPServerTransportMixin: + """Methods of :class:`tools.mcp_tool.MCPServerTask` (mixed in; relies on its attributes).""" + + __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: keep the old always-call-list_tools behavior). + """ + init_result = self.initialize_result + caps = getattr(init_result, "capabilities", None) if init_result is not None else None + if caps is None: + return True + return getattr(caps, "tools", None) is not None + + 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 error); ``legacy`` is handshake only, no fallback. Both result + types expose ``.capabilities``, so downstream gates work on either. + """ + mode = str((self._config or {}).get("protocol", "auto")).lower().strip() + if mode in ("stateless", "modern", "2026-07-28"): + try: + return await asyncio.wait_for( + session.discover(), timeout=connect_timeout + ) + except asyncio.TimeoutError: + raise + except Exception as exc: + logger.info( + "MCP server '%s': server/discover rejected (%s) despite " + "protocol=%s — falling back to the legacy handshake", + self.name, exc, mode, + ) + return await asyncio.wait_for( + session.initialize(), timeout=connect_timeout + ) + if mode in ("legacy", "handshake"): + return await asyncio.wait_for( + session.initialize(), timeout=connect_timeout + ) + if mode != "auto": + logger.warning( + "MCP server '%s': unknown protocol=%r — treating as 'auto' " + "(valid: auto, stateless, legacy)", self.name, mode, + ) + try: + return await asyncio.wait_for( + session.initialize(), timeout=connect_timeout + ) + except asyncio.TimeoutError: + raise + except Exception as exc: + if not _handshake_rejected_as_modern(exc): + raise + if not hasattr(session, "discover"): + # mcp 1.x has no server/discover client — nothing to fall back to. + raise + logger.info( + "MCP server '%s': legacy handshake rejected (%s) — " + "retrying via server/discover (2026-07-28 stateless server)", + self.name, exc, + ) + return await asyncio.wait_for( + session.discover(), timeout=connect_timeout + ) + + 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. + """ + self.initialize_result = await self._negotiate_session(session, connect_timeout) + self.session = session + if mark_lifecycle: + self._mark_lifecycle_started() + await self._discover_tools() + self._ready.set() + self._ever_connected = True + _core._reset_server_error(self.name) + self._session_proven = False + reason = await self._wait_for_lifecycle_event() + if label and reason == "reconnect": + logger.info( + "MCP server '%s': reconnect requested — tearing down %s session", + self.name, label, + ) + return reason + + async def _run_stdio(self, config: dict): + """Run the server using stdio transport.""" + if config.get("identity_header") is not None: + # No headers on stdio — warn so a copy-pasted HTTP block doesn't mislead. + logger.warning( + "MCP server '%s': identity_header is only supported on " + "HTTP/SSE transports — ignored for stdio servers", self.name, + ) + if not _core._ensure_mcp_sdk(): + raise ImportError( + f"MCP server '{self.name}' requires the 'mcp' Python SDK, but " + "it is not installed. Run `hermes setup` to install MCP support, " + "then retry." + ) + + command = config.get("command") + args = config.get("args", []) + user_env = config.get("env") + + if not command: + raise ValueError( + f"MCP server '{self.name}' has no 'command' in config" + ) + + safe_env = _core._build_safe_env(user_env) + command, safe_env = _core._resolve_stdio_command(command, safe_env) + + # 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 below rewrites + # argv to the supervisor, which would turn the check into a no-op. + from tools.osv_check import check_package_for_malware + try: + malware_error = await asyncio.wait_for( + asyncio.to_thread(check_package_for_malware, command, args), + timeout=_core._OSV_MALWARE_CHECK_TIMEOUT_S, + ) + except asyncio.TimeoutError: + logger.warning( + "MCP server '%s': OSV malware preflight timed out after %.0fs " + "(network slow/unreachable) — proceeding without the check.", + self.name, _core._OSV_MALWARE_CHECK_TIMEOUT_S, + ) + malware_error = None + if malware_error: + raise ValueError( + f"MCP server '{self.name}': {malware_error}" + ) + + # Parent-death watchdog: an ungraceful Hermes exit (kill -9, crash) can't + # leave the child and its descendants running. Clean-exit reaping is + # unchanged. 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", + ) + + sampling_kwargs = self._sampling.session_kwargs() if self._sampling else {} + if self._elicitation: + sampling_kwargs.update(self._elicitation.session_kwargs()) + if _core._MCP_NOTIFICATION_TYPES and _core._MCP_MESSAGE_HANDLER_SUPPORTED: + sampling_kwargs["message_handler"] = self._make_message_handler() + if _core._MCP_LOGGING_CALLBACK_SUPPORTED: + sampling_kwargs["logging_callback"] = self._make_logging_callback() + + # 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) and would stall the loop. + 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. + _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. + new_pids = _filter_mcp_children( + _core._snapshot_child_pids() - pids_before + ) + if new_pids: + # Capture pgid while the child is alive — getpgid fails once it + # exits, and the sweep needs it to reach reparented descendants. + new_pgids: Dict[int, int] = {} + for _pid in new_pids: + try: + new_pgids[_pid] = os.getpgid(_pid) + except (AttributeError, ProcessLookupError, OSError): + # AttributeError: Windows; ProcessLookupError: already exited. + pass + with _core._lock: + 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. + 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, + ) + # 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, **sampling_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. + connect_timeout = float( + config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT) + ) + return await self._serve_session( + session, connect_timeout, mark_lifecycle=True + ) + finally: + # Runs on clean exit, exceptions AND cancellation. Any spawned PID + # 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. + if new_pids: + from gateway.status import _pid_exists + _killpg = getattr(os, "killpg", None) + with _core._lock: + for _pid in new_pids: + _stdio_pids.pop(_pid, None) + for pid in new_pids: + # ``os.kill(pid, 0)`` is NOT a no-op on Windows; use the + # cross-platform check. + pid_alive = _pid_exists(pid) + pgroup_alive = False + pgid = _stdio_pgids.get(pid) + if not pid_alive and pgid is not None and _killpg is not None: + # Child exited but descendants may remain in its pgroup; + # signal 0 succeeds iff any member is alive. + try: + _killpg(pgid, 0) + pgroup_alive = True + except (ProcessLookupError, PermissionError, OSError): + pgroup_alive = False + if pid_alive or pgroup_alive: + _orphan_stdio_pids.add(pid) + _orphan_stdio_pid_servers[pid] = self.name + else: + # Nothing to reap — drop the pgid so PID reuse can't + # surface stale pgroup state later. + _stdio_pgids.pop(pid, None) + + 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. + """ + try: + import httpx as _httpx + except ImportError: + return # No httpx → skip probe; SDK import would have failed first. + + client_kwargs: dict = { + "verify": ssl_verify, + "follow_redirects": True, + "timeout": _httpx.Timeout(timeout), + } + if client_cert is not None: + client_kwargs["cert"] = client_cert + + probe_headers = dict(headers) if headers else {} + try: + async with _httpx.AsyncClient(**client_kwargs) as client: + # HEAD is cheapest; fall back to GET on 405/501. + resp = await client.head(url, headers=probe_headers) + if resp.status_code in (405, 501): + resp = await client.get(url, headers=probe_headers) + + # Non-MCP content type on HEAD/GET: try a JSON-RPC POST before + # rejecting, so POST-only servers aren't false positives. + ct = ( + resp.headers.get("content-type", "") + .split(";")[0] + .strip() + .lower() + ) + if ( + ct + and ct not in self._MCP_CONTENT_TYPES + and 200 <= resp.status_code < 300 + ): + post_resp = await client.post( + url, + headers={ + **probe_headers, + "Content-Type": "application/json", + "Accept": "application/json, text/event-stream", + }, + content=( + '{"jsonrpc":"2.0","id":"_probe",' + '"method":"initialize",' + '"params":{"protocolVersion":"2025-03-26",' + '"capabilities":{},' + '"clientInfo":{"name":"hermes-probe",' + '"version":"0.1"}}}' + ), + ) + if 200 <= post_resp.status_code < 300: + post_ct = ( + post_resp.headers.get("content-type", "") + .split(";")[0] + .strip() + .lower() + ) + if post_ct in self._MCP_CONTENT_TYPES: + resp = post_resp + except _httpx.HTTPError: + return # DNS/connect/timeout/transport error — let the SDK try. + + # Only judge 2xx: a 4xx/5xx may be an auth challenge or transient error + # the real handshake handles correctly. + if not (200 <= resp.status_code < 300): + return + + ct_base = resp.headers.get("content-type", "").split(";")[0].strip().lower() + if not ct_base: + return # No content type advertised — don't second-guess the SDK. + if ct_base in self._MCP_CONTENT_TYPES: + return # Looks like a real MCP endpoint. + + raise NonMcpEndpointError( + f"MCP server '{self.name}' at {url} returned Content-Type " + f"'{ct_base}', not an MCP response (expected one of: " + f"{', '.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/)." + ) + + 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; ``"reconnect"`` rebuilds the session at once. + 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). + """ + if self._shutdown_event.is_set(): + raise eg + fatal, _rest = eg.split((KeyboardInterrupt, SystemExit)) + if fatal is not None: + raise eg + cancelled, _rest = eg.split(asyncio.CancelledError) + if cancelled is not None: + raise eg + if not self._ready.is_set(): + raise eg + logger.debug( + "MCP server '%s': transport TaskGroup exited after a live session " + "(%r) — reconnecting immediately instead of backing off", + self.name, eg, + ) + return "reconnect" + + async def _run_http(self, config: dict): + """Run the server using HTTP/StreamableHTTP transport.""" + _core._ensure_mcp_sdk() + if not _core._MCP_HTTP_AVAILABLE: + raise ImportError( + f"MCP server '{self.name}' requires HTTP transport but " + "mcp.client.streamable_http is not available. " + "Upgrade the mcp package to get HTTP support." + ) + + url = config["url"] + headers = dict(config.get("headers") or {}) + # Portable Agent Plugins v1 (strict_redirect_headers): configured + # headers MUST NOT follow a redirect to a different origin. Capture the + # configured names BEFORE client-generated headers are merged in. + _strict_cfg_headers = bool(config.get("strict_redirect_headers")) + _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). Seeded from the HANDSHAKE + # version, not the latest: the body sent by ``initialize()`` speaks the + # handshake era, and a 2026-07-28 header would route the request onto + # the server's per-request-envelope ladder, which rejects that body. + # The header must agree with what the body actually speaks. + if not any(key.lower() == "mcp-protocol-version" for key in headers): + headers["mcp-protocol-version"] = _core.LATEST_HANDSHAKE_VERSION + connect_timeout = config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT) + ssl_verify = config.get("ssl_verify", True) + client_cert = _resolve_client_cert(self.name, config) + + # OAuth 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_auth = None + if self._auth_type == "oauth": + try: + from tools.mcp_oauth_manager import get_manager + _oauth_auth = get_manager().get_or_build_provider( + self.name, url, config.get("oauth"), + ) + except Exception as exc: + logger.warning("MCP OAuth setup failed for '%s': %s", self.name, exc) + raise + + sampling_kwargs = self._sampling.session_kwargs() if self._sampling else {} + if self._elicitation: + sampling_kwargs.update(self._elicitation.session_kwargs()) + if _core._MCP_NOTIFICATION_TYPES and _core._MCP_MESSAGE_HANDLER_SUPPORTED: + sampling_kwargs["message_handler"] = self._make_message_handler() + if _core._MCP_LOGGING_CALLBACK_SUPPORTED: + sampling_kwargs["logging_callback"] = self._make_logging_callback() + + # SSE transport (``transport: sse`` in the mcp_servers entry). + if config.get("transport") == "sse": + if _strict_cfg_headers: + # Fail closed: SSE cannot enforce the redirect boundary. + raise ValueError( + f"MCP server '{self.name}': strict_redirect_headers is " + "not supported on the SSE transport." + ) + if _core.sse_client is None: + 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 below. + _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 + 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(). + _httpx_mod = _core.sdk_httpx() + + _cert_for_factory = client_cert + _verify_for_factory = ssl_verify + + def _mcp_http_client_factory( + headers=None, timeout=None, auth=None, + ): + kwargs: dict = { + "follow_redirects": True, + "verify": _verify_for_factory, + } + if timeout is not None: + kwargs["timeout"] = timeout + else: + kwargs["timeout"] = _httpx_mod.Timeout(30.0, read=300.0) + if headers is not None: + kwargs["headers"] = headers + if auth is not None: + kwargs["auth"] = auth + if _cert_for_factory is not None: + kwargs["cert"] = _cert_for_factory + return _httpx_mod.AsyncClient(**kwargs) + + _sse_kwargs["httpx_client_factory"] = _mcp_http_client_factory + try: + async with _core.sse_client(**_sse_kwargs) as (read_stream, write_stream): + async with _core.ClientSession( + read_stream, write_stream, **sampling_kwargs + ) as session: + reason = await self._serve_session( + session, float(connect_timeout), "SSE" + ) + except BaseExceptionGroup as _eg: + # Transport TaskGroup dropped: reconnect instead of backoff/park. + reason = self._reconnect_or_reraise_group(_eg) + return reason + + if _core._MCP_NEW_HTTP: + # mcp >= 1.24.0: 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(). + httpx = _core.sdk_httpx() + + _original_url = httpx.URL(url) + + _strip_auth_on_cross_origin_redirect = _make_redirect_header_stripper( + _original_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 + if _oauth_auth is not None: + client_kwargs["auth"] = _oauth_auth + if client_cert is not None: + client_kwargs["cert"] = client_cert + + # Caller owns the client lifecycle — the SDK skips cleanup when + # http_client is provided. + try: + async with httpx.AsyncClient(**client_kwargs) as http_client: + # Unpacked positionally: mcp 1.x yields (read, write, + # get_session_id), 2.x yields (read, write). + async with _core.streamable_http_client(url, http_client=http_client) as _streams: + read_stream, write_stream = _streams[0], _streams[1] + async with _core.ClientSession(read_stream, write_stream, **sampling_kwargs) as session: + reason = await self._serve_session( + session, float(connect_timeout), "HTTP" + ) + except BaseExceptionGroup as _eg: + # Transport TaskGroup dropped: reconnect instead of backoff/park. + reason = self._reconnect_or_reraise_group(_eg) + return reason + else: + # Deprecated API (mcp < 1.24.0): the SDK owns the httpx client. + if _strict_cfg_headers: + # Fail closed: without an owned client we cannot hook redirects, + # so the cross-origin header boundary cannot be enforced. + raise ImportError( + f"MCP server '{self.name}' requires mcp >= 1.24.0 to " + "enforce the portable redirect-header boundary " + "(strict_redirect_headers). Upgrade the mcp package." + ) + _http_kwargs: dict = { + "headers": headers, + "timeout": float(connect_timeout), + "verify": ssl_verify, + } + if _oauth_auth is not None: + _http_kwargs["auth"] = _oauth_auth + try: + async with _core.streamablehttp_client(url, **_http_kwargs) as ( + read_stream, write_stream, _get_session_id, + ): + async with _core.ClientSession(read_stream, write_stream, **sampling_kwargs) as session: + reason = await self._serve_session( + session, float(connect_timeout), "legacy HTTP" + ) + except BaseExceptionGroup as _eg: + # Transport TaskGroup dropped: reconnect instead of backoff/park. + reason = self._reconnect_or_reraise_group(_eg) + return reason + + 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 — skip the call when + ``tools`` isn't advertised. + """ + # Fresh transport: re-probe with cheap ``ping`` in case the server gained + # support across the reconnect. + self._ping_unsupported = False + if self.session is None: + return + if not self._advertises_tools(): + logger.info( + "MCP server '%s': does not advertise 'tools' capability — " + "skipping tools/list (prompts/resources remain available)", + self.name, + ) + self._tools = [] + self._register_discovered_tools_if_needed() + return + async with self._rpc_lock: + self._list_cache_meta = {} + self._tools = await _core._paginate_full_list( + self.session.list_tools, "tools", self.name, + cache_meta_out=self._list_cache_meta, + ) + 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. + """ + if self._registered_tool_names: + return + if not self._ready.is_set(): + with _core._lock: + 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. + with _core._lock: + if _core._servers.get(self.name) is self: + _core._server_connect_errors.pop(self.name, None)