refactor(tools): compact MCP facade/discovery/handlers/run/registration/schema/cache modules

This commit is contained in:
Teknium
2026-09-02 22:28:12 -07:00
parent 113f04616b
commit 5e4f30bc9d
8 changed files with 653 additions and 1081 deletions

View File

@@ -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

View File

@@ -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.<name>.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.<name>.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
# ---------------------------------------------------------------------------

View File

@@ -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.<name>.timeout``
> ``timeouts.mcp.tool_call`` > the 300s default; values are platform-clamped
by ``resolve_timeout``."""
"""Per-server tool-call timeout. Precedence: ``mcp_servers.<name>.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):

View File

@@ -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 ``<prefix> 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())

View File

@@ -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):

View File

@@ -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)

View File

@@ -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__<server>__<tool>``: 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__<server>__<tool>``: 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",
}

View File

@@ -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", [])):