refactor(tools): MCP discovery/run/handlers defensive collapse, shared helpers, blank squeeze

This commit is contained in:
Teknium
2026-09-02 23:07:18 -07:00
parent 47be9a3629
commit 7f7fd4a533
8 changed files with 77 additions and 180 deletions
-2
View File
@@ -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
View File
@@ -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()
+2 -5
View File
@@ -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
View File
@@ -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
View File
@@ -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
+3 -11
View File
@@ -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)
-2
View File
@@ -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:
+22 -54
View File
@@ -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)