refactor(tools): compact MCP facade/discovery/handlers/run/registration/schema/cache modules
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
|
||||
@@ -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", [])):
|
||||
|
||||
Reference in New Issue
Block a user