refactor(tools): MCP discovery/run/handlers defensive collapse, shared helpers, blank squeeze
This commit is contained in:
@@ -21,7 +21,6 @@ _cache_lock = threading.Lock()
|
||||
|
||||
def _cache_path() -> Path:
|
||||
from hermes_constants import get_hermes_home
|
||||
|
||||
return get_hermes_home() / "cache" / _CACHE_FILENAME
|
||||
|
||||
|
||||
@@ -53,7 +52,6 @@ def _load_all() -> Dict[str, Any]:
|
||||
|
||||
def _save_all(data: Dict[str, Any]) -> None:
|
||||
from utils import atomic_json_write
|
||||
|
||||
# 0o600: the cache file is trusted input on the lazy registration path, keep it user-only.
|
||||
atomic_json_write(_cache_path(), data, mode=0o600)
|
||||
|
||||
|
||||
+17
-33
@@ -113,16 +113,6 @@ 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).
|
||||
_MCP_SDK_LAZY_SYMBOLS = frozenset({
|
||||
"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) bound in this order to
|
||||
# _MCP_SAMPLING_TYPES / _MCP_ELICITATION_TYPES / _MCP_NOTIFICATION_TYPES. Each is gated
|
||||
# separately so an older SDK only loses that feature, not MCP.
|
||||
@@ -136,6 +126,12 @@ _OPTIONAL_TYPE_FAMILIES = (
|
||||
"ResourceListChangedNotification"),
|
||||
"MCP notification types not available -- dynamic tool discovery disabled"),
|
||||
)
|
||||
# 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"}
|
||||
| {n for _mod, names, _msg in _OPTIONAL_TYPE_FAMILIES for n in names})
|
||||
|
||||
|
||||
def __getattr__(name: str):
|
||||
@@ -173,7 +169,6 @@ def _ensure_mcp_sdk() -> bool:
|
||||
global _MCP_SAMPLING_TYPES, _MCP_NOTIFICATION_TYPES, _MCP_ELICITATION_TYPES, sse_client
|
||||
global _MCP_MESSAGE_HANDLER_SUPPORTED, _MCP_LOGGING_CALLBACK_SUPPORTED, LATEST_HANDSHAKE_VERSION
|
||||
global _JSONRPC_METHOD_NOT_FOUND
|
||||
|
||||
if not _MCP_AVAILABLE:
|
||||
return False
|
||||
if _MCP_SDK_IMPORT_ATTEMPTED or ClientSession is not None:
|
||||
@@ -201,13 +196,11 @@ def _ensure_mcp_sdk() -> bool:
|
||||
_import_sdk_names(*family) for family in _OPTIONAL_TYPE_FAMILIES]
|
||||
else:
|
||||
logger.debug("mcp package not installed -- MCP tool support disabled")
|
||||
|
||||
if _MCP_AVAILABLE:
|
||||
try:
|
||||
_JSONRPC_METHOD_NOT_FOUND = importlib.import_module("mcp.types").METHOD_NOT_FOUND
|
||||
except Exception: # pragma: no cover — SDK without the constant
|
||||
pass
|
||||
|
||||
_MCP_MESSAGE_HANDLER_SUPPORTED = _client_session_accepts("message_handler")
|
||||
if _MCP_AVAILABLE and not _MCP_MESSAGE_HANDLER_SUPPORTED:
|
||||
logger.debug("MCP SDK does not support message_handler -- dynamic tool discovery disabled")
|
||||
@@ -236,15 +229,13 @@ def sdk_httpx():
|
||||
_SDK_HTTPX_MOD = getattr(_transport, "httpx2", None) or getattr(_transport, "httpx", None)
|
||||
except ImportError:
|
||||
_SDK_HTTPX_MOD = None
|
||||
if _SDK_HTTPX_MOD is None:
|
||||
for fallback in ("httpx2", "httpx"):
|
||||
if _SDK_HTTPX_MOD is not None:
|
||||
break
|
||||
try:
|
||||
import httpx2 as _fallback
|
||||
_SDK_HTTPX_MOD = importlib.import_module(fallback)
|
||||
except ImportError:
|
||||
try:
|
||||
import httpx as _fallback # type: ignore[no-redef]
|
||||
except ImportError:
|
||||
return None
|
||||
_SDK_HTTPX_MOD = _fallback
|
||||
pass
|
||||
return _SDK_HTTPX_MOD
|
||||
|
||||
|
||||
@@ -314,21 +305,16 @@ async def _paginate_full_list(list_method, items_attr: str, server_name: str,
|
||||
# mcp 2.0 takes params=PaginatedRequestParams, 1.x takes cursor=.
|
||||
try:
|
||||
import mcp.types as _types # late: keeps the SDK import lazy
|
||||
|
||||
_params_cls = getattr(_types, "PaginatedRequestParams", None)
|
||||
if _params_cls is not None:
|
||||
result = await list_method(params=_params_cls(cursor=cursor))
|
||||
else:
|
||||
result = await list_method(cursor=cursor)
|
||||
result = await (list_method(params=_params_cls(cursor=cursor)) if _params_cls is not None
|
||||
else list_method(cursor=cursor))
|
||||
except TypeError:
|
||||
result = await list_method(cursor=cursor)
|
||||
if cache_meta_out is not None and not items:
|
||||
_ttl = mcp_field(result, "ttl_ms", "ttlMs")
|
||||
_scope = mcp_field(result, "cache_scope", "cacheScope")
|
||||
if _ttl is not None:
|
||||
cache_meta_out["ttl_ms"] = _ttl
|
||||
if _scope is not None:
|
||||
cache_meta_out["cache_scope"] = _scope
|
||||
for key, snake, camel in (("ttl_ms", "ttl_ms", "ttlMs"), ("cache_scope", "cache_scope", "cacheScope")):
|
||||
hint = mcp_field(result, snake, camel)
|
||||
if hint is not None:
|
||||
cache_meta_out[key] = hint
|
||||
items.extend(getattr(result, items_attr, None) or [])
|
||||
cursor = mcp_field(result, "next_cursor", "nextCursor")
|
||||
# Cursor is an opaque string; anything else (incl. mocks) = last page.
|
||||
@@ -517,11 +503,9 @@ def _mcp_registry_scope() -> Optional[str]:
|
||||
"""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():
|
||||
return None
|
||||
from tools.registry import registry
|
||||
|
||||
return registry.current_scope_key()
|
||||
|
||||
|
||||
|
||||
@@ -21,7 +21,6 @@ class _OriginProxy:
|
||||
|
||||
def __getattr__(self, name: str):
|
||||
from tools import mcp_tool
|
||||
|
||||
return getattr(mcp_tool, name)
|
||||
|
||||
|
||||
@@ -51,7 +50,6 @@ def _resolve_tool_timeout(config: dict) -> float:
|
||||
return per_server
|
||||
try:
|
||||
from agent.deadline import resolve_timeout
|
||||
|
||||
resolved = resolve_timeout("mcp.tool_call", default=_DEFAULT_TOOL_TIMEOUT)
|
||||
if resolved is not None:
|
||||
return resolved
|
||||
@@ -151,9 +149,8 @@ def _get_lifecycle_seconds(config: dict, key: str) -> Optional[float]:
|
||||
"""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):
|
||||
raw = lifecycle.get(key)
|
||||
if raw is None and isinstance(config.get("lifecycle"), dict):
|
||||
raw = config["lifecycle"].get(key)
|
||||
if raw is None:
|
||||
return None
|
||||
try:
|
||||
|
||||
+22
-51
@@ -72,9 +72,7 @@ async def _connect_server(name: str, config: dict) -> _core.MCPServerTask:
|
||||
|
||||
def _request_lazy_reconnect(server_name: str, server: _core.MCPServerTask) -> bool:
|
||||
"""Wake a recycled stdio server and wait briefly for a fresh session."""
|
||||
if not server._is_recycled_stdio():
|
||||
return False
|
||||
loop = _core._running_loop()
|
||||
loop = _core._running_loop() if server._is_recycled_stdio() else None
|
||||
if loop is None:
|
||||
return False
|
||||
|
||||
@@ -122,6 +120,13 @@ def _note_connect_success(name: str) -> None:
|
||||
_core._clear_connect_failure(name)
|
||||
|
||||
|
||||
def _adopt_server(name: str, server: _core.MCPServerTask) -> None:
|
||||
"""Publish *server* into ``_servers`` with its owning registry scope (under ``_lock``)."""
|
||||
with _core._lock:
|
||||
_core._servers[name] = server
|
||||
_core._server_scope_keys[name] = _core._mcp_registry_scope()
|
||||
|
||||
|
||||
def _ensure_lazy_server_connected(server_name: str) -> bool:
|
||||
"""Connect a lazily-registered server on demand (sync; blocks the caller).
|
||||
|
||||
@@ -139,21 +144,15 @@ def _ensure_lazy_server_connected(server_name: str) -> bool:
|
||||
return False
|
||||
_core._server_connecting.add(server_name)
|
||||
_core._server_connect_errors.pop(server_name, None)
|
||||
|
||||
logger.info("MCP server '%s': lazy start on first use", server_name)
|
||||
_core._ensure_mcp_loop()
|
||||
connect_timeout = config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT)
|
||||
|
||||
async def _connect():
|
||||
return await _core._discover_and_register_server(server_name, config)
|
||||
|
||||
try:
|
||||
_core._run_on_mcp_loop(_connect, timeout=float(connect_timeout) + 30.0)
|
||||
_core._run_on_mcp_loop(lambda: _core._discover_and_register_server(server_name, config),
|
||||
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, _note_connect_failure(server_name, exc))
|
||||
return False
|
||||
|
||||
_note_connect_success(server_name)
|
||||
with _core._lock:
|
||||
_core._lazy_server_configs.pop(server_name, None)
|
||||
@@ -165,7 +164,6 @@ def _ensure_lazy_server_connected(server_name: str) -> bool:
|
||||
phantom_names = [n for n in cached_names if n not in live_names]
|
||||
if phantom_names:
|
||||
from tools.registry import registry
|
||||
|
||||
for tool_name in phantom_names:
|
||||
registry.deregister(tool_name, scope=_core._server_registry_scope(server_name))
|
||||
_core._forget_mcp_tool_server(tool_name)
|
||||
@@ -184,25 +182,23 @@ def _get_connected_server_for_call(server_name: str) -> Optional[_core.MCPServer
|
||||
is_lazy = server_name in _core._lazy_server_configs
|
||||
if is_lazy and (server is None or server.session is None):
|
||||
_core._ensure_lazy_server_connected(server_name)
|
||||
with _core._lock:
|
||||
server = _core._servers.get(server_name)
|
||||
return server
|
||||
if server is not None and server.session is None and server._is_recycled_stdio():
|
||||
elif server is not None and server.session is None and server._is_recycled_stdio():
|
||||
_core._request_lazy_reconnect(server_name, server)
|
||||
with _core._lock:
|
||||
server = _core._servers.get(server_name)
|
||||
return server
|
||||
else:
|
||||
return server
|
||||
with _core._lock:
|
||||
return _core._servers.get(server_name)
|
||||
|
||||
|
||||
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.
|
||||
claimed: List[_core.MCPServerTask] = []
|
||||
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=config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT))
|
||||
except BaseException:
|
||||
server = claimed[0] if claimed else None
|
||||
task = server._task if server is not None else None
|
||||
@@ -211,21 +207,16 @@ async def _discover_and_register_server(name: str, config: dict) -> List[str]:
|
||||
and not task.done() and not task_cancelling):
|
||||
# Recoverable park: the run task 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()
|
||||
_adopt_server(name, server)
|
||||
elif server is not None:
|
||||
await server.shutdown()
|
||||
raise
|
||||
finally:
|
||||
_core._connect_server_claim.reset(claim_token)
|
||||
|
||||
with _core._lock:
|
||||
_core._server_connecting.discard(name)
|
||||
_core._server_connect_errors.pop(name, None)
|
||||
_core._servers[name] = server
|
||||
_core._server_scope_keys[name] = _core._mcp_registry_scope()
|
||||
|
||||
_adopt_server(name, server)
|
||||
registered_names = _core._register_server_tools(name, server, config)
|
||||
server._registered_tool_names = list(registered_names)
|
||||
logger.info("MCP server '%s' (%s): registered %d tool(s): %s", name,
|
||||
@@ -258,7 +249,6 @@ def _select_new_servers(servers: Dict[str, dict]) -> Dict[str, dict]:
|
||||
_core._parallel_safe_servers.add(srv_name)
|
||||
else:
|
||||
_core._parallel_safe_servers.discard(srv_name)
|
||||
|
||||
for srv in stale_cached:
|
||||
_core._signal_reconnect(srv)
|
||||
return new_servers
|
||||
@@ -366,23 +356,19 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]:
|
||||
if not _core._ensure_mcp_sdk():
|
||||
logger.debug("MCP SDK not available -- skipping explicit MCP registration")
|
||||
return []
|
||||
|
||||
servers = _core._filter_suspicious_mcp_servers(servers)
|
||||
if not servers:
|
||||
logger.debug("No explicit MCP servers provided")
|
||||
return []
|
||||
|
||||
new_servers = _select_new_servers(servers)
|
||||
if not new_servers:
|
||||
return _core._existing_tool_names()
|
||||
|
||||
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)
|
||||
return _core._existing_tool_names()
|
||||
|
||||
_core._ensure_mcp_loop()
|
||||
_run_discovery_pass(new_servers)
|
||||
_log_summary("MCP: registered", new_servers, lazy_tools=lazy_registered, lazy_servers=lazy_server_count)
|
||||
@@ -402,7 +388,6 @@ def _acquire_discovery_lock_with_retry():
|
||||
cookie = _core._try_acquire_mcp_discovery_lock()
|
||||
if cookie is not None:
|
||||
break
|
||||
|
||||
if cookie is None:
|
||||
logger.warning("MCP discovery lock still held after %d retries -- running discovery unguarded",
|
||||
_core._MCP_DISCOVERY_LOCK_MAX_RETRIES)
|
||||
@@ -419,19 +404,16 @@ def discover_mcp_tools() -> List[str]:
|
||||
if not servers:
|
||||
logger.debug("No MCP servers configured")
|
||||
return []
|
||||
|
||||
# SDK import deferred to here so a config without servers never pays it.
|
||||
if not _core._ensure_mcp_sdk():
|
||||
logger.debug("MCP SDK not available -- skipping MCP tool discovery")
|
||||
return []
|
||||
|
||||
cookie = _acquire_discovery_lock_with_retry()
|
||||
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 _enabled(cfg)]
|
||||
|
||||
tool_names = _core.register_mcp_servers(servers)
|
||||
if new_server_names:
|
||||
_log_summary(" MCP:", new_server_names)
|
||||
@@ -461,7 +443,6 @@ def get_mcp_status() -> List[dict]:
|
||||
configured = _core._load_mcp_config()
|
||||
if not configured:
|
||||
return []
|
||||
|
||||
with _core._lock:
|
||||
active_servers = dict(_core._servers)
|
||||
connecting = set(_core._server_connecting)
|
||||
@@ -474,7 +455,6 @@ def get_mcp_status() -> List[dict]:
|
||||
result: List[dict] = []
|
||||
for name, cfg in configured.items():
|
||||
transport = cfg.get("transport", "http") if "url" in cfg else "stdio"
|
||||
enabled = _enabled(cfg)
|
||||
server = active_servers.get(name)
|
||||
if server and server.session is not None:
|
||||
entry = _entry(name, transport, "connected", connected=True)
|
||||
@@ -482,7 +462,7 @@ def get_mcp_status() -> List[dict]:
|
||||
else len(server._tools))
|
||||
if server._sampling:
|
||||
entry["sampling"] = dict(server._sampling.metrics)
|
||||
elif not enabled:
|
||||
elif not _enabled(cfg):
|
||||
entry = _entry(name, transport, "disabled")
|
||||
elif name in connecting:
|
||||
entry = _entry(name, transport, "connecting")
|
||||
@@ -491,7 +471,6 @@ def get_mcp_status() -> List[dict]:
|
||||
else:
|
||||
entry = _entry(name, transport, "configured")
|
||||
result.append(entry)
|
||||
|
||||
return result
|
||||
|
||||
|
||||
@@ -500,17 +479,10 @@ def probe_mcp_server_tools() -> Dict[str, List[tuple]]:
|
||||
without registering anything. Failed servers are omitted."""
|
||||
if not _core._ensure_mcp_sdk():
|
||||
return {}
|
||||
|
||||
servers_config = _core._load_mcp_config()
|
||||
if not servers_config:
|
||||
return {}
|
||||
|
||||
enabled = {k: v for k, v in servers_config.items() if _enabled(v)}
|
||||
enabled = {k: v for k, v in (_core._load_mcp_config() or {}).items() if _enabled(v)}
|
||||
if not enabled:
|
||||
return {}
|
||||
|
||||
_core._ensure_mcp_loop()
|
||||
|
||||
result: Dict[str, List[tuple]] = {}
|
||||
probed_servers: List[_core.MCPServerTask] = []
|
||||
|
||||
@@ -533,7 +505,6 @@ def probe_mcp_server_tools() -> Dict[str, List[tuple]]:
|
||||
logger.debug("MCP probe failed: %s", exc)
|
||||
finally:
|
||||
_core._stop_mcp_loop_if_idle()
|
||||
|
||||
return result
|
||||
|
||||
|
||||
|
||||
+11
-22
@@ -33,7 +33,6 @@ def _trust_gate_check(server_name: str, tool_name: str) -> Optional[str]:
|
||||
# Lazy import: tools.approval routes the prompt to whichever surface owns the session.
|
||||
try:
|
||||
from tools.approval import request_elicitation_consent
|
||||
|
||||
answer = request_elicitation_consent(
|
||||
f"MCP tool '{tool_name}' on UNTRUSTED server '{server_name}' wants to run. This "
|
||||
f"tool is write-capable (no readOnlyHint=true annotation) and may modify external state.",
|
||||
@@ -113,8 +112,7 @@ def _strike(server_name: str, message: str, **extra) -> str:
|
||||
|
||||
|
||||
def _mcp_loop_running() -> bool:
|
||||
loop = _core._mcp_loop
|
||||
return loop is not None and loop.is_running()
|
||||
return _core._mcp_loop is not None and _core._mcp_loop.is_running()
|
||||
|
||||
|
||||
def _lookup_reconnectable_server(server_name: str, require_loop: bool = False):
|
||||
@@ -153,12 +151,8 @@ def _handle_auth_error_and_retry(server_name: str, exc: BaseException, retry_cal
|
||||
return None
|
||||
from tools.mcp_oauth_manager import get_manager
|
||||
manager = get_manager()
|
||||
|
||||
async def _recover():
|
||||
return await manager.handle_401(server_name, None)
|
||||
|
||||
try:
|
||||
recovered = _core._run_on_mcp_loop(_recover, timeout=10)
|
||||
recovered = _core._run_on_mcp_loop(lambda: manager.handle_401(server_name, None), timeout=10)
|
||||
except Exception as rec_exc:
|
||||
logger.warning("MCP OAuth '%s': recovery attempt failed: %s", server_name, rec_exc)
|
||||
recovered = False
|
||||
@@ -223,7 +217,6 @@ def _handle_stdio_child_exited_and_retry(server_name: str, exc: Exception, retry
|
||||
# No MCP loop to wait on (non-async adapters, tests) — still request the respawn
|
||||
# so the next call lands on a live transport.
|
||||
_core._signal_reconnect(srv)
|
||||
|
||||
if not reconnected:
|
||||
return _strike(
|
||||
server_name,
|
||||
@@ -277,9 +270,8 @@ def _invoke_with_recovery(server_name: str, call_once: Callable[[], str], op: st
|
||||
|
||||
def _mark_server_call_started(server: Any) -> None:
|
||||
"""Record a user-visible MCP operation when the server supports it."""
|
||||
mark_tool_call = getattr(server, "mark_tool_call", None)
|
||||
if callable(mark_tool_call):
|
||||
mark_tool_call()
|
||||
if callable(getattr(server, "mark_tool_call", None)):
|
||||
server.mark_tool_call()
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
@@ -320,7 +312,6 @@ async def _call_tool_racing_stdio_death(server, server_name: str, tool_name: str
|
||||
if not (_watch_children is not None and inspect.iscoroutinefunction(_watch_children) and asyncio.iscoroutine(_call_coro)):
|
||||
# Stubbed sessions return a non-awaitable, or there is no child-watcher to race: plain await.
|
||||
return await _call_coro if asyncio.iscoroutine(_call_coro) else _call_coro
|
||||
|
||||
rpc_task = asyncio.ensure_future(_call_coro)
|
||||
watch_task = asyncio.ensure_future(_watch_children())
|
||||
try:
|
||||
@@ -388,7 +379,6 @@ def _render_call_tool_result(result, server_name: str) -> str:
|
||||
``.is_error`` is ``.isError`` before mcp 2.0."""
|
||||
if mcp_field(result, "is_error", "isError", False):
|
||||
return tool_error(_sanitize_error(_truncate_mcp_text_result(_error_result_text(result) or "MCP tool returned an error")))
|
||||
|
||||
text_result = _render_content_blocks(result, server_name)
|
||||
structured = _capped_structured_content(result)
|
||||
meta = _strip_reserved_meta_keys(mcp_field(result, "meta", "meta"))
|
||||
@@ -436,9 +426,8 @@ 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.
|
||||
_mark_proven = getattr(server, "_mark_session_proven", None)
|
||||
if _mark_proven is not None:
|
||||
_mark_proven()
|
||||
if getattr(server, "_mark_session_proven", None) is not None:
|
||||
server._mark_session_proven()
|
||||
return _render_call_tool_result(result, server_name)
|
||||
|
||||
def _on_failure(exc):
|
||||
@@ -502,9 +491,9 @@ def _render_resource_list(all_resources, server_name: str) -> dict:
|
||||
if "uri" in entry:
|
||||
entry["uri"] = str(entry["uri"])
|
||||
# Key stays camelCase — this is the tool's own JSON output shape.
|
||||
_mime = mcp_field(r, "mime_type", "mimeType")
|
||||
if _mime:
|
||||
entry["mimeType"] = _mime
|
||||
mime = mcp_field(r, "mime_type", "mimeType")
|
||||
if mime:
|
||||
entry["mimeType"] = mime
|
||||
resources.append(entry)
|
||||
return {"resources": resources}
|
||||
|
||||
@@ -539,8 +528,7 @@ def _render_get_prompt(result, server_name: str) -> dict:
|
||||
for msg in getattr(result, "messages", []):
|
||||
entry = _pick(msg, ("role", "role"))
|
||||
if hasattr(msg, "content"):
|
||||
content = msg.content
|
||||
entry["content"] = strip_unicode_tags(content.text if hasattr(content, "text") else str(content))
|
||||
entry["content"] = strip_unicode_tags(msg.content.text if hasattr(msg.content, "text") else str(msg.content))
|
||||
messages.append(entry)
|
||||
resp = {"messages": messages}
|
||||
if getattr(result, "description", None):
|
||||
@@ -552,6 +540,7 @@ def _utility_factory(op: str, log_label: str, rpc, render, required: Optional[st
|
||||
"""``(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
|
||||
|
||||
|
||||
|
||||
@@ -154,13 +154,9 @@ class _CachedMCPTool:
|
||||
@classmethod
|
||||
def from_cache_dicts(cls, raws: Iterable[Any]) -> List["_CachedMCPTool"]:
|
||||
"""Cached rows -> stand-ins; rows that are not dicts or lack a name are dropped."""
|
||||
out = []
|
||||
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")))
|
||||
return out
|
||||
return [cls(raw["name"], raw.get("description") or "",
|
||||
raw["inputSchema"] if isinstance(raw.get("inputSchema"), dict) else {}, raw.get("annotations"))
|
||||
for raw in raws if isinstance(raw, dict) and raw.get("name")]
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -224,7 +220,6 @@ def _resolve_name_collisions(name: str, candidates: List[_Candidate]) -> List[_C
|
||||
seen.add((c.registry_name, c.origin))
|
||||
unique.append(c)
|
||||
origins_by_name.setdefault(c.registry_name, set()).add(c.origin)
|
||||
|
||||
ambiguous: Dict[str, List[str]] = {}
|
||||
shadowed: set[tuple[str, str]] = set()
|
||||
for registry_name, origins in origins_by_name.items():
|
||||
@@ -273,7 +268,6 @@ def _register_candidates(name: str, candidates: List[_Candidate], *, check_fn: C
|
||||
``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}"
|
||||
registered: List[str] = []
|
||||
for c in candidates:
|
||||
@@ -301,7 +295,6 @@ def _write_schema_cache(name: str, server: "MCPServerTask", config: dict, should
|
||||
lazily without spawning it. Never raises."""
|
||||
try:
|
||||
from tools.mcp_schema_cache import config_fingerprint, write_cache_entry
|
||||
|
||||
tools_payload = []
|
||||
for t in server._tools:
|
||||
if should_register(t.name):
|
||||
@@ -345,7 +338,6 @@ def _register_from_cache_sync(name: str, config: dict, entry: dict) -> List[str]
|
||||
``_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)
|
||||
cached_tools = _CachedMCPTool.from_cache_dicts(tools_from_cache_entry(entry))
|
||||
_record_tool_trust_metadata(name, config, cached_tools)
|
||||
|
||||
@@ -101,12 +101,10 @@ def _normalize_mcp_input_schema(schema: dict | None) -> dict:
|
||||
if not schema:
|
||||
return dict(_EMPTY_OBJECT_SCHEMA)
|
||||
from tools.schema_sanitizer import collapse_const_unions, strip_nullable_unions
|
||||
|
||||
normalized = _rewrite_local_refs(schema)
|
||||
normalized = strip_nullable_unions(normalized, keep_nullable_hint=True)
|
||||
normalized = collapse_const_unions(normalized)
|
||||
normalized = _repair_object_shape(normalized)
|
||||
|
||||
if not isinstance(normalized, dict):
|
||||
return dict(_EMPTY_OBJECT_SCHEMA)
|
||||
if normalized.get("type") == "object" and "properties" not in normalized:
|
||||
|
||||
@@ -35,6 +35,11 @@ class MCPServerRunMixin:
|
||||
except (asyncio.CancelledError, Exception):
|
||||
pass
|
||||
|
||||
def _event_waiters(self) -> tuple:
|
||||
"""Fresh ``(shutdown, reconnect)`` wait tasks; cancel them via ``_cancel_waiters``."""
|
||||
return (asyncio.ensure_future(self._shutdown_event.wait()),
|
||||
asyncio.ensure_future(self._reconnect_event.wait()))
|
||||
|
||||
def _recycle_if_due(self) -> bool:
|
||||
"""Latch a stdio idle/lifetime recycle when its deadline has passed."""
|
||||
recycle_reason = self._stdio_recycle_reason()
|
||||
@@ -56,26 +61,21 @@ class MCPServerRunMixin:
|
||||
keepalive_interval = max(
|
||||
_core._MIN_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())
|
||||
shutdown_task, reconnect_task = self._event_waiters()
|
||||
try:
|
||||
while True:
|
||||
if self._recycle_if_due():
|
||||
return "recycle"
|
||||
|
||||
timeout = keepalive_interval
|
||||
recycle_deadline = self._next_stdio_recycle_deadline()
|
||||
if recycle_deadline is not None:
|
||||
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)
|
||||
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).
|
||||
@@ -98,7 +98,6 @@ class MCPServerRunMixin:
|
||||
self._mark_session_proven()
|
||||
finally:
|
||||
await self._cancel_waiters(shutdown_task, reconnect_task)
|
||||
|
||||
if self._shutdown_event.is_set():
|
||||
self._fail_inflight_calls("shutdown")
|
||||
return "shutdown"
|
||||
@@ -112,8 +111,7 @@ class MCPServerRunMixin:
|
||||
"""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())
|
||||
shutdown_task, reconnect_task = self._event_waiters()
|
||||
try:
|
||||
await asyncio.wait({shutdown_task, reconnect_task}, return_when=asyncio.FIRST_COMPLETED, timeout=timeout)
|
||||
finally:
|
||||
@@ -154,10 +152,8 @@ class MCPServerRunMixin:
|
||||
self._auth_type = (config.get("auth") or "").lower().strip()
|
||||
self._idle_timeout_seconds = _core._get_lifecycle_seconds(config, "idle_timeout_seconds")
|
||||
self._max_lifetime_seconds = _core._get_lifecycle_seconds(config, "max_lifetime_seconds")
|
||||
|
||||
# The _MCP_*_TYPES flags are False until the lazy SDK import runs.
|
||||
_core._ensure_mcp_sdk()
|
||||
|
||||
sampling_config = config.get("sampling", {})
|
||||
self._sampling = (_core.SamplingHandler(self.name, sampling_config)
|
||||
if sampling_config.get("enabled", True) and _core._MCP_SAMPLING_TYPES else None)
|
||||
@@ -166,12 +162,10 @@ class MCPServerRunMixin:
|
||||
elicitation_config = config.get("elicitation", {})
|
||||
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)
|
||||
|
||||
if not self._is_http():
|
||||
return True
|
||||
try:
|
||||
@@ -206,17 +200,12 @@ class MCPServerRunMixin:
|
||||
"""
|
||||
if not await self._prepare_run(config):
|
||||
return
|
||||
|
||||
self._reconnect_retries = 0
|
||||
budget = _RetryBudget()
|
||||
|
||||
while True:
|
||||
try:
|
||||
if self._is_http():
|
||||
lifecycle_reason = await self._run_http(config)
|
||||
else:
|
||||
lifecycle_reason = await self._run_stdio(config)
|
||||
if not await self._on_clean_return(lifecycle_reason, budget):
|
||||
run_transport = self._run_http if self._is_http() else self._run_stdio
|
||||
if not await self._on_clean_return(await run_transport(config), budget):
|
||||
break
|
||||
except asyncio.CancelledError:
|
||||
# Not a connection failure: re-raise so cancellation reaches asyncio and
|
||||
@@ -253,11 +242,9 @@ class MCPServerRunMixin:
|
||||
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
|
||||
self._teardown_race, budget.backoff = False, 1.0
|
||||
elif self._session_proven:
|
||||
self._reconnect_retries = 0
|
||||
budget.backoff = 1.0
|
||||
self._reconnect_retries, budget.backoff = 0, 1.0
|
||||
else:
|
||||
self._reconnect_retries += 1
|
||||
if self._reconnect_retries > _core._MAX_RECONNECT_RETRIES:
|
||||
@@ -280,8 +267,7 @@ class MCPServerRunMixin:
|
||||
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
|
||||
self._reconnect_retries, budget.backoff = _core._MAX_RECONNECT_RETRIES, 1.0
|
||||
return True
|
||||
|
||||
async def _park_initial_failure(self, exc: Exception, revival_reason: str, budget: "_RetryBudget") -> bool:
|
||||
@@ -291,8 +277,7 @@ class MCPServerRunMixin:
|
||||
self._ready.set()
|
||||
if await self._park(revival_reason):
|
||||
return False
|
||||
budget.initial_retries = 0
|
||||
self._reconnect_retries = 0
|
||||
budget.initial_retries = self._reconnect_retries = 0
|
||||
budget.backoff = 1.0
|
||||
self._error = None
|
||||
self._ready.clear()
|
||||
@@ -314,20 +299,16 @@ class MCPServerRunMixin:
|
||||
"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).
|
||||
if not self._ever_connected:
|
||||
return await self._on_initial_connect_error(exc, root, failure_class, budget)
|
||||
|
||||
if self._shutdown_event.is_set():
|
||||
logger.debug("MCP server '%s' disconnected during shutdown: %s: %s",
|
||||
self.name, type(root).__name__, root)
|
||||
return False
|
||||
|
||||
if failure_class == "permanent":
|
||||
return await self._on_permanent_error(root, budget)
|
||||
|
||||
self._reconnect_retries += 1
|
||||
if self._reconnect_retries > _core._MAX_RECONNECT_RETRIES:
|
||||
logger.warning(
|
||||
@@ -337,7 +318,6 @@ class MCPServerRunMixin:
|
||||
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,
|
||||
@@ -351,19 +331,12 @@ class MCPServerRunMixin:
|
||||
# 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)
|
||||
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)
|
||||
detail = (f"authentication, parking until credentials change; re-authenticate with "
|
||||
f"`hermes mcp login {self.name}`" if _core._is_auth_error(root)
|
||||
else "connection with a permanent error, parking without retries")
|
||||
logger.warning("MCP server '%s' failed initial %s (state: connecting → parked): %s: %s",
|
||||
self.name, detail, type(root).__name__, root)
|
||||
return await self._park_initial_failure(exc, "after permanent initial failure", budget)
|
||||
|
||||
budget.initial_retries += 1
|
||||
if budget.initial_retries > _core._MAX_INITIAL_CONNECT_RETRIES:
|
||||
logger.warning(
|
||||
@@ -372,7 +345,6 @@ class MCPServerRunMixin:
|
||||
"requested (state: connecting → parked): %s: %s",
|
||||
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,
|
||||
@@ -381,8 +353,7 @@ class MCPServerRunMixin:
|
||||
if self._shutdown_event.is_set():
|
||||
self._error = exc
|
||||
self._ready.set()
|
||||
return False
|
||||
return True
|
||||
return not self._shutdown_event.is_set()
|
||||
|
||||
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
|
||||
@@ -395,8 +366,7 @@ class MCPServerRunMixin:
|
||||
"healthy session — marking suspect and forcing "
|
||||
"one reconnect instead of parking (state: connected → suspect): %s: %s",
|
||||
self.name, type(root).__name__, root)
|
||||
self._reconnect_retries = 0
|
||||
budget.backoff = 1.0
|
||||
self._reconnect_retries, budget.backoff = 0, 1.0
|
||||
await asyncio.sleep(_core._jittered(1.0))
|
||||
return not self._shutdown_event.is_set()
|
||||
# Deterministic failure on a working server: park now.
|
||||
@@ -426,9 +396,8 @@ class MCPServerRunMixin:
|
||||
async def shutdown(self):
|
||||
"""Signal the Task to exit and wait for clean resource teardown."""
|
||||
self._shutdown_event.set()
|
||||
# _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".
|
||||
# Also set reconnect: closes any race where _wait_for_lifecycle_event misses the
|
||||
# shutdown flag after returning "reconnect".
|
||||
self._reconnect_event.set()
|
||||
if self._task and not self._task.done():
|
||||
try:
|
||||
@@ -453,7 +422,6 @@ class MCPServerRunMixin:
|
||||
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", [])):
|
||||
registry.deregister(tool_name, scope=_core._server_registry_scope(self.name))
|
||||
_core._forget_mcp_tool_server(tool_name)
|
||||
|
||||
Reference in New Issue
Block a user