diff --git a/tools/mcp_schema_cache.py b/tools/mcp_schema_cache.py index 1218a23593..cdf4e0abff 100644 --- a/tools/mcp_schema_cache.py +++ b/tools/mcp_schema_cache.py @@ -1,10 +1,7 @@ -"""Persistent MCP tool-schema cache for lazy server startup. - -Stores per-server tool manifests on disk so Hermes can register MCP tools -into the agent snapshot without spawning the stdio child process at idle -dashboard startup. Cache entries are keyed by server name + a fingerprint -of the connection config (command/args/url/tools filters). -""" +"""Persistent MCP tool-schema cache for lazy server startup: per-server tool manifests on +disk so Hermes can register MCP tools into the agent snapshot without spawning the stdio +child at idle dashboard startup. Entries are keyed by server name + a fingerprint of the +connection config (command/args/url/tools filters).""" from __future__ import annotations @@ -58,43 +55,31 @@ def _load_all() -> Dict[str, Any]: def _save_all(data: Dict[str, Any]) -> None: from utils import atomic_json_write - # 0o600 (as tools/registry.py _save_discovery_cache): the cache file is - # trusted input on the lazy registration path, so keep it user-only. + # 0o600: the cache file is trusted input on the lazy registration path, keep it user-only. atomic_json_write(_cache_path(), data, mode=0o600) def get_cached_entry(server_name: str, fingerprint: str) -> Optional[dict]: - """Return cached entry when fingerprint matches (and TTL holds), else None. - ``tools/list`` results may carry ``ttlMs`` (SEP-2549); an entry older than a - recorded TTL is a miss so the next startup re-probes instead of serving a - stale manifest forever. Entries without a TTL never expire. ``cacheScope`` - is irrelevant: this cache is per-user local disk, satisfying even ``private``.""" + """Return cached entry when fingerprint matches (and TTL holds), else None. ``tools/list`` + results may carry ``ttlMs`` (SEP-2549); an entry older than a recorded TTL is a miss so the + next startup re-probes instead of serving a stale manifest forever. Entries without a TTL + never expire. ``cacheScope`` is irrelevant: this cache is per-user local disk.""" with _cache_lock: entry = _load_all().get(server_name) if not isinstance(entry, dict) or entry.get("fingerprint") != fingerprint: return None ttl_ms = entry.get("ttl_ms") written_at = entry.get("written_at") - expired = ( - isinstance(ttl_ms, (int, float)) - and isinstance(written_at, (int, float)) - and (time.time() - written_at) * 1000.0 >= float(ttl_ms) - ) + expired = (isinstance(ttl_ms, (int, float)) and isinstance(written_at, (int, float)) + and (time.time() - written_at) * 1000.0 >= float(ttl_ms)) return None if expired else entry -def write_cache_entry( - server_name: str, - fingerprint: str, - *, - tools: List[dict], - utility_tools: Optional[List[dict]] = None, - ttl_ms: Optional[float] = None, - cache_scope: Optional[str] = None, -) -> None: - """Persist tool schemas after a successful live connect. ``ttl_ms`` / - ``cache_scope`` are the server's ``tools/list`` SEP-2549 hints; - ``written_at`` anchors TTL expiry in :func:`get_cached_entry`.""" +def write_cache_entry(server_name: str, fingerprint: str, *, tools: List[dict], + utility_tools: Optional[List[dict]] = None, ttl_ms: Optional[float] = None, + cache_scope: Optional[str] = None) -> None: + """Persist tool schemas after a successful live connect. ``ttl_ms`` / ``cache_scope`` are + the server's ``tools/list`` SEP-2549 hints; ``written_at`` anchors TTL expiry.""" entry = {"fingerprint": fingerprint, "tools": tools, "utility_tools": utility_tools or []} if isinstance(ttl_ms, (int, float)): entry["ttl_ms"] = ttl_ms @@ -103,10 +88,9 @@ def write_cache_entry( entry["cache_scope"] = cache_scope with _cache_lock: data = _load_all() - # Write-through fires on every registration (reconnects, list_changed); - # skip the load-all+rewrite churn when the entry is byte-identical on - # disk. TTL'd entries always rewrite: written_at must advance or the - # entry would expire at its ORIGINAL write time regardless of reconnects. + # Write-through fires on every registration (reconnects, list_changed); skip the + # rewrite when the entry is byte-identical on disk. TTL'd entries always rewrite: + # written_at must advance or the entry would expire at its ORIGINAL write time. if "written_at" not in entry and data.get(server_name) == entry: return data[server_name] = entry diff --git a/tools/mcp_tool.py b/tools/mcp_tool.py index 197f855b48..05512ebae6 100644 --- a/tools/mcp_tool.py +++ b/tools/mcp_tool.py @@ -1,56 +1,17 @@ #!/usr/bin/env python3 -""" -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. +"""MCP (Model Context Protocol) client: connects to the ``mcp_servers`` configured in +~/.hermes/config.yaml 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. -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 (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: {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 - sampling: {enabled: true, model: "gemini-3-flash", max_tokens_cap: 4096, - timeout: 30, max_rpm: 10, allowed_models: [], max_tool_rounds: 5} - -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``. - -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_server_run`` -/ ``mcp_tool_transport`` / ``mcp_tool_health`` (MCPServerTask mixins: run state -machine, transport bring-up, keepalive/liveness), ``mcp_tool_loop`` (discovery -lock, loop thread, cross-thread scheduling and reconnect signalling), -``mcp_tool_discovery`` (connect, lazy start, ``discover_mcp_tools`` and the status -API), ``mcp_tool_lifecycle`` (shutdown, orphan reaping), ``mcp_tool_agent`` -(live-agent tool list refresh). This module keeps the SDK loader, the -``MCPServerTask`` shell and every piece of shared module state (siblings read -it back through ``tools.mcp_tool`` at call time, never by value). +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``; every mutation of ``_servers`` / loop handles holds ``_lock``. +This module keeps the SDK loader, the ``MCPServerTask`` shell and all shared module state; +the ``mcp_tool_*`` siblings read it back through ``tools.mcp_tool`` at call time (never by +value) and every one of their names is re-exported here so ``from tools.mcp_tool import X`` +and ``mock.patch("tools.mcp_tool.X")`` keep working. """ import asyncio @@ -67,46 +28,39 @@ from typing import Any, Callable, Dict, List, Optional, Set logger = logging.getLogger(__name__) -# 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 - _DEFAULT_TOOL_TIMEOUT, _env_ref_name, _exc_str, _get_lifecycle_seconds, - _jittered, _parse_boolish, _resolve_tool_timeout, _safe_numeric, - _sanitize_error, mcp_field, + _DEFAULT_TOOL_TIMEOUT, _env_ref_name, _exc_str, _get_lifecycle_seconds, _jittered, + _parse_boolish, _resolve_tool_timeout, _safe_numeric, _sanitize_error, mcp_field, ) from tools.mcp_tool_schema import ( # noqa: F401 - MCP_TOOL_NAME_PREFIX, _build_utility_schemas, _convert_mcp_schema, - _normalize_mcp_input_schema, _scan_mcp_description, matches_name_filter, - mcp_prefixed_tool_name, sanitize_mcp_name_component, + MCP_TOOL_NAME_PREFIX, _build_utility_schemas, _convert_mcp_schema, _normalize_mcp_input_schema, + _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, _truncate_mcp_text_result, + _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, + _truncate_mcp_text_result, ) from tools.mcp_tool_errors import ( # noqa: F401 InvalidMcpUrlError, NonMcpEndpointError, _EXC_TRAVERSAL_MAX_NODES, - _JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, _classify_mcp_failure, - _format_connect_error, _handshake_rejected_as_modern, _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, + _JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, _classify_mcp_failure, _format_connect_error, + _handshake_rejected_as_modern, _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, _build_safe_env, _filter_suspicious_mcp_servers, - _get_mcp_stderr_log, _interpolate_env_vars, _load_mcp_config, - _resolve_stdio_command, _warn_hidden_whitespace, _whitespace_warned, - _workspace_folder, _wrap_command_with_watchdog, _write_stderr_log_header, + _ENV_VAR_PATTERN, _build_safe_env, _filter_suspicious_mcp_servers, _get_mcp_stderr_log, + _interpolate_env_vars, _load_mcp_config, _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 _handle_auth_error_and_retry, _handle_session_expired_and_retry, _make_check_fn, - _make_get_prompt_handler, _make_list_prompts_handler, - _make_list_resources_handler, _make_read_resource_handler, _make_tool_handler, + _make_get_prompt_handler, _make_list_prompts_handler, _make_list_resources_handler, + _make_read_resource_handler, _make_tool_handler, ) from tools.mcp_tool_registration import ( # noqa: F401 _annotation_read_only_hint, _existing_tool_names, _forget_mcp_tool_server, @@ -116,8 +70,7 @@ from tools.mcp_tool_registration import ( # noqa: F401 from tools.mcp_tool_lifecycle import ( # noqa: F401 _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, + _snapshot_child_pids, _stdio_pgids, _stdio_pids, _stop_mcp_loop_if_idle, shutdown_mcp_servers, ) from tools.mcp_tool_agent import ( # noqa: F401 _reinject_post_build_tools, persist_agent_tool_names, refresh_agent_mcp_tools, @@ -126,28 +79,24 @@ from tools.mcp_tool_agent import ( # noqa: F401 from tools.mcp_tool_transport import MCPServerTransportMixin from tools.mcp_tool_server_run import MCPServerRunMixin from tools.mcp_tool_health import MCPServerHealthMixin -from tools.mcp_tool_loop import ( # noqa: F401 -- re-exported for callers and test patches - _running_loop, - _LockCookie, _acquire_lock_on_fh, _try_acquire_mcp_discovery_lock, - _mcp_loop_exception_handler, _wrap_with_home_override, - _wrap_with_dashboard_oauth_flow, _run_on_mcp_loop, _signal_reconnect, - reconnect_mcp_server, _wait_for_server_session_ready, +from tools.mcp_tool_loop import ( # noqa: F401 + _running_loop, _LockCookie, _acquire_lock_on_fh, _try_acquire_mcp_discovery_lock, + _mcp_loop_exception_handler, _wrap_with_home_override, _wrap_with_dashboard_oauth_flow, + _run_on_mcp_loop, _signal_reconnect, reconnect_mcp_server, _wait_for_server_session_ready, _signal_reconnect_and_wait, _ensure_mcp_loop, _stop_mcp_loop, ) -from tools.mcp_tool_discovery import ( # noqa: F401 -- re-exported for callers and test patches - _record_connect_failure, _clear_connect_failure, _connect_cooldown_active, - _connect_server, _request_lazy_reconnect, _resolve_server_lazy, - _ensure_lazy_server_connected, _get_connected_server_for_call, - _discover_and_register_server, register_mcp_servers, discover_mcp_tools, - is_mcp_tool_parallel_safe, get_mcp_status, probe_mcp_server_tools, +from tools.mcp_tool_discovery import ( # noqa: F401 + _record_connect_failure, _clear_connect_failure, _connect_cooldown_active, _connect_server, + _request_lazy_reconnect, _resolve_server_lazy, _ensure_lazy_server_connected, + _get_connected_server_for_call, _discover_and_register_server, register_mcp_servers, + discover_mcp_tools, is_mcp_tool_parallel_safe, get_mcp_status, probe_mcp_server_tools, has_registered_mcp_tools, get_registered_mcp_server_names, ) -# 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). +# 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. _OSV_MALWARE_CHECK_TIMEOUT_S = 12.0 @@ -165,19 +114,17 @@ _MCP_ELICITATION_TYPES = False _MCP_MESSAGE_HANDLER_SUPPORTED = False _MCP_LOGGING_CALLBACK_SUPPORTED = False sse_client = None -# 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). +# 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" -# 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. +# 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. LATEST_HANDSHAKE_VERSION = LATEST_PROTOCOL_VERSION -# 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. +# 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: _MCP_AVAILABLE = importlib.util.find_spec("mcp") is not None except Exception: @@ -189,19 +136,24 @@ ClientSession: Any = None _MCP_SDK_IMPORT_ATTEMPTED = False _MCP_SDK_IMPORT_LOCK = threading.Lock() -# 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). +# 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", - "CreateMessageResult", "CreateMessageResultWithTools", "ErrorData", - "SamplingCapability", "SamplingToolsCapability", "TextContent", - "ToolUseContent", "ElicitRequestParams", "ElicitResult", - "ServerNotification", "ToolListChangedNotification", + "StdioServerParameters", "stdio_client", "streamablehttp_client", "streamable_http_client", + "CreateMessageResult", "CreateMessageResultWithTools", "ErrorData", "SamplingCapability", + "SamplingToolsCapability", "TextContent", "ToolUseContent", "ElicitRequestParams", + "ElicitResult", "ServerNotification", "ToolListChangedNotification", "PromptListChangedNotification", "ResourceListChangedNotification", }) +# Optional SDK type families: (module, names, debug message when absent). Each is gated +# separately so an older SDK only loses that feature, not MCP. +_SAMPLING_TYPE_NAMES = ("CreateMessageResult", "CreateMessageResultWithTools", "ErrorData", + "SamplingCapability", "SamplingToolsCapability", "TextContent", "ToolUseContent") +_NOTIFICATION_TYPE_NAMES = ("ServerNotification", "ToolListChangedNotification", + "PromptListChangedNotification", "ResourceListChangedNotification") + def __getattr__(name: str): if name in _MCP_SDK_LAZY_SYMBOLS: @@ -214,11 +166,8 @@ def __getattr__(name: str): def _import_sdk_names(module: str, names: tuple, missing_msg: Optional[str] = None) -> bool: - """Bind ``names`` from the SDK ``module`` into this module's globals. - - False (plus an optional debug line) when this SDK build lacks the module or - any of the names; nothing is bound in that case. - """ + """Bind ``names`` from SDK ``module`` into this module's globals; False (nothing bound, + optional debug line) when this SDK build lacks the module or any of the names.""" try: mod = importlib.import_module(module) values = {n: getattr(mod, n) for n in names} @@ -233,10 +182,9 @@ def _import_sdk_names(module: str, names: tuple, missing_msg: Optional[str] = No def _ensure_mcp_sdk() -> bool: """Import the optional ``mcp`` SDK on first use; return availability. - 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). Optional type families - are gated separately so an older SDK only loses that feature, not MCP. + 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 @@ -251,44 +199,30 @@ def _ensure_mcp_sdk() -> bool: with _MCP_SDK_IMPORT_LOCK: if _MCP_SDK_IMPORT_ATTEMPTED or ClientSession is not None: return _MCP_AVAILABLE - if ( - _import_sdk_names("mcp", ("ClientSession", "StdioServerParameters")) - and _import_sdk_names("mcp.client.stdio", ("stdio_client",)) - ): + if (_import_sdk_names("mcp", ("ClientSession", "StdioServerParameters")) + and _import_sdk_names("mcp.client.stdio", ("stdio_client",))): _MCP_AVAILABLE = True - # mcp >= 1.24 ships streamable_http_client; 2.0 dropped the - # deprecated streamablehttp_client alias. Either one gives HTTP. + # mcp >= 1.24 ships streamable_http_client; 2.0 dropped the deprecated + # streamablehttp_client alias. Either one gives HTTP. _MCP_NEW_HTTP = _import_sdk_names("mcp.client.streamable_http", ("streamable_http_client",)) _MCP_LEGACY_HTTP = _import_sdk_names("mcp.client.streamable_http", ("streamablehttp_client",)) _MCP_HTTP_AVAILABLE = _MCP_NEW_HTTP or _MCP_LEGACY_HTTP - _import_sdk_names( - "mcp.types", ("LATEST_PROTOCOL_VERSION",), - "mcp.types.LATEST_PROTOCOL_VERSION not available -- using fallback protocol version", - ) + _import_sdk_names("mcp.types", ("LATEST_PROTOCOL_VERSION",), + "mcp.types.LATEST_PROTOCOL_VERSION not available -- using fallback protocol version") if not _import_sdk_names("mcp.client.session", ("LATEST_HANDSHAKE_VERSION",)): # Pre-2.x SDKs: newest revision IS the handshake revision. LATEST_HANDSHAKE_VERSION = LATEST_PROTOCOL_VERSION - if not _import_sdk_names( - "mcp.client.sse", ("sse_client",), - "mcp.client.sse.sse_client not available -- SSE transport disabled", - ): + if not _import_sdk_names("mcp.client.sse", ("sse_client",), + "mcp.client.sse.sse_client not available -- SSE transport disabled"): sse_client = None _MCP_SAMPLING_TYPES = _import_sdk_names( - "mcp.types", - ("CreateMessageResult", "CreateMessageResultWithTools", "ErrorData", - "SamplingCapability", "SamplingToolsCapability", "TextContent", "ToolUseContent"), - "MCP sampling types not available -- sampling disabled", - ) + "mcp.types", _SAMPLING_TYPE_NAMES, "MCP sampling types not available -- sampling disabled") _MCP_ELICITATION_TYPES = _import_sdk_names( "mcp.types", ("ElicitRequestParams", "ElicitResult"), - "MCP elicitation types not available -- elicitation disabled", - ) + "MCP elicitation types not available -- elicitation disabled") _MCP_NOTIFICATION_TYPES = _import_sdk_names( - "mcp.types", - ("ServerNotification", "ToolListChangedNotification", - "PromptListChangedNotification", "ResourceListChangedNotification"), - "MCP notification types not available -- dynamic tool discovery disabled", - ) + "mcp.types", _NOTIFICATION_TYPE_NAMES, + "MCP notification types not available -- dynamic tool discovery disabled") else: logger.debug("mcp package not installed -- MCP tool support disabled") @@ -312,25 +246,21 @@ _SDK_HTTPX_MOD = None def sdk_httpx(): """Return the httpx module the *installed* MCP SDK is built against. - 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. + 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. Resolved from the SDK's transport module, not a version number; falls + back to the newest module present. ``None`` only when neither module is importable. """ global _SDK_HTTPX_MOD if _SDK_HTTPX_MOD is not None: return _SDK_HTTPX_MOD try: from mcp.client import streamable_http as _transport - _SDK_HTTPX_MOD = getattr(_transport, "httpx2", None) or getattr( - _transport, "httpx", None - ) + _SDK_HTTPX_MOD = getattr(_transport, "httpx2", None) or getattr(_transport, "httpx", None) except ImportError: _SDK_HTTPX_MOD = None if _SDK_HTTPX_MOD is None: - # Transport module missing / renamed its import: newest present wins. try: import httpx2 as _fallback except ImportError: @@ -343,11 +273,8 @@ def sdk_httpx(): def _client_session_accepts(kwarg: str) -> bool: - """Whether this SDK's ``ClientSession.__init__`` takes ``kwarg``. - - Older SDKs lack ``message_handler`` (no list_changed notifications) and - ``logging_callback`` (server ``notifications/message`` silently dropped). - """ + """Whether this SDK's ``ClientSession.__init__`` takes ``kwarg`` (older SDKs lack + ``message_handler`` and ``logging_callback``).""" if not _MCP_AVAILABLE: return False try: @@ -358,40 +285,32 @@ def _client_session_accepts(kwarg: str) -> bool: # MCP logging levels (RFC 5424 syslog severities) -> Python logging levels. _MCP_LOG_LEVEL_MAP = { - "debug": logging.DEBUG, - "info": logging.INFO, - "notice": logging.INFO, - "warning": logging.WARNING, - "error": logging.ERROR, - "critical": logging.ERROR, - "alert": logging.ERROR, - "emergency": logging.ERROR, + "debug": logging.DEBUG, "info": logging.INFO, "notice": logging.INFO, + "warning": logging.WARNING, "error": logging.ERROR, "critical": logging.ERROR, + "alert": logging.ERROR, "emergency": logging.ERROR, } # --------------------------------------------------------------------------- # Reconnect / keepalive tuning # --------------------------------------------------------------------------- - _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 -# 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 +# 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 _RECYCLED_RECONNECT_TIMEOUT = 15.0 -# 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. +# 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. _STDIO_RESPAWN_WAIT_SEC = 15.0 - -# 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 +# Servers may expire idle sessions on any TTL, so the client MUST ping faster than that TTL; +# short-TTL servers (~15s) need a smaller configured ``keepalive_interval``. The floor stops +# a tiny interval from busy-looping. +_DEFAULT_KEEPALIVE_INTERVAL = 180 +_MIN_KEEPALIVE_INTERVAL = 5 # One bounded cancellation cycle for pending loop tasks at final shutdown, so # cancellation-resistant tasks cannot hang process exit. @@ -401,20 +320,16 @@ _MCP_LOOP_DRAIN_TIMEOUT = 3.0 # _ensure_mcp_sdk() overrides it from mcp.types once the SDK is loaded. _JSONRPC_METHOD_NOT_FOUND = -32601 -# 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. +# 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 ``list_*`` call by following ``nextCursor``. - - 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. - """ + """Drain a paginated ``list_*`` call by following ``nextCursor`` (the SDK fetches one + page per call). ``cache_meta_out`` receives the first page's SEP-2549 hints (``ttl_ms``, + ``cache_scope``). Callers must hold the server's ``_rpc_lock`` for a consistent snapshot.""" items: list = [] cursor = None for _ in range(_MCP_LIST_MAX_PAGES): @@ -443,11 +358,8 @@ async def _paginate_full_list(list_method, items_attr: str, server_name: str, if not isinstance(cursor, str) or not cursor: break else: - logger.warning( - "MCP server '%s': %s pagination exceeded %d pages; " - "truncating at %d items", - server_name, items_attr, _MCP_LIST_MAX_PAGES, len(items), - ) + logger.warning("MCP server '%s': %s pagination exceeded %d pages; truncating at %d items", + server_name, items_attr, _MCP_LIST_MAX_PAGES, len(items)) return items @@ -462,28 +374,19 @@ def _mcp_types(): # --------------------------------------------------------------------------- class MCPServerTask(MCPServerRunMixin, MCPServerTransportMixin, MCPServerHealthMixin): - """One MCP server connection living in one long-lived asyncio Task. - - 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``. - """ + """One MCP server connection living in one long-lived asyncio Task, so the transport's + anyio cancel scopes are entered and exited in the same Task. Run state machine in + ``MCPServerRunMixin``, transport bring-up in ``MCPServerTransportMixin``, keepalive / + refresh / liveness in ``MCPServerHealthMixin``.""" __slots__ = ( - "name", "session", "tool_timeout", - "_task", "_ready", "_shutdown_event", "_reconnect_event", - "_tools", "_error", "_config", - "_sampling", "_elicitation", - "_registered_tool_names", "_auth_type", "_refresh_lock", - "_rpc_lock", "_pending_refresh_tasks", - "_pending_call_context", - "_lifecycle_started_at", "_last_tool_call_at", - "_idle_timeout_seconds", "_max_lifetime_seconds", "_recycled_reason", - "initialize_result", "_ping_unsupported", "_list_cache_meta", - "_reconnect_retries", "_session_proven", "_was_parked", - "_inflight_tasks", "_reconnecting", "_suspect_reason", - "_teardown_race", "_permanent_grace_used", "_stdio_child_pids", + "name", "session", "tool_timeout", "_task", "_ready", "_shutdown_event", "_reconnect_event", + "_tools", "_error", "_config", "_sampling", "_elicitation", "_registered_tool_names", + "_auth_type", "_refresh_lock", "_rpc_lock", "_pending_refresh_tasks", "_pending_call_context", + "_lifecycle_started_at", "_last_tool_call_at", "_idle_timeout_seconds", "_max_lifetime_seconds", + "_recycled_reason", "initialize_result", "_ping_unsupported", "_list_cache_meta", + "_reconnect_retries", "_session_proven", "_was_parked", "_inflight_tasks", "_reconnecting", + "_suspect_reason", "_teardown_race", "_permanent_grace_used", "_stdio_child_pids", "_ever_connected", ) @@ -494,8 +397,8 @@ class MCPServerTask(MCPServerRunMixin, MCPServerTransportMixin, MCPServerHealthM self._task: Optional[asyncio.Task] = None self._ready = asyncio.Event() self._shutdown_event = asyncio.Event() - # When set, _run_http/_run_stdio exit their async-with cleanly and - # run() re-enters the transport (auth recovery, manual refresh, ...). + # 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 @@ -504,48 +407,43 @@ class MCPServerTask(MCPServerRunMixin, MCPServerTransportMixin, MCPServerHealthM self._elicitation: Optional[ElicitationHandler] = None self._registered_tool_names: list[str] = [] self._reconnect_retries: int = 0 - # 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. + # 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 parks. self._session_proven: bool = False - # 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. + # Set once tools were ever registered, never cleared (unlike _ready, which clears every + # reconnect cycle): separates first-connect from reconnect failures in run()'s ladders. self._ever_connected: bool = False - # True from park until the session proves healthy again; logs the - # parked->revived transition exactly once. + # True from park until the session proves healthy again; logs the revival once. self._was_parked: bool = False - # In-flight RPC tasks, 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 fails in-flight calls; lets - # _track_inflight_rpc turn 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 - # Latched by races (teardown-vs-keepalive, auth-lock corruption); - # verified lazily by ensure_healthy() before the next call. + # Latched by races (teardown-vs-keepalive, auth-lock corruption); verified lazily by + # ensure_healthy() before the next call. self._suspect_reason: Optional[str] = None - # A teardown that failed >=1 in-flight call makes the next reconnect a - # RACE RECOVERY: it must not charge the rapid-drop budget. + # 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). + # One-time grace: an auth/permanent-classified failure on a previously PROVEN session + # gets one suspect+reconnect cycle before the park ladder applies. self._permanent_grace_used: bool = False - # Children of the current stdio transport: lets in-flight calls fail - # FAST when the child dies instead of riding out the tool timeout. + # Children of the current stdio transport: 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() - # 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). + # 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 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. + # 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 prompt correctly. self._pending_call_context: Optional[contextvars.Context] = None now = time.monotonic() self._lifecycle_started_at: float = now @@ -553,19 +451,16 @@ class MCPServerTask(MCPServerRunMixin, MCPServerTransportMixin, MCPServerHealthM self._idle_timeout_seconds: Optional[float] = None self._max_lifetime_seconds: Optional[float] = None self._recycled_reason: Optional[str] = None - # InitializeResult from the handshake: the server's REAL advertised - # capabilities, used instead of assuming every ClientSession method - # maps to a supported server method. + # InitializeResult from the handshake: the server's REAL advertised capabilities. self.initialize_result: Optional[Any] = None # SEP-2549 cache hints from the last tools/list (ttl_ms, cache_scope). self._list_cache_meta: dict = {} - # 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. + # Latched when keepalive ``ping`` returns -32601 (optional utility not implemented); + # later keepalives use list_tools instead. Reset on every fresh transport connection. self._ping_unsupported: bool = False - # 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. + # 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") @@ -574,55 +469,48 @@ class MCPServerTask(MCPServerRunMixin, MCPServerTransportMixin, MCPServerHealthM # --------------------------------------------------------------------------- _servers: Dict[str, MCPServerTask] = {} -# Profile registry scope owning each live connection (None outside multiplex): -# a multiplexed /reload-mcp tears down only its own profile's servers. +# 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 startup: servers registered from the on-disk schema cache without -# connecting; 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 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) +# 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)) -# 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. +# 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 _CONNECT_RETRY_MAX_BACKOFF_SEC = 600.0 - -# 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. +# 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. _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 (``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. +# 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). 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]] = {} @@ -644,12 +532,11 @@ def _reset_server_error(server_name: str) -> None: _server_breaker_opened_at.pop(server_name, None) -# 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. +# 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() -# registry tool name -> raw server name, captured at registration. The -# generated name is lossy (punctuation -> ``_``), so never re-parse it. +# 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. @@ -660,11 +547,8 @@ _lock = threading.Lock() def _mcp_registry_scope() -> Optional[str]: - """Registry scope for MCP registrations from the current context. - - Under a profile multiplexer each profile's MCP tools live in its own - registry overlay; single-profile processes stay process-global (None). - """ + """Registry scope for MCP registrations: under a profile multiplexer each profile's MCP + tools live in its own registry overlay; single-profile processes stay global (None).""" from agent.secret_scope import is_multiplex_active if not is_multiplex_active(): @@ -675,34 +559,18 @@ 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 runs on the MCP loop without the discovering profile's context, - so the scope captured at adoption into ``_servers`` is authoritative. - """ + """Scope owning server *name*'s tools: recorded at connect, else current. 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] return _mcp_registry_scope() -# --------------------------------------------------------------------------- -# Cross-process MCP discovery guard: advisory file lock so gateway + CLI + TUI -# don't all run discovery at once. -# --------------------------------------------------------------------------- +# Cross-process MCP discovery guard: advisory file lock so gateway + CLI + TUI don't all +# run discovery at once. _LOCK_UNAVAILABLE: Any = object() # sentinel: locking broken/unavailable _MCP_DISCOVERY_LOCK_PATH: Optional[str] = None # resolved lazily # Bounded wait when another process holds the lock. _MCP_DISCOVERY_LOCK_MAX_RETRIES: int = 240 _MCP_DISCOVERY_LOCK_RETRY_DELAY_S: float = 0.5 - - -# --------------------------------------------------------------------------- -# Connecting, lazy start, discovery -# --------------------------------------------------------------------------- - - -# --------------------------------------------------------------------------- -# Public API -# --------------------------------------------------------------------------- - - diff --git a/tools/mcp_tool_common.py b/tools/mcp_tool_common.py index 0e96d04200..125e3ffaf4 100644 --- a/tools/mcp_tool_common.py +++ b/tools/mcp_tool_common.py @@ -1,6 +1,5 @@ -"""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.""" +"""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 @@ -13,11 +12,10 @@ 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).""" + """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 initialising).""" __slots__ = () @@ -32,11 +30,9 @@ _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).""" + """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.""" value = getattr(obj, snake, _MISSING) if value is _MISSING: value = getattr(obj, camel, _MISSING) @@ -47,9 +43,9 @@ _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 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 @@ -64,8 +60,7 @@ def _resolve_tool_timeout(config: dict) -> float: 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). +# Jitter on reconnect backoff so servers that lost the same backend don't retry in lockstep. _BACKOFF_JITTER = 0.2 # +/-20% @@ -104,8 +99,8 @@ def _sanitize_error(text: str) -> str: 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.""" + """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) @@ -115,9 +110,7 @@ def _prepend_path(env: dict, directory: str) -> dict: updated = dict(env or {}) if not directory: return updated - - existing = updated.get("PATH", "") - parts = [part for part in existing.split(os.pathsep) if part] + parts = [part for part in updated.get("PATH", "").split(os.pathsep) if part] if directory not in parts: parts = [directory, *parts] updated["PATH"] = os.pathsep.join(parts) if parts else directory @@ -125,8 +118,8 @@ def _prepend_path(env: dict, directory: str) -> dict: 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.""" + """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): @@ -157,8 +150,8 @@ def _parse_boolish(value: Any, default: bool = True) -> bool: def _get_lifecycle_seconds(config: dict, key: str) -> Optional[float]: - """Return an optional positive lifecycle timeout from top-level/nested config - (``0`` disables; negatives and non-numbers are warned about and ignored).""" + """Optional positive lifecycle timeout from top-level/nested ``lifecycle`` config (``0`` + disables; negatives and non-numbers are warned about and ignored).""" raw = config.get(key) lifecycle = config.get("lifecycle") if raw is None and isinstance(lifecycle, dict): diff --git a/tools/mcp_tool_discovery.py b/tools/mcp_tool_discovery.py index 38aa0be6d1..7b8f298937 100644 --- a/tools/mcp_tool_discovery.py +++ b/tools/mcp_tool_discovery.py @@ -1,8 +1,7 @@ """Connecting and discovery for tools.mcp_tool: per-server connect cooldown, connect / -lazy-start / recycled-stdio wake-up, ``register_mcp_servers`` / ``discover_mcp_tools`` -and the status / probe public API. Split from tools/mcp_tool.py; origin state -(``_servers``, ``_lock``, the loop, patchable helpers) is read through ``_core`` so -``mock.patch("tools.mcp_tool.X")`` keeps working.""" +lazy-start / recycled-stdio wake-up, ``register_mcp_servers`` / ``discover_mcp_tools`` and +the status / probe public API. Origin state (``_servers``, ``_lock``, the loop, patchable +helpers) is read through ``_core`` so ``mock.patch("tools.mcp_tool.X")`` keeps working.""" from __future__ import annotations @@ -19,10 +18,7 @@ def _record_connect_failure(server_name: str) -> None: """Stamp a geometric, capped retry cooldown after a failed connect (under ``_lock``).""" n = _core._server_connect_failures.get(server_name, 0) + 1 _core._server_connect_failures[server_name] = n - backoff = min( - _core._CONNECT_RETRY_BASE_BACKOFF_SEC * (2 ** (n - 1)), - _core._CONNECT_RETRY_MAX_BACKOFF_SEC, - ) + backoff = min(_core._CONNECT_RETRY_BASE_BACKOFF_SEC * (2 ** (n - 1)), _core._CONNECT_RETRY_MAX_BACKOFF_SEC) _core._server_connect_retry_after[server_name] = time.monotonic() + backoff @@ -33,42 +29,40 @@ def _clear_connect_failure(server_name: str) -> None: def _connect_cooldown_active(server_name: str) -> bool: - """Return True if ``server_name`` is still within its retry cooldown.""" + """True if ``server_name`` is still within its retry cooldown.""" deadline = _core._server_connect_retry_after.get(server_name) return deadline is not None and time.monotonic() < deadline -async def _connect_server(name: str, config: dict) -> _core.MCPServerTask: - """Create an MCPServerTask, start it and return once ready. +def _enabled(cfg: dict) -> bool: + return _core._parse_boolish(cfg.get("enabled", True), default=True) - Tear it down with ``server.shutdown()`` on the same loop. Raises on bad - config, missing HTTP support, or connect/initialize failure. - """ + +async def _connect_server(name: str, config: dict) -> _core.MCPServerTask: + """Create an MCPServerTask, start it and return once ready. Tear it down with + ``server.shutdown()`` on the same loop. Raises on bad config, missing HTTP support, + or connect/initialize failure.""" server = _core.MCPServerTask(name) claim = _core._connect_server_claim.get() claim_token = None if claim is not None: claim(server) - # 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. + # 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 = _core._connect_server_claim.set(None) try: await server.start(config) except asyncio.CancelledError: - # start() already reaps server._task; a shutdown() here could - # swallow the cancellation. + # start() already reaps server._task; a shutdown() here could swallow the cancellation. raise except BaseException: - # Discovery owns claimed tasks (recoverable park vs terminal failure); - # standalone probes have no revival owner and must reap 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() except Exception as shutdown_exc: # noqa: BLE001 -- best-effort reap, don't mask the real error - logger.debug( - "MCP server '%s' shutdown during orphan-reap failed: %s", - name, shutdown_exc, - ) + logger.debug("MCP server '%s' shutdown during orphan-reap failed: %s", name, shutdown_exc) raise finally: if claim_token is not None: @@ -80,7 +74,6 @@ def _request_lazy_reconnect(server_name: str, server: _core.MCPServerTask) -> bo """Wake a recycled stdio server and wait briefly for a fresh session.""" if not server._is_recycled_stdio(): return False - loop = _core._running_loop() if loop is None: return False @@ -102,10 +95,7 @@ def _request_lazy_reconnect(server_name: str, server: _core.MCPServerTask) -> bo try: return bool(_core._run_on_mcp_loop(_await_ready, timeout=_core._RECYCLED_RECONNECT_TIMEOUT)) except Exception as exc: - logger.warning( - "MCP server '%s': lazy reconnect after stdio recycle failed: %s", - server_name, exc, - ) + logger.warning("MCP server '%s': lazy reconnect after stdio recycle failed: %s", server_name, exc) return False @@ -135,20 +125,17 @@ def _note_connect_success(name: str) -> None: def _ensure_lazy_server_connected(server_name: str) -> bool: """Connect a lazily-registered server on demand (sync; blocks the caller). - 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. + 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 _core._lock: server = _core._servers.get(server_name) if server is not None and server.session is not None: return True config = _core._lazy_server_configs.get(server_name) - if not config: - return False - if _core._connect_cooldown_active(server_name): - return False - if server_name in _core._server_connecting: + if (not config or _core._connect_cooldown_active(server_name) + or server_name in _core._server_connecting): return False _core._server_connecting.add(server_name) _core._server_connect_errors.pop(server_name, None) @@ -164,9 +151,7 @@ def _ensure_lazy_server_connected(server_name: str) -> bool: _core._run_on_mcp_loop(_connect, timeout=float(connect_timeout) + 30.0) except BaseException as exc: message = _note_connect_failure(server_name, exc) - logger.warning( - "Lazy MCP connect failed for '%s': %s", server_name, message, - ) + logger.warning("Lazy MCP connect failed for '%s': %s", server_name, message) return False _note_connect_success(server_name) @@ -175,11 +160,8 @@ def _ensure_lazy_server_connected(server_name: str) -> bool: stale_fingerprint = _core._lazy_server_fingerprints.pop(server_name, None) cached_names = _core._lazy_server_tool_names.pop(server_name, None) or [] server = _core._servers.get(server_name) - live_names = set( - getattr(server, "_registered_tool_names", []) or [] - ) - # The cached manifest may advertise tools the live server no longer - # serves; deregister those phantoms. + live_names = set(getattr(server, "_registered_tool_names", []) or []) + # The cached manifest may advertise tools the live server no longer serves. phantom_names = [n for n in cached_names if n not in live_names] if phantom_names: from tools.registry import registry @@ -188,17 +170,15 @@ def _ensure_lazy_server_connected(server_name: str) -> bool: registry.deregister(tool_name, scope=_core._server_registry_scope(server_name)) _core._forget_mcp_tool_server(tool_name) logger.info( - "MCP server '%s': deregistered %d phantom cached tool(s) not " - "served live (stale schema-cache fingerprint %s): %s", - server_name, len(phantom_names), stale_fingerprint, - ", ".join(phantom_names), - ) + "MCP server '%s': deregistered %d phantom cached tool(s) not served live (stale " + "schema-cache fingerprint %s): %s", + server_name, len(phantom_names), stale_fingerprint, ", ".join(phantom_names)) return server is not None and server.session is not None def _get_connected_server_for_call(server_name: str) -> Optional[_core.MCPServerTask]: - """Return a connected server; the single first-use connect point for lazy - servers and the wake-up point for recycled stdio ones.""" + """Return a connected server; the single first-use connect point for lazy servers and + the wake-up point for recycled stdio ones.""" with _core._lock: server = _core._servers.get(server_name) is_lazy = server_name in _core._lazy_server_configs @@ -217,36 +197,20 @@ def _get_connected_server_for_call(server_name: str) -> Optional[_core.MCPServer async def _discover_and_register_server(name: str, config: dict) -> List[str]: """Connect one server, register its tools; return the registered names.""" connect_timeout = config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT) - # The claim callback runs inside _connect_server while this frame is - # suspended; a list append avoids a nonlocal rebind. + # The claim callback runs inside _connect_server while this frame is suspended; a list + # append avoids a nonlocal rebind. claimed: List[_core.MCPServerTask] = [] - - def _claim_server(created: _core.MCPServerTask) -> None: - claimed.append(created) - - claim_token = _core._connect_server_claim.set(_claim_server) + claim_token = _core._connect_server_claim.set(claimed.append) try: - server = await asyncio.wait_for( - _core._connect_server(name, config), - timeout=connect_timeout, - ) + server = await asyncio.wait_for(_core._connect_server(name, config), timeout=connect_timeout) except BaseException: server = claimed[0] if claimed else None task = server._task if server is not None else None - task_cancelling = ( - task.cancelling() - if task is not None and hasattr(task, "cancelling") - else 0 - ) - if ( - server is not None - and server._error is not None - and task is not None - and not task.done() - and not task_cancelling - ): - # Recoverable park: the run task stays alive to self-probe, so - # adopt it for shutdown/revival. + task_cancelling = task.cancelling() if task is not None and hasattr(task, "cancelling") else 0 + if (server is not None and server._error is not None and task is not None + and not task.done() and not task_cancelling): + # Recoverable park: the run task stays alive to self-probe, so adopt it for + # shutdown/revival. with _core._lock: _core._servers[name] = server _core._server_scope_keys[name] = _core._mcp_registry_scope() @@ -264,40 +228,28 @@ async def _discover_and_register_server(name: str, config: dict) -> List[str]: registered_names = _core._register_server_tools(name, server, config) server._registered_tool_names = list(registered_names) - - transport_type = "HTTP" if "url" in config else "stdio" - logger.info( - "MCP server '%s' (%s): registered %d tool(s): %s", - name, transport_type, len(registered_names), - ", ".join(registered_names), - ) + logger.info("MCP server '%s' (%s): registered %d tool(s): %s", name, + "HTTP" if "url" in config else "stdio", len(registered_names), ", ".join(registered_names)) return registered_names def _select_new_servers(servers: Dict[str, dict]) -> Dict[str, dict]: """Pick connect candidates and refresh per-server bookkeeping (under ``_lock``). - Candidates: enabled, not connected, not connecting (dedups concurrent - discovery entry points), not lazily registered, not in backoff. Known - servers without a live session are parked or mid-reconnect; their tools are - deregistered so nothing else can nudge them — signal a reconnect here. + Candidates: enabled, not connected, not connecting (dedups concurrent discovery entry + points), not lazily registered, not in backoff. Known servers without a live session are + parked or mid-reconnect; their tools are deregistered so nothing else can nudge them — + signal a reconnect here. """ with _core._lock: connecting = set(_core._server_connecting) new_servers = { - k: v - for k, v in servers.items() - if k not in _core._servers - and k not in connecting - and k not in _core._lazy_server_configs - and _core._parse_boolish(v.get("enabled", True), default=True) - and not _core._connect_cooldown_active(k) + k: v for k, v in servers.items() + if k not in _core._servers and k not in connecting and k not in _core._lazy_server_configs + and _enabled(v) and not _core._connect_cooldown_active(k) } - stale_cached = [ - _core._servers[k] - for k in servers - if k in _core._servers and getattr(_core._servers[k], "session", None) is None - ] + stale_cached = [_core._servers[k] for k in servers + if k in _core._servers and getattr(_core._servers[k], "session", None) is None] _core._server_connecting.update(new_servers) for srv_name in new_servers: _core._server_connect_errors.pop(srv_name, None) @@ -314,10 +266,9 @@ def _select_new_servers(servers: Dict[str, dict]) -> Dict[str, dict]: def _register_lazy_from_cache(new_servers: Dict[str, dict]) -> Tuple[Dict[str, dict], int, int]: - """Register ``lazy: true`` servers with a valid schema-cache entry without - connecting; a missing/stale entry (or a failed registration) falls back to - eager. Returns (servers still needing an eager connect, lazy tool count, - lazy server count).""" + """Register ``lazy: true`` servers with a valid schema-cache entry without connecting; a + missing/stale entry (or a failed registration) falls back to eager. Returns (servers still + needing an eager connect, lazy tool count, lazy server count).""" eager_servers: Dict[str, dict] = dict(new_servers) lazy_registered = 0 lazy_server_count = 0 @@ -336,9 +287,7 @@ def _register_lazy_from_cache(new_servers: Dict[str, dict]) -> Tuple[Dict[str, d try: names = _core._register_from_cache_sync(name, cfg, entry) except Exception as exc: - logger.warning( - "Failed lazy MCP registration for '%s': %s", name, exc, - ) + logger.warning("Failed lazy MCP registration for '%s': %s", name, exc) with _core._lock: _core._server_connecting.add(name) continue @@ -350,30 +299,24 @@ def _register_lazy_from_cache(new_servers: Dict[str, dict]) -> Tuple[Dict[str, d async def _discover_all(new_servers: Dict[str, dict]) -> None: """Connect every candidate concurrently; record per-server outcome.""" - server_names = list(new_servers.keys()) results = await asyncio.gather( *(_core._discover_and_register_server(name, cfg) for name, cfg in new_servers.items()), - return_exceptions=True, - ) - for name, result in zip(server_names, results): + return_exceptions=True) + for name, result in zip(new_servers, results): if isinstance(result, BaseException): command = new_servers.get(name, {}).get("command") message = _note_connect_failure(name, result) - logger.warning( - "Failed to connect to MCP server '%s'%s: %s", - name, - f" (command={command})" if command else "", - message, - ) + logger.warning("Failed to connect to MCP server '%s'%s: %s", + name, f" (command={command})" if command else "", message) else: _note_connect_success(name) def _run_discovery_pass(new_servers: Dict[str, dict]) -> None: - """Run ``_discover_all`` on the MCP loop with the interrupt flag parked and - the ``_server_connecting`` set cleaned up when the pass dies early.""" - # Clear a stale interrupt flag (executor threads are reused) so a prior - # session's interrupt cannot cancel this discovery pass. + """Run ``_discover_all`` on the MCP loop with the interrupt flag parked and the + ``_server_connecting`` set cleaned up when the pass dies early.""" + # 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: @@ -381,22 +324,16 @@ def _run_discovery_pass(new_servers: Dict[str, dict]) -> None: try: _core._run_on_mcp_loop(lambda: _discover_all(new_servers), timeout=120) except (TimeoutError, InterruptedError) as _e: - # Entries stranded in _server_connecting would block future - # reconnect attempts. + # Entries stranded in _server_connecting would block future reconnect attempts. how = "timed out" if isinstance(_e, TimeoutError) else "interrupted" with _core._lock: stale = [n for n in new_servers if n in _core._server_connecting] if stale: - logger.warning( - "MCP discovery %s while %d server(s) were still " - "connecting; clearing stale connecting set: %s", - how, len(stale), ", ".join(stale), - ) + logger.warning("MCP discovery %s while %d server(s) were still connecting; " + "clearing stale connecting set: %s", how, len(stale), ", ".join(stale)) _core._server_connecting.difference_update(stale) for _sn in stale: - _core._server_connect_errors.setdefault( - _sn, f"Connection attempt {how} during discovery", - ) + _core._server_connect_errors.setdefault(_sn, f"Connection attempt {how} during discovery") raise finally: if _was_interrupted: @@ -404,29 +341,29 @@ def _run_discovery_pass(new_servers: Dict[str, dict]) -> None: def _connected_summary(names, *, lazy_tools: int = 0, lazy_servers: int = 0) -> Tuple[int, int, int]: - """(tool count, connected server count, failed count) for a set of - candidate names, folding in lazily registered servers.""" + """(tool count, connected server count, failed count) for a set of candidate names, + folding in lazily registered servers.""" with _core._lock: - connected = [ - n - for n in names - if n in _core._servers and n not in _core._server_connect_errors - ] - tool_count = sum( - len(getattr(_core._servers[n], "_registered_tool_names", [])) - for n in connected - ) + connected = [n for n in names if n in _core._servers and n not in _core._server_connect_errors] + tool_count = sum(len(getattr(_core._servers[n], "_registered_tool_names", [])) for n in connected) failed = len(names) - len(connected) return tool_count + lazy_tools, len(connected) + lazy_servers, failed -def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: - """Connect the given ``{name: config}`` servers and register their tools. +def _log_summary(prefix: str, names, **lazy) -> None: + """Log `` N tool(s) from M server(s) (K failed)`` when anything happened.""" + new_tool_count, connected_count, failed = _connected_summary(names, **lazy) + if new_tool_count or failed: + summary = f"{prefix} {new_tool_count} tool(s) from {connected_count} server(s)" + if failed: + summary += f" ({failed} failed)" + logger.info(summary) - Idempotent for connected names; ``enabled: false`` servers are skipped - without disconnecting existing sessions. Returns every registered MCP - tool name. - """ + +def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: + """Connect the given ``{name: config}`` servers and register their tools. Idempotent for + connected names; ``enabled: false`` servers are skipped without disconnecting existing + sessions. Returns every registered MCP tool name.""" if not _core._ensure_mcp_sdk(): logger.debug("MCP SDK not available -- skipping explicit MCP registration") return [] @@ -443,39 +380,24 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]: new_servers, lazy_registered, lazy_server_count = _register_lazy_from_cache(new_servers) if not new_servers: if lazy_registered: - logger.info( - "MCP: registered %d lazy tool(s) from schema cache " - "(no processes spawned)", - lazy_registered, - ) + logger.info("MCP: registered %d lazy tool(s) from schema cache (no processes spawned)", + lazy_registered) return _core._existing_tool_names() _core._ensure_mcp_loop() _run_discovery_pass(new_servers) - - new_tool_count, connected_count, failed = _connected_summary( - new_servers, lazy_tools=lazy_registered, lazy_servers=lazy_server_count, - ) - if new_tool_count or failed: - summary = f"MCP: registered {new_tool_count} tool(s) from {connected_count} server(s)" - if failed: - summary += f" ({failed} failed)" - logger.info(summary) - + _log_summary("MCP: registered", new_servers, lazy_tools=lazy_registered, lazy_servers=lazy_server_count) return _core._existing_tool_names() def _acquire_discovery_lock_with_retry(): - """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). Returns the cookie (None / _LOCK_UNAVAILABLE when - 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). Returns the + cookie (None / _LOCK_UNAVAILABLE when unguarded).""" cookie = _core._try_acquire_mcp_discovery_lock() if cookie is not None: return cookie - logger.debug( - "Another process holds MCP discovery lock -- retrying with backoff" - ) + logger.debug("Another process holds MCP discovery lock -- retrying with backoff") for _ in range(_core._MCP_DISCOVERY_LOCK_MAX_RETRIES): time.sleep(_core._MCP_DISCOVERY_LOCK_RETRY_DELAY_S) cookie = _core._try_acquire_mcp_discovery_lock() @@ -483,22 +405,17 @@ def _acquire_discovery_lock_with_retry(): break if cookie is None: - logger.warning( - "MCP discovery lock still held after %d retries -- " - "running discovery unguarded", - _core._MCP_DISCOVERY_LOCK_MAX_RETRIES, - ) + logger.warning("MCP discovery lock still held after %d retries -- running discovery unguarded", + _core._MCP_DISCOVERY_LOCK_MAX_RETRIES) elif cookie is not _core._LOCK_UNAVAILABLE: logger.debug("Retry succeeded -- acquired MCP discovery lock") return cookie def discover_mcp_tools() -> List[str]: - """Entry point: load config, connect servers, register tools. - - Safe without the ``mcp`` package (returns []). Idempotent: only servers - missing from a previous call are retried. Returns all MCP tool names. - """ + """Entry point: load config, connect servers, register tools. Safe without the ``mcp`` + package (returns []). Idempotent: only servers missing from a previous call are retried. + Returns all MCP tool names.""" servers = _core._load_mcp_config() if not servers: logger.debug("No MCP servers configured") @@ -513,38 +430,21 @@ def discover_mcp_tools() -> List[str]: try: with _core._lock: connecting = set(_core._server_connecting) - new_server_names = [ - name - for name, cfg in servers.items() - if name not in _core._servers - and name not in connecting - and _core._parse_boolish(cfg.get("enabled", True), default=True) - ] + new_server_names = [name for name, cfg in servers.items() + if name not in _core._servers and name not in connecting and _enabled(cfg)] tool_names = _core.register_mcp_servers(servers) - if not new_server_names: - return tool_names - - new_tool_count, connected_count, failed_count = _connected_summary(new_server_names) - if new_tool_count or failed_count: - summary = f" MCP: {new_tool_count} tool(s) from {connected_count} server(s)" - if failed_count: - summary += f" ({failed_count} failed)" - logger.info(summary) - + if new_server_names: + _log_summary(" MCP:", new_server_names) return tool_names - finally: if cookie not in (None, _core._LOCK_UNAVAILABLE): cookie.release() def is_mcp_tool_parallel_safe(tool_name: str) -> bool: - """True when the tool's server opted into ``supports_parallel_tool_calls``. - - Uses the provenance captured at registration, never the (ambiguous) - ``mcp__{server}__{tool}`` string shape. - """ + """True when the tool's server opted into ``supports_parallel_tool_calls``. Uses the + provenance captured at registration, never the (ambiguous) ``mcp__{server}__{tool}`` shape.""" if not tool_name.startswith(_core.MCP_TOOL_NAME_PREFIX): return False with _core._lock: @@ -555,9 +455,9 @@ def is_mcp_tool_parallel_safe(tool_name: str) -> bool: def get_mcp_status() -> List[dict]: """Status of every configured server for banner/TUI display. - 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. + 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. """ configured = _core._load_mcp_config() if not configured: @@ -569,28 +469,18 @@ def get_mcp_status() -> List[dict]: connect_errors = dict(_core._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, - } + 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 = _core._parse_boolish(cfg.get("enabled", True), default=True) + enabled = _enabled(cfg) server = active_servers.get(name) if server and server.session is not None: entry = _entry(name, transport, "connected", connected=True) - entry["tools"] = ( - len(server._registered_tool_names) - if hasattr(server, "_registered_tool_names") - else len(server._tools) - ) + 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) elif not enabled: @@ -607,8 +497,8 @@ def get_mcp_status() -> List[dict]: def probe_mcp_server_tools() -> Dict[str, List[tuple]]: - """Connect to each enabled server, list ``(tool_name, description)`` and - disconnect, without registering anything. Failed servers are omitted.""" + """Connect to each enabled server, list ``(tool_name, description)`` and disconnect, + without registering anything. Failed servers are omitted.""" if not _core._ensure_mcp_sdk(): return {} @@ -616,10 +506,7 @@ def probe_mcp_server_tools() -> Dict[str, List[tuple]]: if not servers_config: return {} - enabled = { - k: v for k, v in servers_config.items() - if _core._parse_boolish(v.get("enabled", True), default=True) - } + enabled = {k: v for k, v in servers_config.items() if _enabled(v)} if not enabled: return {} @@ -629,29 +516,17 @@ def probe_mcp_server_tools() -> Dict[str, List[tuple]]: probed_servers: List[_core.MCPServerTask] = [] async def _probe_all(): - names = list(enabled.keys()) - coros = [] - for name, cfg in enabled.items(): - ct = cfg.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT) - coros.append(asyncio.wait_for(_core._connect_server(name, cfg), timeout=ct)) - + coros = [asyncio.wait_for(_core._connect_server(name, cfg), + timeout=cfg.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT)) + for name, cfg in enabled.items()] outcomes = await asyncio.gather(*coros, return_exceptions=True) - - for name, outcome in zip(names, outcomes): + for name, outcome in zip(enabled, outcomes): if isinstance(outcome, Exception): logger.debug("Probe: failed to connect to '%s': %s", name, outcome) continue probed_servers.append(outcome) - tools = [] - for t in outcome._tools: - desc = getattr(t, "description", "") or "" - tools.append((t.name, desc)) - result[name] = tools - - await asyncio.gather( - *(s.shutdown() for s in probed_servers), - return_exceptions=True, - ) + result[name] = [(t.name, getattr(t, "description", "") or "") for t in outcome._tools] + await asyncio.gather(*(s.shutdown() for s in probed_servers), return_exceptions=True) try: _core._run_on_mcp_loop(_probe_all, timeout=120) @@ -664,17 +539,15 @@ def probe_mcp_server_tools() -> Dict[str, List[tuple]]: def has_registered_mcp_tools() -> bool: - """True if any MCP server has registered tools (cheap; no registry walk). - - Checks registered TOOLS, not connected servers, so the per-turn refresh - hook stays idle for zero-tool servers. - """ + """True if any MCP server has registered tools (cheap; no registry walk). Checks + registered TOOLS, not connected servers, so the per-turn refresh hook stays idle for + zero-tool servers.""" with _core._lock: return bool(_core._mcp_tool_server_names) def get_registered_mcp_server_names() -> set: - """Server names that registered at least one tool (the live, filtered - signal — not merely what config.yaml lists).""" + """Server names that registered at least one tool (the live, filtered signal — not + merely what config.yaml lists).""" with _core._lock: return set(_core._mcp_tool_server_names.values()) diff --git a/tools/mcp_tool_handlers.py b/tools/mcp_tool_handlers.py index 3beebc7ddb..338433469b 100644 --- a/tools/mcp_tool_handlers.py +++ b/tools/mcp_tool_handlers.py @@ -1,4 +1,6 @@ -"""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.""" +"""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.""" import logging import asyncio @@ -12,7 +14,10 @@ from typing import Any, Callable, 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_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") @@ -21,8 +26,8 @@ logger = logging.getLogger("tools.mcp_tool") # --------------------------------------------------------------- pre-call gates def _trust_gate_check(server_name: str, tool_name: str) -> Optional[str]: - """Approval gate for write-capable tools on ``trust: untrusted`` servers. - None to proceed, else a ``tool_error``. Fail-closed: approval-system errors block.""" + """Approval gate for write-capable tools on ``trust: untrusted`` servers. None to proceed, + else a ``tool_error``. Fail-closed: approval-system errors block.""" trust = _core._server_trust_levels.get(server_name, _core._TRUST_FULL) if trust != _core._TRUST_UNTRUSTED or _core._tool_read_only_hints.get(server_name, {}).get(tool_name) is True: return None @@ -35,8 +40,7 @@ def _trust_gate_check(server_name: str, tool_name: str) -> Optional[str]: f"tool is write-capable (no readOnlyHint=true annotation) and may modify external 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}", - ) + 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}' was blocked: the approval " @@ -67,10 +71,9 @@ def _check_circuit_breaker(server_name: str) -> Optional[str]: def _acquire_call_server(server_name: str, tool_timeout: float): """``(server, None)`` when a call may be dispatched, else ``(None, error)``. No session: a reconnect may be completing (fresh session swaps in asynchronously), so wait - briefly before charging a breaker strike. Still down → reconnecting or parked (e.g. dead + briefly before charging a breaker strike. 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 and return a clean "reconnecting" error — the breaker resets once the - fresh session initializes.""" + task to rebuild and return a clean "reconnecting" error.""" not_connected = tool_error(f"MCP server '{server_name}' is not connected") server = _core._get_connected_server_for_call(server_name) if not server: @@ -110,6 +113,11 @@ def _strike(server_name: str, message: str, **extra) -> str: return tool_error(message, **extra) +def _mcp_loop_running() -> bool: + loop = _core._mcp_loop + return loop is not None and loop.is_running() + + def _lookup_reconnectable_server(server_name: str, require_loop: bool = False): """The registered server object when it can be signalled to reconnect, else None. With *require_loop*, also None unless the MCP loop is running (nothing to wait on).""" @@ -120,11 +128,6 @@ def _lookup_reconnectable_server(server_name: str, require_loop: bool = False): return srv -def _mcp_loop_running() -> bool: - loop = _core._mcp_loop - return loop is not None and loop.is_running() - - def _retry_once(server_name: str, retry_call, op_description: str, what: str): """Re-run ``retry_call`` after a recovery step. Returns the result (closing the breaker) when it is not an error payload; None when the retry raised or errored (caller falls through).""" @@ -198,9 +201,8 @@ def _handle_session_expired_and_retry(server_name: str, exc: BaseException, retr 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).""" + """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).""" def _handle_stdio_child_exited_and_retry(server_name: str, exc: Exception, retry_call, op_description: str): @@ -249,17 +251,12 @@ def _handle_stdio_child_exited_and_retry(server_name: str, exc: Exception, retry f"{type(retry_exc).__name__}: {_exc_str(retry_exc)}")) -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 _invoke_with_recovery(server_name: str, call_once: Callable[[], str], op: str, recoverers, on_final_failure: Callable[[BaseException], None], record_outcome: bool = False) -> str: """Run ``call_once``, walking the recovery ladder on failure. Each recoverer ``(server_name, exc, retry_call, op) -> Optional[str]`` returns None when the exception is - not its kind; order matters: dead stdio child → auth → session expiry. Unrecovered + not its kind; order matters: dead stdio child -> auth -> session expiry. Unrecovered exceptions go through ``on_final_failure`` (breaker strike / logging) and become the generic call-failed error. ``record_outcome`` applies breaker bookkeeping to the FIRST attempt only; retries own their bookkeeping inside the recoverers.""" @@ -267,7 +264,7 @@ def _invoke_with_recovery(server_name: str, call_once: Callable[[], str], op: st result = call_once() return _record_call_outcome(server_name, result) if record_outcome else result except InterruptedError: - return _interrupted_call_result() + return tool_error("MCP call interrupted: user sent a new message") except Exception as exc: for recover in recoverers: recovered = recover(server_name, exc, call_once, op) @@ -386,7 +383,7 @@ def _capped_structured_content(result): def _render_call_tool_result(result, server_name: str) -> str: - """Pure: ``CallToolResult`` → the handler's JSON string. ``content`` is the primary + """Pure: ``CallToolResult`` -> the handler's JSON string. ``content`` is the primary (model-oriented) payload; ``structuredContent`` supplements it (or becomes ``result`` when there is no text). Server-level ``_meta`` is surfaced minus protocol-reserved keys. ``.is_error`` is ``.isError`` before mcp 2.0.""" @@ -440,7 +437,6 @@ def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float): 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() @@ -453,18 +449,17 @@ def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float): return _invoke_with_recovery( server_name, lambda: _core._run_on_mcp_loop(_call, timeout=tool_timeout), op, (_handle_stdio_child_exited_and_retry, _handle_auth_error_and_retry, _handle_session_expired_and_retry), - _on_failure, record_outcome=True, - ) + _on_failure, record_outcome=True) return _handler def _make_utility_handler(server_name: str, tool_timeout: float, op: str, log_label: str, rpc, render, required: Optional[str] = None): - """Shared shape of the four utility handlers (resources/prompts): ``rpc(session, args)`` - is awaited under ``_rpc_lock``, ``render(result, server_name)`` builds the JSON-able payload, - ``required`` names a parameter validated before any transport work. The wrapper owns the - connected check and the auth / session-expired recovery ladder.""" + """Shared shape of the four utility handlers (resources/prompts): ``rpc(session, args, + server_name)`` is awaited under ``_rpc_lock``, ``render(result, server_name)`` builds the + JSON-able payload, ``required`` names a parameter validated before any transport work. 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) @@ -476,14 +471,13 @@ def _make_utility_handler(server_name: str, tool_timeout: float, op: str, log_la async def _call(): _mark_server_call_started(server) async with server._rpc_lock: - result = await rpc(server.session, args) + result = await rpc(server.session, args, server_name) return json.dumps(render(result, server_name), ensure_ascii=False) return _invoke_with_recovery( server_name, lambda: _core._run_on_mcp_loop(_call, timeout=tool_timeout), op, (_handle_auth_error_and_retry, _handle_session_expired_and_retry), - lambda exc: logger.error("MCP %s/%s failed: %s", server_name, log_label, exc), - ) + lambda exc: logger.error("MCP %s/%s failed: %s", server_name, log_label, exc)) return _handler @@ -536,8 +530,7 @@ def _render_prompt_list(all_prompts, server_name: str) -> dict: if getattr(p, "arguments", None): entry["arguments"] = [ {"name": a.name, **_pick(a, ("description", "description", True), ("required", "required"))} - for a in p.arguments - ] + for a in p.arguments] prompts.append(entry) return {"prompts": prompts} @@ -556,31 +549,28 @@ def _render_get_prompt(result, server_name: str) -> dict: return resp -def _make_list_resources_handler(server_name: str, tool_timeout: float): - """Sync handler that lists resources from an MCP server.""" - return _make_utility_handler(server_name, tool_timeout, "resources/list", "list_resources", - lambda session, args: _core._paginate_full_list(session.list_resources, "resources", server_name), - _render_resource_list) +def _utility_factory(op: str, log_label: str, rpc, render, required: Optional[str] = None): + """``(server_name, tool_timeout) -> sync handler`` for one utility tool.""" + def _factory(server_name: str, tool_timeout: float): + return _make_utility_handler(server_name, tool_timeout, op, log_label, rpc, render, required) + return _factory -def _make_read_resource_handler(server_name: str, tool_timeout: float): - """Sync handler that reads a resource by URI from an MCP server.""" - return _make_utility_handler(server_name, tool_timeout, "resources/read", "read_resource", - lambda session, args: session.read_resource(args["uri"]), _render_read_resource, required="uri") - - -def _make_list_prompts_handler(server_name: str, tool_timeout: float): - """Sync handler that lists prompts from an MCP server.""" - return _make_utility_handler(server_name, tool_timeout, "prompts/list", "list_prompts", - lambda session, args: _core._paginate_full_list(session.list_prompts, "prompts", server_name), - _render_prompt_list) - - -def _make_get_prompt_handler(server_name: str, tool_timeout: float): - """Sync handler that gets a prompt by name from an MCP server.""" - return _make_utility_handler(server_name, tool_timeout, "prompts/get", "get_prompt", - lambda session, args: session.get_prompt(args["name"], arguments=args.get("arguments", {})), - _render_get_prompt, required="name") +_make_list_resources_handler = _utility_factory( + "resources/list", "list_resources", + lambda session, args, sn: _core._paginate_full_list(session.list_resources, "resources", sn), + _render_resource_list) +_make_read_resource_handler = _utility_factory( + "resources/read", "read_resource", + lambda session, args, sn: session.read_resource(args["uri"]), _render_read_resource, required="uri") +_make_list_prompts_handler = _utility_factory( + "prompts/list", "list_prompts", + lambda session, args, sn: _core._paginate_full_list(session.list_prompts, "prompts", sn), + _render_prompt_list) +_make_get_prompt_handler = _utility_factory( + "prompts/get", "get_prompt", + lambda session, args, sn: session.get_prompt(args["name"], arguments=args.get("arguments", {})), + _render_get_prompt, required="name") def _make_check_fn(server_name: str): diff --git a/tools/mcp_tool_registration.py b/tools/mcp_tool_registration.py index 06e7096a12..1acf01d469 100644 --- a/tools/mcp_tool_registration.py +++ b/tools/mcp_tool_registration.py @@ -1,17 +1,21 @@ -"""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. - -Both entry points (``_register_server_tools`` for a live server, -``_register_from_cache_sync`` for a lazy cached manifest) build ``_Candidate`` +"""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. Both entry points +(``_register_server_tools`` live, ``_register_from_cache_sync`` lazy) build ``_Candidate`` records and feed the single ``_register_candidates`` loop.""" import logging from dataclasses import dataclass from typing import TYPE_CHECKING, Any, Callable, Dict, Iterable, List, Optional from tools.mcp_tool_common import _parse_boolish, _core, _resolve_tool_timeout -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_schema import _UTILITY_CAPABILITY_ATTRS, _UTILITY_CAPABILITY_METHODS, _build_utility_schemas, _normalize_name_filter, matches_name_filter +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_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 @@ -21,30 +25,27 @@ logger = logging.getLogger("tools.mcp_tool") _UTILITY_ORIGIN_PREFIX = "generated utility " # 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, + "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 _normalize_server_trust(value: Any) -> str: - """Config ``trust`` -> tier. None -> ``full`` (backward-compatible default); - an unrecognized string -> ``untrusted`` so a misspelled tier fails closed.""" + """Config ``trust`` -> tier. None -> ``full`` (backward-compatible default); an + unrecognized string -> ``untrusted`` so a misspelled tier fails closed.""" if value is None: return _core._TRUST_FULL text = str(value).strip().lower() if text in (_core._TRUST_FULL, _core._TRUST_UNTRUSTED): return text logger.warning( - "MCP trust: unrecognized trust value %r — treating as 'untrusted' (valid values: full, untrusted)", value, - ) + "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``; unknown metadata means write-capable.""" + """True only when annotations (SDK object or schema-cache dict) carry ``readOnlyHint is + True``; unknown metadata means write-capable.""" annotations = getattr(mcp_tool, "annotations", None) if isinstance(annotations, dict): return annotations.get("readOnlyHint") is True @@ -52,9 +53,9 @@ def _annotation_read_only_hint(mcp_tool: Any) -> bool: def _record_tool_trust_metadata(server_name: str, config: dict, tools: List[Any]) -> None: - """Capture per-server trust and per-tool readOnlyHint at discovery — the - security boundary: the call-time gate classifies from data we control, - never re-read server-supplied state.""" + """Capture per-server trust and per-tool readOnlyHint at discovery — the security + boundary: the call-time gate classifies from data we control, never re-read + server-supplied state.""" 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, {}) @@ -77,14 +78,12 @@ def _forget_mcp_tool_server(tool_name: str) -> None: def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: dict) -> List[dict]: - """Utility schemas allowed by config (``tools.resources``/``tools.prompts``) - and by the server's advertised capabilities. - - ``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 — ClientSession - defines all four methods). When no initialize_result was captured (test - fixtures, older paths) fall back to that legacy session-method check.""" + """Utility schemas allowed by config (``tools.resources``/``tools.prompts``) and by the + server's advertised capabilities. ``initialize_result.capabilities`` is the source of + truth: its sub-objects are non-None iff the server advertises that request family (a + ``hasattr(server.session, ...)`` gate never filters anything — ClientSession defines all + four methods). Without an initialize_result (test fixtures, older paths) fall back to + that legacy session-method check.""" tools_filter = config.get("tools") or {} enabled = {f: _parse_boolish(tools_filter.get(f), default=True) for f in ("resources", "prompts")} init_result = getattr(server, "initialize_result", None) @@ -112,8 +111,8 @@ def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: d def _existing_tool_names() -> List[str]: - """Tool names for all currently connected servers plus lazy (cache-registered) - servers, whose tools live only in the registry.""" + """Tool names for all currently connected servers plus lazy (cache-registered) servers, + whose tools live only in the registry.""" names: List[str] = [] for _sname, server in _core._servers.items(): if hasattr(server, "_registered_tool_names"): @@ -121,17 +120,15 @@ def _existing_tool_names() -> List[str]: else: names.extend(_core._convert_mcp_schema(server.name, t)["name"] for t in server._tools) with _core._lock: - names.extend( - 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(n for sname, tool_names in _core._lazy_server_tool_names.items() + if sname not in _core._servers for n in tool_names) return names def _make_tool_filter(name: str, config: dict) -> Callable[[str], bool]: - """Include/exclude predicate for a server's tool names: ``tools.include`` is a - whitelist (``[]`` = register nothing), ``tools.exclude`` a blacklist; entries - are exact names or fnmatch globs; include wins over exclude.""" + """Include/exclude predicate for a server's tool names: ``tools.include`` is a whitelist + (``[]`` = register nothing), ``tools.exclude`` a blacklist; entries are exact names or + fnmatch globs; include wins over exclude.""" 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") @@ -147,8 +144,8 @@ def _make_tool_filter(name: str, config: dict) -> Callable[[str], bool]: class _CachedMCPTool: - """Stand-in for MCP Tool objects loaded from the schema cache. Missing or - non-dict ``annotations`` (older cache files) fail closed to write-capable.""" + """Stand-in for MCP Tool objects loaded from the schema cache. Missing or non-dict + ``annotations`` (older cache files) fail closed to write-capable.""" __slots__ = ("name", "description", "inputSchema", "annotations") @@ -165,17 +162,15 @@ class _CachedMCPTool: for raw in raws: if isinstance(raw, dict) and raw.get("name"): schema = raw.get("inputSchema") - out.append(cls( - raw["name"], raw.get("description") or "", - schema if isinstance(schema, dict) else {}, raw.get("annotations"), - )) + out.append(cls(raw["name"], raw.get("description") or "", + schema if isinstance(schema, dict) else {}, raw.get("annotations"))) return out @dataclass class _Candidate: - """One registration attempt: a native tool or a generated utility. - ``origin`` is the provenance text used in collision diagnostics.""" + """One registration attempt: a native tool or a generated utility. ``origin`` is the + provenance text used in collision diagnostics.""" registry_name: str origin: str @@ -187,11 +182,10 @@ class _Candidate: return self.origin.startswith(_UTILITY_ORIGIN_PREFIX) -def _tool_candidates( - name: str, tools: Iterable[Any], should_register: Callable[[str], bool], tool_timeout, -) -> List[_Candidate]: - """Native tools (live SDK objects or ``_CachedMCPTool``) -> candidates. The - injection scan runs on BOTH paths: the cache file is user-writable JSON.""" +def _tool_candidates(name: str, tools: Iterable[Any], should_register: Callable[[str], bool], + tool_timeout) -> List[_Candidate]: + """Native tools (live SDK objects or ``_CachedMCPTool``) -> candidates. The injection scan + runs on BOTH paths: the cache file is user-writable JSON.""" out: List[_Candidate] = [] for t in tools: if not should_register(t.name): @@ -218,21 +212,18 @@ def _utility_candidates(name: str, entries: Iterable[Any], tool_timeout) -> List def _resolve_name_collisions(name: str, candidates: List[_Candidate]) -> List[_Candidate]: - """Preflight registry-name collisions among one server's candidates. - - Exact duplicates (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). Returns the survivors in order.""" + """Preflight registry-name collisions among one server's candidates. Exact duplicates + (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). Returns the + survivors in order.""" unique: List[_Candidate] = [] seen: set[tuple[str, str]] = set() origins_by_name: Dict[str, set[str]] = {} for c in candidates: if (c.registry_name, c.origin) in seen: - logger.debug( - "MCP server '%s': duplicate registration candidate %s for '%s'; keeping one", - name, c.origin, c.registry_name, - ) + logger.debug("MCP server '%s': duplicate registration candidate %s for '%s'; keeping one", + name, c.origin, c.registry_name) continue seen.add((c.registry_name, c.origin)) unique.append(c) @@ -251,16 +242,14 @@ def _resolve_name_collisions(name: str, candidates: List[_Candidate]) -> List[_C "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], - ) + name, ", ".join(utility_origins), native_origins[0]) continue ambiguous[registry_name] = sorted(origins) for registry_name, origins in sorted(ambiguous.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), - ) + name, registry_name, ", ".join(origins)) return [c for c in unique if c.registry_name not in ambiguous and (c.registry_name, c.origin) not in shadowed] @@ -268,30 +257,25 @@ def _log_foreign_owner(name: str, c: _Candidate, existing_toolset: str, lazy: bo """Diagnostics for a name already owned by another toolset (skipped to preserve the owner).""" if lazy: if not c.is_utility: - logger.warning( - "MCP server '%s' (lazy): cached tool '%s' collides with toolset '%s' — skipping", - name, c.registry_name, existing_toolset, - ) + logger.warning("MCP server '%s' (lazy): cached tool '%s' collides with toolset '%s' — skipping", + name, c.registry_name, existing_toolset) return if existing_toolset.startswith("mcp-"): log, fmt = logger.error, ( "MCP server '%s': %s normalizes to '%s', already owned by MCP toolset '%s' " - "— skipping to preserve the existing owner" - ) + "— skipping to preserve the existing owner") else: log, fmt = logger.warning, ( - "MCP server '%s': %s (→ '%s') collides with built-in tool in toolset '%s' — skipping to preserve built-in" - ) + "MCP server '%s': %s (→ '%s') collides with built-in tool in toolset '%s' — skipping to preserve built-in") log(fmt, name, c.origin, c.registry_name, existing_toolset) -def _register_candidates( - name: str, candidates: List[_Candidate], *, check_fn: Callable, scope: Callable[[], Optional[str]], lazy: bool, -) -> List[str]: - """Register candidates under toolset ``mcp-{name}``; returns the names that - landed. The ownership pre-check is advisory only — servers connect in - parallel, so ``ToolRegistry.register()`` is the atomic ownership gate and - its verdict is re-read after every call.""" +def _register_candidates(name: str, candidates: List[_Candidate], *, check_fn: Callable, + scope: Callable[[], Optional[str]], lazy: bool) -> List[str]: + """Register candidates under toolset ``mcp-{name}``; returns the names that landed. The + ownership pre-check is advisory only — servers connect in parallel, so + ``ToolRegistry.register()`` is the atomic ownership gate and its verdict is re-read after + every call.""" from tools.registry import registry toolset_name = f"mcp-{name}" @@ -303,15 +287,11 @@ def _register_candidates( continue registry.register( name=c.registry_name, toolset=toolset_name, schema=c.schema, handler=c.handler, check_fn=check_fn, - is_async=False, description=c.schema.get("description") or "", scope=scope(), - ) + is_async=False, description=c.schema.get("description") or "", scope=scope()) if registry.get_toolset_for_tool(c.registry_name) != toolset_name: if not lazy: - logger.error( - "MCP server '%s': registration of %s as '%s' was rejected by the registry; " - "skipping provenance/count updates", - name, c.origin, c.registry_name, - ) + logger.error("MCP server '%s': registration of %s as '%s' was rejected by the registry; " + "skipping provenance/count updates", name, c.origin, c.registry_name) continue _core._track_mcp_tool_server(c.registry_name, name) registered.append(c.registry_name) @@ -321,8 +301,8 @@ def _register_candidates( 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.""" + """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 @@ -337,46 +317,38 @@ def _write_schema_cache(name: str, server: "MCPServerTask", config: dict, should # Persisted so the lazy path trust-gates identically next startup. "annotations": {"readOnlyHint": _annotation_read_only_hint(t)}, }) - utility_payload = [ - {"schema": e["schema"], "handler_key": e["handler_key"]} for e in _select_utility_schemas(name, server, config) - ] + utility_payload = [{"schema": e["schema"], "handler_key": e["handler_key"]} + for e 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"), - ) + 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); used by - initial discovery and list_changed refresh. Returns the registered names. - - 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. Generated utilities share - the namespace and join the same preflight.""" + """Register an already-connected server's tools (plus utility tools); used by initial + discovery and list_changed refresh. Returns the registered names. 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.""" should_register = _make_tool_filter(name, config) _record_tool_trust_metadata(name, config, server._tools) candidates = _tool_candidates(name, server._tools, should_register, server.tool_timeout) candidates += _utility_candidates(name, _select_utility_schemas(name, server, config), server.tool_timeout) registered = _register_candidates( name, _resolve_name_collisions(name, candidates), - check_fn=_make_check_fn(name), scope=lambda: _core._server_registry_scope(name), lazy=False, - ) + check_fn=_make_check_fn(name), scope=lambda: _core._server_registry_scope(name), lazy=False) if registered: _write_schema_cache(name, server, config, should_register) return registered 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 goes through - ``_get_connected_server_for_call`` -> ``_ensure_lazy_server_connected``. - Trust metadata is recorded first so the call-time gate is identical whether - the server was spawned live or registered from cache.""" + """Lazy startup: register a server's tools from a cached manifest with no child process; + the first real call goes through ``_get_connected_server_for_call`` -> + ``_ensure_lazy_server_connected``. Trust metadata is recorded first so the call-time gate + is identical whether the server was spawned live or registered from cache.""" from tools.mcp_schema_cache import config_fingerprint, tools_from_cache_entry, utility_tools_from_cache_entry tool_timeout = _resolve_tool_timeout(config) @@ -385,8 +357,7 @@ def _register_from_cache_sync(name: str, config: dict, entry: dict) -> List[str] candidates = _tool_candidates(name, cached_tools, _make_tool_filter(name, config), tool_timeout) candidates += _utility_candidates(name, utility_tools_from_cache_entry(entry), tool_timeout) registered = _register_candidates( - name, candidates, check_fn=_make_check_fn(name), scope=_core._mcp_registry_scope, lazy=True, - ) + name, candidates, check_fn=_make_check_fn(name), scope=_core._mcp_registry_scope, lazy=True) if registered: with _core._lock: _core._lazy_server_configs[name] = dict(config) diff --git a/tools/mcp_tool_schema.py b/tools/mcp_tool_schema.py index c654861636..47d6409c65 100644 --- a/tools/mcp_tool_schema.py +++ b/tools/mcp_tool_schema.py @@ -1,6 +1,6 @@ """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.""" +compatibility, mcp__server__tool naming, utility-tool schemas, include/exclude filters and +description injection scanning.""" import logging import fnmatch @@ -11,8 +11,8 @@ 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. +# 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(pattern, re.I), reason) for pattern, reason in ( @@ -31,16 +31,14 @@ _MCP_INJECTION_PATTERNS = [ 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.""" + """Scan a tool description for injection patterns; returns finding strings (empty = + clean) and logs a warning when any match.""" if not description: return [] findings = [reason for pattern, reason in _MCP_INJECTION_PATTERNS if pattern.search(description)] if findings: - logger.warning( - "MCP server '%s' tool '%s': suspicious description content — %s. Description: %.200s", - server_name, tool_name, "; ".join(findings), description, - ) + logger.warning("MCP server '%s' tool '%s': suspicious description content — %s. Description: %.200s", + server_name, tool_name, "; ".join(findings), description) return findings @@ -48,11 +46,11 @@ _EMPTY_OBJECT_SCHEMA = {"type": "object", "properties": {}} def _rewrite_local_refs(node): - """Promote legacy ``definitions`` to ``$defs`` (Moonshot rejects the draft-07 - form) — ONLY where it is a JSON Schema meta-keyword, never as a property NAME - inside ``properties``/``patternProperties``: a parameter legitimately named - ``definitions`` rewritten to ``$defs`` would 400 the whole tool array - (Anthropic/OpenAI forbid ``$`` in property names).""" + """Promote legacy ``definitions`` to ``$defs`` (Moonshot rejects the draft-07 form) — ONLY + where it is a JSON Schema meta-keyword, never as a property NAME inside + ``properties``/``patternProperties``: a parameter legitimately named ``definitions`` + rewritten to ``$defs`` would 400 the whole tool array (Anthropic/OpenAI forbid ``$`` in + property names).""" if isinstance(node, list): return [_rewrite_local_refs(item) for item in node] if not isinstance(node, dict): @@ -70,9 +68,9 @@ def _rewrite_local_refs(node): def _repair_object_shape(node): - """Recursively fill a missing object ``type``, ensure ``properties`` (so - ``required`` can't dangle) and prune ``required`` to names present in - ``properties`` (Gemini 400s otherwise).""" + """Recursively fill a missing object ``type``, ensure ``properties`` (so ``required`` + can't dangle) and prune ``required`` to names present in ``properties`` (Gemini 400s + otherwise).""" if isinstance(node, list): return [_repair_object_shape(item) for item in node] if not isinstance(node, dict): @@ -96,13 +94,12 @@ def _repair_object_shape(node): def _normalize_mcp_input_schema(schema: dict | None) -> dict: - """Normalize MCP input schemas so one form is valid on OpenAI, Anthropic, - Gemini and Moonshot. Order matters: ``definitions`` -> ``$defs``; nullable - ``anyOf`` unions collapsed to the non-null branch (Anthropic rejects nullable - branches; optionality lives in the parent's ``required``; the ``nullable: - true`` hint is kept so runtime coercion can map a model-emitted ``"null"`` - string to ``None``); same-typed const unions -> enum (AFTER the nullable - strip); then object-shape repair.""" + """Normalize MCP input schemas so one form is valid on OpenAI, Anthropic, Gemini and + Moonshot. Order matters: ``definitions`` -> ``$defs``; nullable ``anyOf`` unions collapsed + to the non-null branch (Anthropic rejects nullable branches; optionality lives in the + parent's ``required``; the ``nullable: true`` hint is kept so runtime coercion can map a + model-emitted ``"null"`` string to ``None``); same-typed const unions -> enum (AFTER the + nullable strip); then object-shape repair.""" if not schema: return dict(_EMPTY_OBJECT_SCHEMA) from tools.schema_sanitizer import collapse_const_unions, strip_nullable_unions @@ -120,14 +117,14 @@ def _normalize_mcp_input_schema(schema: dict | None) -> dict: 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.""" + """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____``: 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 = "__" @@ -139,8 +136,8 @@ def mcp_prefixed_tool_name(server_name: str, tool_name: str) -> str: 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.""" + """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}"), @@ -148,9 +145,9 @@ def _convert_mcp_schema(server_name: str, mcp_tool) -> dict: } -# Utility tools generated per server: handler_key -> (description template, -# parameter properties, required names). Schemas are FROZEN wire bytes — the -# key order emitted by ``_build_utility_schemas`` must not change. +# Utility tools generated per server: handler_key -> (description template, parameter +# properties, required names). Schemas are FROZEN wire bytes — the key order emitted by +# ``_build_utility_schemas`` must not change. _UTILITY_TOOL_SPECS = ( ("list_resources", "List available resources from MCP server '{server}'", {}, None), ("read_resource", "Read a resource by URI from MCP server '{server}'", @@ -200,9 +197,9 @@ def _normalize_name_filter(value: Any, label: str) -> set[str]: 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).""" + """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: @@ -210,17 +207,15 @@ def matches_name_filter(tool_name: str, patterns: set[str]) -> bool: return any(fnmatch.fnmatchcase(tool_name, p) for p in patterns if "*" in p or "?" in p or "[" in p) -# Utility handler -> ClientSession method it needs (legacy gate when no -# initialize_result was captured). +# Utility handler -> ClientSession method it needs (legacy gate when no initialize_result +# was captured). _UTILITY_CAPABILITY_METHODS = {key: key for key, *_ in _UTILITY_TOOL_SPECS} -# 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 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", + "list_resources": "resources", "read_resource": "resources", + "list_prompts": "prompts", "get_prompt": "prompts", } diff --git a/tools/mcp_tool_server_run.py b/tools/mcp_tool_server_run.py index 08b6a173c1..91f4f9ddc3 100644 --- a/tools/mcp_tool_server_run.py +++ b/tools/mcp_tool_server_run.py @@ -1,7 +1,7 @@ """Lifecycle of :class:`tools.mcp_tool.MCPServerTask`: the long-lived ``run`` state machine (connect -> serve -> reconnect/park/recycle), keepalive-driven lifecycle waits, start/shutdown -and tool deregistration. Split from tools/mcp_tool.py; origin state and patchable helpers are -read through ``_core`` so ``mock.patch("tools.mcp_tool.X")`` keeps working.""" +and tool deregistration. Origin state and patchable helpers are read through ``_core`` so +``mock.patch("tools.mcp_tool.X")`` keeps working.""" import asyncio import logging @@ -15,8 +15,8 @@ logger = logging.getLogger("tools.mcp_tool") @dataclass class _RetryBudget: - """Per-run() retry counters shared by the branch helpers (``_reconnect_retries`` - lives on the task because handlers and tests read it).""" + """Per-run() retry counters shared by the branch helpers (``_reconnect_retries`` lives on + the task because handlers and tests read it).""" initial_retries: int = 0 backoff: float = 1.0 @@ -46,21 +46,16 @@ class MCPServerRunMixin: async def _wait_for_lifecycle_event(self) -> str: """Serve the connection until a lifecycle event; return its kind. - ``"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. - - 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. + ``"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. 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. """ keepalive_interval = max( _core._MIN_KEEPALIVE_INTERVAL, - float(self._config.get("keepalive_interval", _core._DEFAULT_KEEPALIVE_INTERVAL)), - ) + float(self._config.get("keepalive_interval", _core._DEFAULT_KEEPALIVE_INTERVAL))) shutdown_task = asyncio.create_task(self._shutdown_event.wait()) reconnect_task = asyncio.create_task(self._reconnect_event.wait()) @@ -75,22 +70,17 @@ class MCPServerRunMixin: timeout = max(0.0, min(timeout, recycle_deadline - time.monotonic())) done, _pending = await asyncio.wait( - {shutdown_task, reconnect_task}, - timeout=timeout, - return_when=asyncio.FIRST_COMPLETED, - ) + {shutdown_task, reconnect_task}, timeout=timeout, return_when=asyncio.FIRST_COMPLETED) if done: break if self._recycle_if_due(): return "recycle" - # 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). + # 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 - ): + if self._rpc_lock.locked() or any(not t.done() for t in self._inflight_tasks): continue try: async with self._rpc_lock: @@ -100,11 +90,8 @@ class MCPServerRunMixin: logger.warning( "MCP server '%s' keepalive failed, triggering " "reconnect (state: connected → degraded): %s: %s", - self.name, type(root).__name__, root, - ) - self.mark_suspect( - f"keepalive failed: {type(root).__name__}: {root}" - ) + self.name, type(root).__name__, root) + self.mark_suspect(f"keepalive failed: {type(root).__name__}: {root}") self._reconnect_event.set() break # Survived a full keepalive interval: real proof of health. @@ -115,29 +102,20 @@ class MCPServerRunMixin: if self._shutdown_event.is_set(): self._fail_inflight_calls("shutdown") return "shutdown" - # Deliberate teardown: fail in-flight RPCs NOW rather than letting - # them ride the dying transport to the full tool timeout. + # 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" - async def _wait_for_reconnect_or_shutdown( - self, timeout: Optional[float] = None - ) -> str: - """Wait, while parked, for a reconnect request or shutdown. - - Returns ``"shutdown"`` or ``"reconnect"`` (explicit request or, with - ``timeout``, the periodic self-probe); the reconnect event is cleared - first. Shutdown wins a tie. - """ + async def _wait_for_reconnect_or_shutdown(self, timeout: Optional[float] = None) -> str: + """Wait, while parked, for a reconnect request or shutdown. 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()) try: - await asyncio.wait( - {shutdown_task, reconnect_task}, - return_when=asyncio.FIRST_COMPLETED, - timeout=timeout, - ) + await asyncio.wait({shutdown_task, reconnect_task}, return_when=asyncio.FIRST_COMPLETED, timeout=timeout) finally: await self._cancel_waiters(shutdown_task, reconnect_task) if self._shutdown_event.is_set(): @@ -148,36 +126,28 @@ class MCPServerRunMixin: async def _park(self, revival_reason: str) -> bool: """Drop this server's tools and wait for a reconnect request. - 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. + 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. True when shutdown was + requested instead. """ self._was_parked = True self._deregister_tools() self._reconnect_event.clear() - parked = await self._wait_for_reconnect_or_shutdown( - timeout=_core._PARKED_RETRY_INTERVAL - ) - if parked == "shutdown": + if await self._wait_for_reconnect_or_shutdown(timeout=_core._PARKED_RETRY_INTERVAL) == "shutdown": return True - logger.debug( - "MCP server '%s': attempting revival %s (self-probe or explicit " - "reconnect request); rebuilding transport.", - self.name, revival_reason, - ) + logger.debug("MCP server '%s': attempting revival %s (self-probe or explicit " + "reconnect request); rebuilding transport.", self.name, revival_reason) return False async def _prepare_run(self, config: dict) -> bool: """Bind config, build sampling/elicitation handlers, validate HTTP. - 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. + 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 = _core._resolve_tool_timeout(config) @@ -189,49 +159,34 @@ class MCPServerRunMixin: _core._ensure_mcp_sdk() sampling_config = config.get("sampling", {}) - if sampling_config.get("enabled", True) and _core._MCP_SAMPLING_TYPES: - self._sampling = _core.SamplingHandler(self.name, sampling_config) - else: - self._sampling = None - - # elicitation/create lets a server ask for structured input mid-call; - # the handler routes it through Hermes' approval system. + self._sampling = (_core.SamplingHandler(self.name, sampling_config) + if sampling_config.get("enabled", True) and _core._MCP_SAMPLING_TYPES else None) + # 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 _core._MCP_ELICITATION_TYPES: - self._elicitation = _core.ElicitationHandler(self.name, elicitation_config, owner=self) - else: - self._elicitation = None + self._elicitation = (_core.ElicitationHandler(self.name, elicitation_config, owner=self) + if elicitation_config.get("enabled", True) and _core._MCP_ELICITATION_TYPES else None) if "url" in config and "command" in config: - logger.warning( - "MCP server '%s' has both 'url' and 'command' in config. " - "Using HTTP transport ('url'). Remove 'command' to silence " - "this warning.", - self.name, - ) + logger.warning("MCP server '%s' has both 'url' and 'command' in config. " + "Using HTTP transport ('url'). Remove 'command' to silence " + "this warning.", self.name) if not self._is_http(): return True try: _core._validate_remote_mcp_url(self.name, config.get("url")) - # 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" - ): + # 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"): await self._preflight_content_type( - config["url"], - headers=dict(config.get("headers") or {}), + config["url"], headers=dict(config.get("headers") or {}), ssl_verify=config.get("ssl_verify", True), - client_cert=_core._resolve_client_cert(self.name, config), - ) + client_cert=_core._resolve_client_cert(self.name, config)) except (_core.InvalidMcpUrlError, _core.NonMcpEndpointError) as exc: # Fail fast and non-retryably: publish the error to start(). logger.warning("%s", exc) @@ -243,12 +198,11 @@ class MCPServerRunMixin: 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. The branch helpers return True to - keep looping and False to exit the loop. + 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. The branch helpers return True to keep + looping and False to exit the loop. """ if not await self._prepare_run(config): return @@ -265,8 +219,8 @@ class MCPServerRunMixin: if not await self._on_clean_return(lifecycle_reason, budget): break except asyncio.CancelledError: - # Not a connection failure: re-raise so cancellation reaches - # asyncio and shutdown()'s ``await self._task`` completes. + # 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: @@ -279,37 +233,26 @@ class MCPServerRunMixin: self._stdio_child_pids = set() async def _on_clean_return(self, lifecycle_reason: str, budget: "_RetryBudget") -> bool: - """Transport returned cleanly: shutdown, stdio recycle, or a requested - rebuild (auth recovery / manual refresh / keepalive failure). A rebuild - is not a failure for the retry counters.""" + """Transport returned cleanly: 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(): return False if lifecycle_reason == "recycle": - logger.info( - "MCP server '%s': stdio session recycled after %s; " - "waiting for lazy reconnect", - self.name, self._recycled_reason, - ) + logger.info("MCP server '%s': stdio session recycled after %s; " + "waiting for lazy reconnect", self.name, self._recycled_reason) self.session = None # Dormant until a lazy call wakes it (untimed: nothing to self-probe). return await self._wait_for_reconnect_or_shutdown() != "shutdown" # Per-cycle chatter stays DEBUG; WARNINGs mark state transitions. - logger.debug( - "MCP server '%s': reconnecting (OAuth recovery or " - "manual refresh)", - self.name, - ) - # 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. + logger.debug("MCP server '%s': reconnecting (OAuth recovery or manual refresh)", self.name) + # 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: - logger.info( - "MCP server '%s': reconnect after teardown race " - "(in-flight calls were failed); not charging the " - "rapid-drop budget", - self.name, - ) + logger.info("MCP server '%s': reconnect after teardown race " + "(in-flight calls were failed); not charging the " + "rapid-drop budget", self.name) self._teardown_race = False budget.backoff = 1.0 elif self._session_proven: @@ -323,30 +266,27 @@ class MCPServerRunMixin: "without a healthy session (rapid-drop budget " "exhausted), parking; will self-probe every %ds " "until it recovers (state: degraded → parked)", - self.name, _core._MAX_RECONNECT_RETRIES, - _core._PARKED_RETRY_INTERVAL, - ) + self.name, _core._MAX_RECONNECT_RETRIES, _core._PARKED_RETRY_INTERVAL) if not await self._park_and_rearm("from parked state", budget): return False - # Clear readiness too: a stale _ready lets handler-side - # recovery mistake the old session for a fresh one. + # Clear readiness too: a stale _ready lets handler-side recovery mistake the old + # session for a fresh one. self._ready.clear() self.session = None return True async def _park_and_rearm(self, revival_reason: str, budget: "_RetryBudget") -> bool: - """Park; on revival leave a budget of ONE probe per wake so a still-dead - server parks again instead of burning 5 rapid retries. False on shutdown.""" + """Park; on revival leave a budget of ONE probe per wake so a still-dead server parks + again instead of burning 5 rapid retries. False on shutdown.""" if await self._park(revival_reason): return False self._reconnect_retries = _core._MAX_RECONNECT_RETRIES budget.backoff = 1.0 return True - async def _park_initial_failure(self, exc: Exception, revival_reason: str, - budget: "_RetryBudget") -> bool: - """Publish ``exc`` to the waiting ``start()``, park, and on revival reset - every counter so the ladder starts fresh. False on shutdown.""" + async def _park_initial_failure(self, exc: Exception, revival_reason: str, budget: "_RetryBudget") -> bool: + """Publish ``exc`` to the waiting ``start()``, park, and on revival reset every counter + so the ladder starts fresh. False on shutdown.""" self._error = exc self._ready.set() if await self._park(revival_reason): @@ -363,32 +303,26 @@ class MCPServerRunMixin: budget.backoff = min(budget.backoff * 2, _core._MAX_BACKOFF_SECONDS) async def _on_transport_error(self, exc: Exception, budget: "_RetryBudget") -> bool: - """Transport raised: classify, then run the initial-connect or the - reconnect ladder. Returns False when the run loop must exit.""" - # Unwrap anyio TaskGroup wrappers: the group's str() is useless - # and hides the root cause from the classification below. + """Transport raised: classify, then run the initial-connect or the reconnect ladder. + Returns False when the run loop must exit.""" + # Unwrap anyio TaskGroup wrappers: the group's str() is useless and hides the root + # cause from the classification below. root = _core._unwrap_exception_group(exc) failure_class = _core._classify_mcp_failure(root) if self._is_recycled_stdio(): - logger.warning( - "MCP server '%s': lazy reconnect after stdio recycle " - "failed, marking unavailable while retrying: %s: %s", - self.name, type(root).__name__, root, - ) + logger.warning("MCP server '%s': lazy reconnect after stdio recycle " + "failed, marking unavailable while retrying: %s: %s", + self.name, type(root).__name__, root) self._recycled_reason = None - # 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). + # 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: return await self._on_initial_connect_error(exc, root, failure_class, budget) - # If shutdown was requested, don't reconnect if self._shutdown_event.is_set(): - logger.debug( - "MCP server '%s' disconnected during shutdown: %s: %s", - self.name, type(root).__name__, root, - ) + logger.debug("MCP server '%s' disconnected during shutdown: %s: %s", + self.name, type(root).__name__, root) return False if failure_class == "permanent": @@ -400,45 +334,36 @@ class MCPServerRunMixin: "MCP server '%s' failed after %d reconnection attempts, " "parking; will self-probe every %ds until it recovers " "(state: degraded → parked): %s: %s", - self.name, _core._MAX_RECONNECT_RETRIES, - _core._PARKED_RETRY_INTERVAL, - type(root).__name__, root, - ) + self.name, _core._MAX_RECONNECT_RETRIES, _core._PARKED_RETRY_INTERVAL, + type(root).__name__, root) return await self._park_and_rearm("from parked state", budget) - logger.debug( - "MCP server '%s' connection lost (attempt %d/%d), " - "reconnecting in %.0fs: %s: %s", - self.name, self._reconnect_retries, _core._MAX_RECONNECT_RETRIES, - budget.backoff, type(root).__name__, root, - ) + logger.debug("MCP server '%s' connection lost (attempt %d/%d), " + "reconnecting in %.0fs: %s: %s", + self.name, self._reconnect_retries, _core._MAX_RECONNECT_RETRIES, + budget.backoff, type(root).__name__, root) await self._backoff_sleep(budget) - # Check again after sleeping return not self._shutdown_event.is_set() async def _on_initial_connect_error(self, exc: Exception, root: BaseException, failure_class: str, budget: "_RetryBudget") -> bool: if failure_class == "permanent": - # Deterministic failure (bad command, non-MCP URL, - # 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. + # Deterministic failure (bad command, non-MCP URL, 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 _core._is_auth_error(root): logger.warning( "MCP server '%s' failed initial authentication, " "parking until credentials change; re-authenticate " "with `hermes mcp login %s` " "(state: connecting → parked): %s: %s", - self.name, self.name, - type(root).__name__, root, - ) + self.name, self.name, type(root).__name__, root) else: logger.warning( "MCP server '%s' failed initial connection with a " "permanent error, parking without retries " "(state: connecting → parked): %s: %s", - self.name, type(root).__name__, root, - ) + self.name, type(root).__name__, root) return await self._park_initial_failure(exc, "after permanent initial failure", budget) budget.initial_retries += 1 @@ -447,20 +372,15 @@ class MCPServerRunMixin: "MCP server '%s' failed initial connection after " "%d attempts, parking until a reconnect is " "requested (state: connecting → parked): %s: %s", - self.name, _core._MAX_INITIAL_CONNECT_RETRIES, - type(root).__name__, root, - ) + self.name, _core._MAX_INITIAL_CONNECT_RETRIES, type(root).__name__, root) return await self._park_initial_failure(exc, "after initial connection failures", budget) logger.debug( "MCP server '%s' initial connection failed " "(attempt %d/%d), retrying in %.0fs: %s: %s", - self.name, budget.initial_retries, - _core._MAX_INITIAL_CONNECT_RETRIES, budget.backoff, - type(root).__name__, root, - ) + self.name, budget.initial_retries, _core._MAX_INITIAL_CONNECT_RETRIES, budget.backoff, + type(root).__name__, root) await self._backoff_sleep(budget) - # Check if shutdown was requested during the sleep if self._shutdown_event.is_set(): self._error = exc self._ready.set() @@ -468,25 +388,17 @@ class MCPServerRunMixin: return True async def _on_permanent_error(self, root: BaseException, budget: "_RetryBudget") -> bool: - # 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 ( - _core._is_auth_error(root) - and self._session_proven - and not self._permanent_grace_used - ): + # 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 _core._is_auth_error(root) and self._session_proven and not self._permanent_grace_used: self._permanent_grace_used = True - self.mark_suspect( - f"auth error on proven session: {root}" - ) + self.mark_suspect(f"auth error on proven session: {root}") logger.warning( "MCP server '%s': auth error on a previously " "healthy session — marking suspect and forcing " "one reconnect instead of parking (state: " "connected → suspect): %s: %s", - self.name, type(root).__name__, root, - ) + self.name, type(root).__name__, root) self._reconnect_retries = 0 budget.backoff = 1.0 await asyncio.sleep(_core._jittered(1.0)) @@ -496,9 +408,7 @@ class MCPServerRunMixin: "MCP server '%s' hit a permanent error, parking " "without retries; will self-probe every %ds " "(state: connected → parked): %s: %s", - self.name, _core._PARKED_RETRY_INTERVAL, - type(root).__name__, root, - ) + self.name, _core._PARKED_RETRY_INTERVAL, type(root).__name__, root) return await self._park_and_rearm("from parked state (permanent error)", budget) async def start(self, config: dict): @@ -507,13 +417,11 @@ class MCPServerRunMixin: try: await self._ready.wait() except asyncio.CancelledError: - # The caller's connect timeout (discover_mcp_tools wraps start() - # in asyncio.wait_for) cancels *this* coroutine, but the - # ensure_future'd run() task is independent and would otherwise - # keep running detached — parked on a hung transport with no - # owner to reap it (#59349). Propagate the cancellation so the - # transport context managers unwind and their finally blocks - # release the child process / FDs. + # The caller's connect timeout (discover_mcp_tools wraps start() in + # asyncio.wait_for) cancels *this* coroutine, but the ensure_future'd run() task + # is independent and would otherwise keep running detached — parked on a hung + # transport with no owner to reap it. Propagate the cancellation so the transport + # context managers unwind and release the child process / FDs. if self._task and not self._task.done(): self._task.cancel() raise @@ -523,20 +431,15 @@ class MCPServerRunMixin: async def shutdown(self): """Signal the Task to exit and wait for clean resource teardown.""" self._shutdown_event.set() - # Defensive: if _wait_for_lifecycle_event is blocking, we need ANY - # event to unblock it. _shutdown_event alone is sufficient (the - # helper checks shutdown first), but setting reconnect too ensures - # there's no race where the helper misses the shutdown flag after - # returning "reconnect". + # _shutdown_event alone unblocks _wait_for_lifecycle_event (it checks shutdown first), + # but setting reconnect too closes any race where the helper misses the shutdown flag + # after returning "reconnect". self._reconnect_event.set() if self._task and not self._task.done(): try: await asyncio.wait_for(self._task, timeout=10) except asyncio.TimeoutError: - logger.warning( - "MCP server '%s' shutdown timed out, cancelling task", - self.name, - ) + logger.warning("MCP server '%s' shutdown timed out, cancelling task", self.name) self._task.cancel() try: await self._task @@ -551,14 +454,9 @@ class MCPServerRunMixin: self.session = None def _deregister_tools(self) -> None: - """Drop this server's tools from the global registry (idempotent). - - Pulls the server's tool schemas out of the registry so the agent - stops advertising them to the model. Called on shutdown AND when the - reconnect budget is exhausted, so a dead server never leaves phantom - tool definitions bloating the prompt cache and producing "not - connected" errors on every turn. - """ + """Drop this server's tools from the global registry (idempotent). Called on shutdown + AND when the reconnect budget is exhausted, so a dead server never leaves phantom tool + definitions bloating the prompt cache and producing "not connected" errors.""" from tools.registry import registry for tool_name in list(getattr(self, "_registered_tool_names", [])):