From faceb453374043da2a38ed3c4981dd8190fdceac Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 15:48:34 -0700 Subject: [PATCH] refactor(mcp): split orphan reaper/config/schema/agent-refresh helpers; table-build utility schemas and injection patterns --- tools/mcp_schema_cache.py | 29 ++-- tools/mcp_stdio_watchdog.py | 19 +-- tools/mcp_tool_agent.py | 140 +++++++--------- tools/mcp_tool_common.py | 23 +-- tools/mcp_tool_config.py | 253 ++++++++++++---------------- tools/mcp_tool_lifecycle.py | 188 +++++++++------------ tools/mcp_tool_schema.py | 318 ++++++++++++++---------------------- 7 files changed, 397 insertions(+), 573 deletions(-) diff --git a/tools/mcp_schema_cache.py b/tools/mcp_schema_cache.py index 0e3377bd9a..48ce08c12f 100644 --- a/tools/mcp_schema_cache.py +++ b/tools/mcp_schema_cache.py @@ -77,19 +77,16 @@ def get_cached_entry(server_name: str, fingerprint: str) -> Optional[dict]: """ with _cache_lock: entry = _load_all().get(server_name) - if not isinstance(entry, dict): - return None - if entry.get("fingerprint") != fingerprint: + if not isinstance(entry, dict) or entry.get("fingerprint") != fingerprint: return None ttl_ms = entry.get("ttl_ms") written_at = entry.get("written_at") - if ( + expired = ( isinstance(ttl_ms, (int, float)) and isinstance(written_at, (int, float)) and (time.time() - written_at) * 1000.0 >= float(ttl_ms) - ): - return None - return entry + ) + return None if expired else entry def write_cache_entry( @@ -107,11 +104,7 @@ def write_cache_entry( ``tools/list`` result (2026-07-28 servers). ``written_at`` anchors TTL expiry in :func:`get_cached_entry`. """ - entry = { - "fingerprint": fingerprint, - "tools": tools, - "utility_tools": utility_tools or [], - } + entry = {"fingerprint": fingerprint, "tools": tools, "utility_tools": utility_tools or []} if isinstance(ttl_ms, (int, float)): entry["ttl_ms"] = ttl_ms entry["written_at"] = time.time() @@ -131,12 +124,16 @@ def write_cache_entry( _save_all(data) +def _list_field(entry: dict, key: str) -> List[dict]: + value = entry.get(key) + return list(value) if isinstance(value, list) else [] + + def tools_from_cache_entry(entry: dict) -> List[dict]: """Return cached MCP tool dicts (name, description, inputSchema).""" - tools = entry.get("tools") - return list(tools) if isinstance(tools, list) else [] + return _list_field(entry, "tools") def utility_tools_from_cache_entry(entry: dict) -> List[dict]: - util = entry.get("utility_tools") - return list(util) if isinstance(util, list) else [] + """Return cached ``{schema, handler_key}`` utility rows.""" + return _list_field(entry, "utility_tools") diff --git a/tools/mcp_stdio_watchdog.py b/tools/mcp_stdio_watchdog.py index 7c4c1896b7..fa0570f9cd 100644 --- a/tools/mcp_stdio_watchdog.py +++ b/tools/mcp_stdio_watchdog.py @@ -78,9 +78,7 @@ def _watchdog_loop(proc: subprocess.Popen, original_ppid: int) -> None: def main(argv: list[str] | None = None) -> int: - parser = argparse.ArgumentParser( - description="Parent-death watchdog for a stdio MCP subprocess.", - ) + parser = argparse.ArgumentParser(description="Parent-death watchdog for a stdio MCP subprocess.") parser.add_argument("--ppid", type=int, required=True) parser.add_argument("command", nargs=argparse.REMAINDER) args = parser.parse_args(argv) @@ -94,13 +92,7 @@ def main(argv: list[str] | None = None) -> int: # New process group so we can killpg() the whole tree the real command may # spawn, without touching our own group or the original parent's. - proc = subprocess.Popen( - real_argv, - stdin=sys.stdin, - stdout=sys.stdout, - stderr=sys.stderr, - start_new_session=True, - ) + proc = subprocess.Popen(real_argv, stdin=sys.stdin, stdout=sys.stdout, stderr=sys.stderr, start_new_session=True) # The server lives in its OWN group, so the parent's shutdown killpg of # *our* group no longer reaches it. Forward SIGTERM/SIGINT to the child's @@ -113,12 +105,7 @@ def main(argv: list[str] | None = None) -> int: signal.signal(signal.SIGTERM, _forward_shutdown) signal.signal(signal.SIGINT, _forward_shutdown) - watchdog = threading.Thread( - target=_watchdog_loop, - args=(proc, args.ppid), - daemon=True, - ) - watchdog.start() + threading.Thread(target=_watchdog_loop, args=(proc, args.ppid), daemon=True).start() try: return proc.wait() diff --git a/tools/mcp_tool_agent.py b/tools/mcp_tool_agent.py index e7a509c7f3..2722d1b35b 100644 --- a/tools/mcp_tool_agent.py +++ b/tools/mcp_tool_agent.py @@ -16,6 +16,38 @@ logger = logging.getLogger("tools.mcp_tool") _agent_tools_lock = threading.Lock() +def _def_name(tool_def: dict) -> str: + return (tool_def.get("function") or {}).get("name", "") + + +def _agent_tool_defs(agent) -> list: + return list(getattr(agent, "tools", None) or []) + + +def _resolve_refresh_toolsets(agent, enabled_override, disabled_override): + """Explicit reloads pass freshly-resolved toolsets (so a server just ENABLED + in config is picked up) and the agent's selection is updated to match; + automatic paths pass nothing and reuse the build-time selection.""" + enabled = getattr(agent, "enabled_toolsets", None) + disabled = getattr(agent, "disabled_toolsets", None) + if enabled_override is not None or disabled_override is not None: + enabled = enabled_override if enabled_override is not None else enabled + disabled = disabled_override if disabled_override is not None else disabled + agent.enabled_toolsets = enabled + agent.disabled_toolsets = disabled + return enabled, disabled + + +def _tool_defs_content_changed(agent, new_defs: list) -> bool: + """Byte-level diff of the serialized tool arrays (dynamic schemas change + CONTENT under stable names); False if either side fails to serialize.""" + try: + dump = lambda defs: json.dumps(defs, sort_keys=True, separators=(",", ":"), default=str) # noqa: E731 + return dump(_agent_tool_defs(agent)) != dump(new_defs) + except Exception: # noqa: BLE001 + return False + + def refresh_agent_mcp_tools( agent, *, @@ -51,17 +83,7 @@ def refresh_agent_mcp_tools( from model_tools import get_tool_definitions from tools.registry import registry - # Explicit reloads pass freshly-resolved toolsets (so a server just ENABLED - # in config is picked up) and the agent's selection is updated to match; - # automatic paths pass nothing and reuse the build-time selection. - if enabled_override is not None or disabled_override is not None: - enabled = enabled_override if enabled_override is not None else getattr(agent, "enabled_toolsets", None) - disabled = disabled_override if disabled_override is not None else getattr(agent, "disabled_toolsets", None) - agent.enabled_toolsets = enabled - agent.disabled_toolsets = disabled - else: - enabled = getattr(agent, "enabled_toolsets", None) - disabled = getattr(agent, "disabled_toolsets", None) + enabled, disabled = _resolve_refresh_toolsets(agent, enabled_override, disabled_override) # Capture the registry generation BEFORE the slow get_tool_definitions call; # at publish time a slower caller holding an OLDER set must not clobber a @@ -70,15 +92,8 @@ def refresh_agent_mcp_tools( # Computed OUTSIDE the lock (can be slow); diff + publish happen together in # one critical section so concurrent callers can't torn-publish. - new_defs = list( - get_tool_definitions( - enabled_toolsets=enabled, - disabled_toolsets=disabled, - quiet_mode=quiet_mode, - ) - or [] - ) - new_names = {t["function"]["name"] for t in new_defs} + new_defs = list(get_tool_definitions(enabled_toolsets=enabled, disabled_toolsets=disabled, quiet_mode=quiet_mode) or []) + new_names = {_def_name(t) for t in new_defs} # Re-append the post-build families on LOCALS only; live agent attributes # are untouched until the single atomic publish below. @@ -102,34 +117,16 @@ def refresh_agent_mcp_tools( published_gen = published_gen_raw if isinstance(published_gen_raw, int) else -1 if snapshot_generation < published_gen: return set() # a newer snapshot already won - current_defs = list(getattr(agent, "tools", None) or []) - current = {t["function"]["name"] for t in current_defs} + current_defs = _agent_tool_defs(agent) + current = {_def_name(t) for t in current_defs} if preserve_prefix: - new_defs, new_names = _merge_preserving_prefix( - current_defs, new_defs, registered_names, - ) - if new_names == current: - # Same NAME set: no change for MCP-reload callers. Content-aware - # callers (compaction boundary) also diff serialized bytes, since - # dynamic schemas change CONTENT under stable names. - content_changed = False - if content_aware: - try: - _stable = json.dumps( - (getattr(agent, "tools", None) or []), - sort_keys=True, separators=(",", ":"), default=str, - ) - _new = json.dumps( - new_defs, sort_keys=True, separators=(",", ":"), - default=str, - ) - content_changed = _stable != _new - except Exception: # noqa: BLE001 - content_changed = False - if not content_changed: - # Record the generation so an in-flight older caller can't clobber. - agent._tool_snapshot_generation = max(published_gen, snapshot_generation) - return set() + new_defs, new_names = _merge_preserving_prefix(current_defs, new_defs, registered_names) + # Same NAME set: no change for MCP-reload callers. Content-aware callers + # (compaction boundary) also diff serialized bytes. + if new_names == current and not (content_aware and _tool_defs_content_changed(agent, new_defs)): + # Record the generation so an in-flight older caller can't clobber. + agent._tool_snapshot_generation = max(published_gen, snapshot_generation) + return set() agent.tools = new_defs agent.valid_tool_names = new_names # Publish context-engine routing names atomically with the snapshot. @@ -163,10 +160,7 @@ def persist_agent_tool_names(agent) -> None: if not db or not session_id: return try: - db.update_session_tool_names( - session_id, - [t["function"]["name"] for t in (getattr(agent, "tools", None) or [])], - ) + db.update_session_tool_names(session_id, [_def_name(t) for t in _agent_tool_defs(agent)]) except Exception: # noqa: BLE001 logger.debug("tool_names persist skipped", exc_info=True) @@ -184,8 +178,8 @@ def restore_agent_tool_prefix(agent, saved_names: list) -> bool: return False from tools.registry import registry - fresh_defs = list(getattr(agent, "tools", None) or []) - fresh = {t["function"]["name"]: t for t in fresh_defs} + fresh_defs = _agent_tool_defs(agent) + fresh = {_def_name(t): t for t in fresh_defs} saved_defs = [] for name in saved_names: entry_def = fresh.get(name) @@ -202,14 +196,12 @@ def restore_agent_tool_prefix(agent, saved_names: list) -> bool: return False agent.tools = merged agent.valid_tool_names = merged_names - if [t["function"]["name"] for t in merged] != list(saved_names): + if [_def_name(t) for t in merged] != list(saved_names): persist_agent_tool_names(agent) return True -def _merge_preserving_prefix( - current_defs: list, new_defs: list, registered_names: set, -) -> tuple[list, set]: +def _merge_preserving_prefix(current_defs: list, new_defs: list, registered_names: set) -> tuple[list, set]: """Fold a fresh tool snapshot into a live one without moving existing bytes. Ordered by ``current_defs`` (the cached request prefix): a name in both @@ -217,22 +209,17 @@ def _merge_preserving_prefix( kept if still registered (``check_fn`` flapped) and dropped if not; a name only in the fresh list is appended at the tail. """ - fresh = {} - for entry in new_defs: - name = (entry.get("function") or {}).get("name", "") - if name: - fresh[name] = entry - + fresh = {_def_name(entry): entry for entry in new_defs if _def_name(entry)} merged = [] for entry in current_defs: - name = (entry.get("function") or {}).get("name", "") + name = _def_name(entry) replacement = fresh.pop(name, None) if replacement is not None: merged.append(replacement) elif name and name in registered_names: merged.append(entry) merged.extend(fresh.values()) - return merged, {(t.get("function") or {}).get("name", "") for t in merged} + return merged, {_def_name(t) for t in merged} def _reinject_post_build_tools(agent, tools_list: list, name_set: set) -> set: @@ -243,14 +230,15 @@ def _reinject_post_build_tools(agent, tools_list: list, name_set: set) -> set: Returns the context-engine routing names THIS rebuild appended: a name already owned by a registry/plugin tool is not claimed, matching agent_init. """ - def _add(schema: dict) -> bool: - name = schema.get("name", "") + def _add(schema) -> bool: + name = schema.get("name", "") if isinstance(schema, dict) else "" if not name or name in name_set: return False tools_list.append({"type": "function", "function": schema}) name_set.add(name) return True + enabled = getattr(agent, "enabled_toolsets", None) try: memory_manager = getattr(agent, "_memory_manager", None) get_mem_schemas = getattr(memory_manager, "get_all_tool_schemas", None) if memory_manager else None @@ -258,13 +246,10 @@ def _reinject_post_build_tools(agent, tools_list: list, name_set: set) -> set: # Same toolset gate inject_memory_provider_tools uses. from agent.memory_manager import memory_provider_tools_enabled if memory_provider_tools_enabled( - getattr(agent, "enabled_toolsets", None), - getattr(agent, "disabled_toolsets", None), - memory_tool_present="memory" in name_set, + enabled, getattr(agent, "disabled_toolsets", None), memory_tool_present="memory" in name_set, ): for schema in get_mem_schemas(): - if isinstance(schema, dict): - _add(schema) + _add(schema) except Exception: logger.debug("Memory-provider tool re-injection skipped", exc_info=True) @@ -273,18 +258,13 @@ def _reinject_post_build_tools(agent, tools_list: list, name_set: set) -> set: # restricted-toolset platform would re-leak tools the build excluded. staged_engine_names: set = set() try: - enabled = getattr(agent, "enabled_toolsets", None) - context_engine_allowed = enabled is None or "context_engine" in enabled compressor = getattr(agent, "context_compressor", None) get_schemas = getattr(compressor, "get_tool_schemas", None) if compressor else None - if context_engine_allowed and callable(get_schemas): + if (enabled is None or "context_engine" in enabled) and callable(get_schemas): for schema in get_schemas(): - if not isinstance(schema, dict): - continue - name = schema.get("name", "") # Claim the routing name only when WE appended the schema. - if _add(schema) and name: - staged_engine_names.add(name) + if _add(schema): + staged_engine_names.add(schema["name"]) except Exception: logger.debug("Context-engine tool re-injection skipped", exc_info=True) diff --git a/tools/mcp_tool_common.py b/tools/mcp_tool_common.py index d0bc4f3438..e127e9154f 100644 --- a/tools/mcp_tool_common.py +++ b/tools/mcp_tool_common.py @@ -44,9 +44,8 @@ def mcp_field(obj, snake: str, camel: str, default=None): either SDK generation (``mcp`` is an optional extra at the user's version). """ value = getattr(obj, snake, _MISSING) - if value is not _MISSING: - return value - value = getattr(obj, camel, _MISSING) + if value is _MISSING: + value = getattr(obj, camel, _MISSING) return default if value is _MISSING else value @@ -78,8 +77,7 @@ _BACKOFF_JITTER = 0.2 # +/-20% def _jittered(seconds: float) -> float: """``seconds`` with +/-20% uniform jitter, floored at 0.""" - return max(0.0, seconds * random.uniform(1.0 - _BACKOFF_JITTER, - 1.0 + _BACKOFF_JITTER)) + return max(0.0, seconds * random.uniform(1.0 - _BACKOFF_JITTER, 1.0 + _BACKOFF_JITTER)) # Credential patterns to strip from error messages. @@ -144,6 +142,10 @@ def _safe_numeric(value, default, coerce=int, minimum=1): return default +_TRUE_WORDS = frozenset({"true", "1", "yes", "on"}) +_FALSE_WORDS = frozenset({"false", "0", "no", "off"}) + + def _parse_boolish(value: Any, default: bool = True) -> bool: """Parse a bool-like config value with safe fallback.""" if value is None: @@ -152,16 +154,17 @@ def _parse_boolish(value: Any, default: bool = True) -> bool: return value if isinstance(value, str): lowered = value.strip().lower() - if lowered in {"true", "1", "yes", "on"}: + if lowered in _TRUE_WORDS: return True - if lowered in {"false", "0", "no", "off"}: + if lowered in _FALSE_WORDS: return False logger.warning("MCP config expected a boolean-ish value, got %r; using default=%s", value, default) return default def _get_lifecycle_seconds(config: dict, key: str) -> Optional[float]: - """Return an optional positive lifecycle timeout from top-level/nested config.""" + """Return an optional positive lifecycle timeout from top-level/nested 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): @@ -173,9 +176,7 @@ def _get_lifecycle_seconds(config: dict, key: str) -> Optional[float]: except (TypeError, ValueError): logger.warning("MCP config %s must be a number of seconds; ignoring %r", key, raw) return None - if seconds == 0: - return None if seconds < 0: logger.warning("MCP config %s must be positive; ignoring %r", key, raw) return None - return seconds + return seconds or None diff --git a/tools/mcp_tool_config.py b/tools/mcp_tool_config.py index 80fd0dff2b..c3205ac1cb 100644 --- a/tools/mcp_tool_config.py +++ b/tools/mcp_tool_config.py @@ -8,7 +8,6 @@ import re import shutil import sys import threading -from typing import Callable from datetime import datetime from typing import Any, Dict, List, Optional, Set, Tuple from tools.mcp_tool_common import _env_ref_name, _prepend_path, _core @@ -17,8 +16,6 @@ logger = logging.getLogger("tools.mcp_tool") _mcp_stderr_log_fh: Optional[Any] = None - - _mcp_stderr_log_lock = threading.Lock() @@ -36,10 +33,9 @@ def _get_mcp_stderr_log() -> Any: from hermes_constants import get_hermes_home log_dir = get_hermes_home() / "logs" log_dir.mkdir(parents=True, exist_ok=True) - log_path = log_dir / "mcp-stderr.log" # Line-buffered so output lands promptly; errors="replace" tolerates # garbled binary from misbehaving servers. - fh = open(log_path, "a", encoding="utf-8", errors="replace", buffering=1) + fh = open(log_dir / "mcp-stderr.log", "a", encoding="utf-8", errors="replace", buffering=1) fh.fileno() # confirm a real fd before committing _mcp_stderr_log_fh = fh except Exception as exc: # pragma: no cover — best-effort fallback @@ -64,44 +60,19 @@ def _write_stderr_log_header(server_name: str) -> None: # Env vars safe to pass to stdio subprocesses (no secrets). -_SAFE_ENV_KEYS = frozenset({ - "PATH", "HOME", "USER", "LANG", "LC_ALL", "TERM", "SHELL", "TMPDIR", -}) - +_SAFE_ENV_KEYS = frozenset({"PATH", "HOME", "USER", "LANG", "LC_ALL", "TERM", "SHELL", "TMPDIR"}) +# Windows process/location vars needed by launcher-style tools (e.g. Docker +# Desktop's MCP plugin discovery); none carry secrets. _SAFE_ENV_KEYS_CASE_INSENSITIVE = frozenset({ - # Windows process/location vars needed by launcher-style tools (e.g. - # Docker Desktop's MCP plugin discovery); none carry secrets. - "ALLUSERSPROFILE", - "APPDATA", - "COMMONPROGRAMFILES", - "COMMONPROGRAMFILES(X86)", - "COMMONPROGRAMW6432", - "COMPUTERNAME", - "COMSPEC", - "HOMEDRIVE", - "HOMEPATH", - "LOCALAPPDATA", - "NUMBER_OF_PROCESSORS", - "OS", - "PATHEXT", - "PROCESSOR_ARCHITECTURE", - "PROGRAMDATA", - "PROGRAMFILES", - "PROGRAMFILES(X86)", - "PROGRAMW6432", - "PUBLIC", - "SYSTEMDRIVE", - "SYSTEMROOT", - "TEMP", - "TMP", - "USERDOMAIN", - "USERNAME", - "USERPROFILE", - "WINDIR", + "ALLUSERSPROFILE", "APPDATA", "COMMONPROGRAMFILES", "COMMONPROGRAMFILES(X86)", + "COMMONPROGRAMW6432", "COMPUTERNAME", "COMSPEC", "HOMEDRIVE", "HOMEPATH", + "LOCALAPPDATA", "NUMBER_OF_PROCESSORS", "OS", "PATHEXT", "PROCESSOR_ARCHITECTURE", + "PROGRAMDATA", "PROGRAMFILES", "PROGRAMFILES(X86)", "PROGRAMW6432", "PUBLIC", + "SYSTEMDRIVE", "SYSTEMROOT", "TEMP", "TMP", "USERDOMAIN", "USERNAME", + "USERPROFILE", "WINDIR", }) - # ${VAR_NAME} interpolation; any non-} chars allowed so MY-VAR / my.var work. _ENV_VAR_PATTERN = re.compile(r"\$\{([^}]+)\}") @@ -120,20 +91,26 @@ def _workspace_folder() -> str: return os.getcwd() +def _workspace_basename() -> str: + root = _core._workspace_folder() + return os.path.basename(root.rstrip("/\\")) or root + + +# Cursor's case-sensitive context vars -> resolver. +_CONTEXT_VAR_RESOLVERS = { + "userHome": lambda: os.path.expanduser("~"), + "workspaceFolder": lambda: _core._workspace_folder(), + "workspaceFolderBasename": _workspace_basename, + "pathSeparator": lambda: os.sep, + "/": lambda: os.sep, +} + + def _context_var_value(ref: str) -> Optional[str]: - """Resolve Cursor's case-sensitive context vars (``userHome``, - ``workspaceFolder``, ``workspaceFolderBasename``, ``pathSeparator``/``/``). - Returns None for anything else so it falls through to env-var lookup.""" - if ref == "userHome": - return os.path.expanduser("~") - if ref == "workspaceFolder": - return _core._workspace_folder() - if ref == "workspaceFolderBasename": - root = _core._workspace_folder() - return os.path.basename(root.rstrip("/\\")) or root - if ref in ("pathSeparator", "/"): - return os.sep - return None + """Resolve a Cursor context var; None for anything else so it falls through + to env-var lookup.""" + resolver = _CONTEXT_VAR_RESOLVERS.get(ref) + return resolver() if resolver else None def _build_safe_env(user_env: Optional[dict]) -> dict: @@ -147,20 +124,59 @@ def _build_safe_env(user_env: Optional[dict]) -> dict: from hermes_cli.env_loader import get_secret_source except Exception: # pragma: no cover — early bootstrap/import fallback get_secret_source = None - env = {} - for key, value in os.environ.items(): - if ( - key in _SAFE_ENV_KEYS - or key.upper() in _SAFE_ENV_KEYS_CASE_INSENSITIVE - or key.startswith("XDG_") - or (get_secret_source is not None and get_secret_source(key)) - ): - env[key] = value + env = { + key: value + for key, value in os.environ.items() + if key in _SAFE_ENV_KEYS + or key.upper() in _SAFE_ENV_KEYS_CASE_INSENSITIVE + or key.startswith("XDG_") + or (get_secret_source is not None and get_secret_source(key)) + } if user_env: env.update(user_env) return env +def _which_with_config_pathext(command: str, path_arg, env: dict): + """``shutil.which`` retried under the config env's PATHEXT (Windows only): + ``which(path=...)`` uses the PARENT's PATHEXT, not the config env's.""" + cfg_pathext = next( + (v for k, v in env.items() if k.upper() == "PATHEXT" and isinstance(v, str) and v.strip()), + None, + ) + if not cfg_pathext or cfg_pathext == os.environ.get("PATHEXT"): + return None + saved = os.environ.get("PATHEXT") + try: + os.environ["PATHEXT"] = cfg_pathext + return shutil.which(command, path=path_arg) + finally: + if saved is None: + os.environ.pop("PATHEXT", None) + else: + os.environ["PATHEXT"] = saved + + +def _node_fallback(command: str) -> str: + """Well-known Node install locations for bare ``npx``/``npm``/``node`` when + PATH lookup failed; returns *command* unchanged when none is executable.""" + home = os.path.expanduser("~") + hermes_home = os.path.expanduser(os.getenv("HERMES_HOME", os.path.join(home, ".hermes"))) + candidates = [ + os.path.join(hermes_home, "node", "bin", command), + os.path.join(home, ".local", "bin", command), + # Canonical Node location for from-source Linux builds, the Hermes Docker + # image and Intel Homebrew. Needed when a user's hand-authored env.PATH + # omits it: npx's shebang re-execs /usr/bin/env node, so a symlink + # workaround fails one layer deeper. + os.path.join(os.sep, "usr", "local", "bin", command), + ] + for candidate in candidates: + if os.path.isfile(candidate) and os.access(candidate, os.X_OK): + return candidate + return command + + def _resolve_stdio_command(command: str, env: dict) -> tuple[str, dict]: """Resolve a stdio command against the exact subprocess env, mainly so bare ``npx``/``npm``/``node`` work under a filtered PATH.""" @@ -171,72 +187,31 @@ def _resolve_stdio_command(command: str, env: dict) -> tuple[str, dict]: path_arg = resolved_env.get("PATH") which_hit = shutil.which(resolved_command, path=path_arg) if which_hit is None and sys.platform == "win32" and resolved_env: - # shutil.which(path=...) uses the PARENT's PATHEXT, not the config - # env's, so retry with the config's PATHEXT (any key casing) applied. - cfg_pathext = next( - (v for k, v in resolved_env.items() - if k.upper() == "PATHEXT" and isinstance(v, str) and v.strip()), - None, - ) - if cfg_pathext and cfg_pathext != os.environ.get("PATHEXT"): - _saved = os.environ.get("PATHEXT") - try: - os.environ["PATHEXT"] = cfg_pathext - which_hit = shutil.which(resolved_command, path=path_arg) - finally: - if _saved is None: - os.environ.pop("PATHEXT", None) - else: - os.environ["PATHEXT"] = _saved + which_hit = _which_with_config_pathext(resolved_command, path_arg, resolved_env) if which_hit: resolved_command = which_hit elif resolved_command in {"npx", "npm", "node"}: - hermes_home = os.path.expanduser( - os.getenv( - "HERMES_HOME", os.path.join(os.path.expanduser("~"), ".hermes") - ) - ) - candidates = [ - os.path.join(hermes_home, "node", "bin", resolved_command), - os.path.join(os.path.expanduser("~"), ".local", "bin", resolved_command), - # Canonical Node location for from-source Linux builds, the - # Hermes Docker image and Intel Homebrew. Needed when a user's - # hand-authored env.PATH omits it: npx's shebang re-execs - # /usr/bin/env node, so a symlink workaround fails one layer deeper. - os.path.join(os.sep, "usr", "local", "bin", resolved_command), - ] - for candidate in candidates: - if os.path.isfile(candidate) and os.access(candidate, os.X_OK): - resolved_command = candidate - break + resolved_command = _node_fallback(resolved_command) command_dir = os.path.dirname(resolved_command) if command_dir: resolved_env = _prepend_path(resolved_env, command_dir) - return resolved_command, resolved_env def _wrap_command_with_watchdog(command: str, args: list) -> tuple[str, list]: - """Wrap a stdio command in the parent-death watchdog (POSIX only; the - watchdog polls ``getppid()`` against our PID). Unchanged on non-POSIX or - if the PID cannot be read — watchdog bookkeeping must never block a connection.""" + """Wrap a stdio command in the parent-death watchdog (POSIX only — it relies + on process groups, same scope as the killpg-based orphan cleanup; the + watchdog polls ``getppid()`` against our PID). Unchanged on non-POSIX or if + the PID cannot be read — watchdog bookkeeping must never block a connection.""" if os.name != "posix": - # Relies on process groups (getpgid/killpg), same scope as the - # killpg-based orphan cleanup. return command, args try: my_pid = os.getpid() except Exception: return command, args - watchdog_args = [ - os.path.join(os.path.dirname(os.path.abspath(__file__)), "mcp_stdio_watchdog.py"), - "--ppid", str(my_pid), - "--", - command, - *args, - ] - return sys.executable, watchdog_args + watchdog = os.path.join(os.path.dirname(os.path.abspath(__file__)), "mcp_stdio_watchdog.py") + return sys.executable, [watchdog, "--ppid", str(my_pid), "--", command, *args] def _interpolate_env_vars(value): @@ -254,8 +229,7 @@ def _interpolate_env_vars(value): ctx = _context_var_value(m.group(1).strip()) if ctx is not None: return ctx - name = _env_ref_name(m.group(1)) - return _get_secret(name, m.group(0)) or m.group(0) + return _get_secret(_env_ref_name(m.group(1)), m.group(0)) or m.group(0) return _ENV_VAR_PATTERN.sub(_replace, value) if isinstance(value, dict): return {k: _interpolate_env_vars(v) for k, v in value.items()} @@ -301,8 +275,7 @@ def _warn_hidden_whitespace(server_name: str, config: dict) -> List[str]: "trailing whitespace — this often causes authentication or " "connection failures. Check for stray spaces/newlines in " "config.yaml (or the referenced env var).", - server_name, - key_path, + server_name, key_path, ) return flagged @@ -310,30 +283,37 @@ def _warn_hidden_whitespace(server_name: str, config: dict) -> List[str]: def _filter_suspicious_mcp_servers(servers: Dict[str, dict]) -> Dict[str, dict]: """Drop exfiltration-shaped MCP configs before any stdio spawn path.""" try: - from hermes_cli.mcp_security import validate_mcp_server_entry as _validate_mcp_server_entry + from hermes_cli.mcp_security import validate_mcp_server_entry except Exception: - _validate_mcp_server_entry: Callable[[str, dict[str, Any]], list[str]] | None = None - - if _validate_mcp_server_entry is None: return servers safe_servers = {} for name, cfg in servers.items(): - if not isinstance(cfg, dict): - safe_servers[name] = cfg - continue - issues = _validate_mcp_server_entry(name, cfg) + issues = validate_mcp_server_entry(name, cfg) if isinstance(cfg, dict) else None if issues: - logger.warning( - "Skipping suspicious MCP server '%s': %s", - name, - "; ".join(issues), - ) + logger.warning("Skipping suspicious MCP server '%s': %s", name, "; ".join(issues)) continue safe_servers[name] = cfg return safe_servers +def _portable_mcp_servers(safe_servers: Dict[str, dict]) -> None: + """Merge plugin-provided (portable) MCP servers into *safe_servers*; native + config wins on a name clash. Never raises.""" + try: + from hermes_cli.plugins import discover_plugins, get_plugin_manager + + discover_plugins() + portable = get_plugin_manager().get_portable_mcp_servers() + for name, cfg in _core._filter_suspicious_mcp_servers(portable).items(): + if name in safe_servers: + logger.warning("Portable MCP server '%s' conflicts with native config; skipping", name) + continue + safe_servers[name] = dict(cfg) + except Exception: + logger.debug("Failed to load portable MCP servers", exc_info=True) + + def _load_mcp_config() -> Dict[str, dict]: """Read ``mcp_servers`` from config.yaml as ``{name: config}`` (empty on error or in safe mode). Entries carry ``command``/``args``/``env`` (stdio) or @@ -345,8 +325,7 @@ def _load_mcp_config() -> Dict[str, dict]: if _env_enabled("HERMES_SAFE_MODE"): return {} - config = load_config() - servers = config.get("mcp_servers") + servers = load_config().get("mcp_servers") if not isinstance(servers, dict): servers = {} # Ensure .env vars are available for interpolation @@ -361,21 +340,7 @@ def _load_mcp_config() -> Dict[str, dict]: if isinstance(interpolated, dict): _warn_hidden_whitespace(name, interpolated) safe_servers[name] = interpolated - try: - from hermes_cli.plugins import discover_plugins, get_plugin_manager - - discover_plugins() - portable = get_plugin_manager().get_portable_mcp_servers() - for name, cfg in _core._filter_suspicious_mcp_servers(portable).items(): - if name in safe_servers: - logger.warning( - "Portable MCP server '%s' conflicts with native config; skipping", - name, - ) - continue - safe_servers[name] = dict(cfg) - except Exception: - logger.debug("Failed to load portable MCP servers", exc_info=True) + _portable_mcp_servers(safe_servers) return safe_servers except Exception as exc: logger.debug("Failed to load MCP config: %s", exc) diff --git a/tools/mcp_tool_lifecycle.py b/tools/mcp_tool_lifecycle.py index 6be36a732e..023179edc8 100644 --- a/tools/mcp_tool_lifecycle.py +++ b/tools/mcp_tool_lifecycle.py @@ -15,16 +15,12 @@ logger = logging.getLogger("tools.mcp_tool") # removed on normal shutdown, so they can be force-killed if SDK teardown fails. _stdio_pids: Dict[int, str] = {} - # PIDs that survived their session context exit (SDK teardown failed to kill # them); detected in _run_stdio's finally, reaped by _kill_orphaned_mcp_children(). # Kept separate from _stdio_pids so cleanup sweeps never race active sessions. _orphan_stdio_pids: set = set() - - _orphan_stdio_pid_servers: Dict[int, str] = {} - # pid -> pgid captured at spawn. The SDK spawns children with # start_new_session=True (PGID == PID); grandchildren inherit that PGID and # keep it after the direct child exits, so killpg still reaches them. Tracked @@ -96,16 +92,19 @@ def _filter_mcp_children(pids: set) -> set: except (psutil.NoSuchProcess, psutil.AccessDenied, OSError): # Raced away or zombie — cannot be our fresh server, unsafe to track. continue - if any( - marker in arg - for arg in argv[1:] - for marker in _NON_MCP_CHILD_CMDLINE_MARKERS - ): + if any(marker in arg for arg in argv[1:] for marker in _NON_MCP_CHILD_CMDLINE_MARKERS): continue filtered.add(pid) return filtered +def _clear_connect_cooldowns() -> None: + """Drop connect-retry cooldowns: a restart must re-attempt every server + immediately, not honour a stale per-server backoff. Caller holds ``_core._lock``.""" + _core._server_connect_retry_after.clear() + _core._server_connect_failures.clear() + + def shutdown_mcp_servers(*, scope: Optional[str] = None): """Close MCP server connections (in parallel) and stop the background loop. @@ -127,8 +126,7 @@ def shutdown_mcp_servers(*, scope: Optional[str] = None): # most likely state for stale backoff entries; a restart must retry at once. if not servers_snapshot: with _core._lock: - _core._server_connect_retry_after.clear() - _core._server_connect_failures.clear() + _clear_connect_cooldowns() _core._stop_mcp_loop(only_if_idle=scope is not None) return @@ -139,26 +137,19 @@ def shutdown_mcp_servers(*, scope: Optional[str] = None): ) for server, result in zip(servers_snapshot, results): if isinstance(result, Exception): - logger.debug( - "Error closing MCP server '%s': %s", server.name, result, - ) + logger.debug("Error closing MCP server '%s': %s", server.name, result) with _core._lock: for name in selected: _core._servers.pop(name, None) _core._server_scope_keys.pop(name, None) - # Drop connect-retry cooldowns too: a restart must re-attempt every - # server immediately, not honour a stale per-server backoff. - _core._server_connect_retry_after.clear() - _core._server_connect_failures.clear() + _clear_connect_cooldowns() with _core._lock: loop = _core._mcp_loop if loop is not None and loop.is_running(): from agent.async_utils import safe_schedule_threadsafe future = safe_schedule_threadsafe( - _shutdown(), loop, - logger=logger, - log_message="MCP shutdown: failed to schedule", + _shutdown(), loop, logger=logger, log_message="MCP shutdown: failed to schedule", ) if future is not None: try: @@ -169,16 +160,63 @@ def shutdown_mcp_servers(*, scope: Optional[str] = None): # Unconditional final sweep: whether ``_shutdown`` ran, timed out, or was # never scheduled, no stale connect-cooldown state may survive shutdown. with _core._lock: - _core._server_connect_retry_after.clear() - _core._server_connect_failures.clear() - + _clear_connect_cooldowns() _core._stop_mcp_loop(only_if_idle=scope is not None) -def _kill_orphaned_mcp_children( - include_active: bool = False, - server_name: Optional[str] = None, -) -> None: +def _take_reapable_pids(include_active: bool, server_name: Optional[str]) -> tuple[Dict[int, str], Dict[int, int]]: + """Pop the PIDs to reap (and their spawn-time pgids) out of the ledgers under + the lock, so a future spawn can't collide with stale state. + Returns ``(pid -> owner, pid -> pgid)``.""" + def _owned(entries: Dict[int, str]) -> Dict[int, str]: + return {pid: owner for pid, owner in entries.items() if server_name is None or owner == server_name} + + with _core._lock: + pids = _owned({opid: _orphan_stdio_pid_servers.get(opid, "orphan") for opid in _orphan_stdio_pids}) + for opid in pids: + _orphan_stdio_pids.discard(opid) + _orphan_stdio_pid_servers.pop(opid, None) + if include_active: + active = _owned(_stdio_pids) + pids.update(active) + for pid in active: + _stdio_pids.pop(pid, None) + pgids = {pid: _stdio_pgids.pop(pid) for pid in pids if pid in _stdio_pgids} + return pids, pgids + + +def _signal_mcp_process(pid: int, sig: int, server_name: str, pgid: Optional[int], my_pgid: Optional[int]) -> None: + """SIGTERM/SIGKILL via the spawn-time pgroup on POSIX (reaches reparented + grandchildren), falling back to a per-pid signal.""" + killpg = getattr(os, "killpg", None) + if pgid is not None and killpg is not None: + if my_pgid is not None and pgid == my_pgid: + # Child shares the gateway's pgroup: killpg would kill the gateway + # too, so use per-pid kill. Warn because per-pid kill can't reach + # grandchildren in this group (inherent trade-off). + logger.warning( + "MCP server '%s' pgid %d matches gateway pgid; skipping " + "killpg to avoid self-kill and using per-pid kill — any " + "grandchildren in this group may not be reaped", + server_name, pgid, + ) + else: + try: + killpg(pgid, sig) + return + except (ProcessLookupError, PermissionError, OSError) as exc: + # Pgroup gone or refused — still try the direct child. + logger.debug( + "killpg(%d, %d) failed for MCP server '%s': %s; falling back to kill(pid)", + pgid, sig, server_name, exc, + ) + try: + os.kill(pid, sig) + except (ProcessLookupError, PermissionError, OSError): + pass + + +def _kill_orphaned_mcp_children(include_active: bool = False, server_name: Optional[str] = None) -> None: """Best-effort reap of stdio MCP subprocesses: SIGTERM, wait 2s, SIGKILL survivors. By default only ``_orphan_stdio_pids`` (PIDs that outlived their session @@ -186,97 +224,35 @@ def _kill_orphaned_mcp_children( ``include_active=True`` also kills every ``_stdio_pids`` entry and is only for final shutdown after the MCP loop has stopped. ``server_name`` limits the sweep to one server (stdio reconnects cleaning up their old transport). - - On POSIX signals go via ``os.killpg`` to the spawn-time pgid when tracked, - so reparented grandchildren are reaped too; falls back to ``os.kill``. """ import signal as _signal - with _core._lock: - pids: Dict[int, str] = {} - for opid in _orphan_stdio_pids: - owner = _orphan_stdio_pid_servers.get(opid, "orphan") - if server_name is not None and owner != server_name: - continue - pids[opid] = owner - for opid in pids: - _orphan_stdio_pids.discard(opid) - _orphan_stdio_pid_servers.pop(opid, None) - if include_active: - active = dict(_stdio_pids) - if server_name is not None: - active = { - pid: owner - for pid, owner in active.items() - if owner == server_name - } - pids.update(active) - for pid in active: - _stdio_pids.pop(pid, None) - # Snapshot pgids for the pids we're about to kill, then drop them so a - # future spawn can't collide with stale state. - pgids: Dict[int, int] = {pid: _stdio_pgids[pid] for pid in pids if pid in _stdio_pgids} - for pid in pgids: - _stdio_pgids.pop(pid, None) - + pids, pgids = _take_reapable_pids(include_active, server_name) # Fast path: nothing to reap — skip the 2s sleep every MCP-free shutdown # would otherwise pay. if not pids: return - # Our own pgid, so _send_signal never killpg()s the gateway itself. + # Our own pgid, so we never killpg() the gateway itself. try: - _my_pgid = os.getpgrp() + my_pgid = os.getpgrp() except (AttributeError, OSError): - _my_pgid = None # Windows or restricted environment + my_pgid = None # Windows or restricted environment - def _send_signal(pid: int, sig: int, server_name: str) -> None: - """SIGTERM/SIGKILL via pgroup on POSIX, fall back to pid signal.""" - pgid = pgids.get(pid) - killpg = getattr(os, "killpg", None) - if pgid is not None and killpg is not None: - if _my_pgid is not None and pgid == _my_pgid: - # Child shares the gateway's pgroup: killpg would kill the - # gateway too, so use per-pid kill. Warn because per-pid kill - # can't reach grandchildren in this group (inherent trade-off). - logger.warning( - "MCP server '%s' pgid %d matches gateway pgid; skipping " - "killpg to avoid self-kill and using per-pid kill — any " - "grandchildren in this group may not be reaped", - server_name, pgid, - ) - else: - try: - killpg(pgid, sig) - return - except (ProcessLookupError, PermissionError, OSError) as exc: - # Pgroup gone or refused — still try the direct child. - logger.debug( - "killpg(%d, %d) failed for MCP server '%s': %s; falling back to kill(pid)", - pgid, sig, server_name, exc, - ) - try: - os.kill(pid, sig) - except (ProcessLookupError, PermissionError, OSError): - pass - - for pid, server_name in pids.items(): - _send_signal(pid, _signal.SIGTERM, server_name) - logger.debug("Sent SIGTERM to orphaned MCP process %d (%s)", pid, server_name) + for pid, owner in pids.items(): + _signal_mcp_process(pid, _signal.SIGTERM, owner, pgids.get(pid), my_pgid) + logger.debug("Sent SIGTERM to orphaned MCP process %d (%s)", pid, owner) time.sleep(2) - _sigkill = getattr(_signal, "SIGKILL", _signal.SIGTERM) + sigkill = getattr(_signal, "SIGKILL", _signal.SIGTERM) # ``os.kill(pid, 0)`` is NOT a no-op on Windows; use the portable check. from gateway.status import _pid_exists - for pid, server_name in pids.items(): + for pid, owner in pids.items(): if not _pid_exists(pid): continue # exited after SIGTERM - _send_signal(pid, _sigkill, server_name) - logger.warning( - "Force-killed MCP process %d (%s) after SIGTERM timeout", - pid, server_name, - ) + _signal_mcp_process(pid, sigkill, owner, pgids.get(pid), my_pgid) + logger.warning("Force-killed MCP process %d (%s) after SIGTERM timeout", pid, owner) def _stop_mcp_loop_if_idle() -> bool: @@ -290,10 +266,7 @@ def _stop_mcp_loop_if_idle() -> bool: return _core._stop_mcp_loop(only_if_idle=True) -async def _drain_mcp_loop_tasks( - *, - timeout: Optional[float] = None, -) -> None: +async def _drain_mcp_loop_tasks(*, timeout: Optional[float] = None) -> None: """Cancel every task still pending on the MCP loop and reap it. ``Task.cancel()`` only schedules the throw, so tasks need a cancellation @@ -322,10 +295,7 @@ async def _drain_mcp_loop_tasks( logger.debug("Pending MCP loop task ended during shutdown: %s", exc) if still_pending: - logger.warning( - "%d MCP loop task(s) still pending after %.1fs drain", - len(still_pending), timeout, - ) + logger.warning("%d MCP loop task(s) still pending after %.1fs drain", len(still_pending), timeout) async def _drain_and_stop_mcp_loop() -> None: diff --git a/tools/mcp_tool_schema.py b/tools/mcp_tool_schema.py index db8f43472b..9afe0ad953 100644 --- a/tools/mcp_tool_schema.py +++ b/tools/mcp_tool_schema.py @@ -15,144 +15,113 @@ 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. _MCP_INJECTION_PATTERNS = [ - (re.compile(r"ignore\s+(all\s+)?previous\s+instructions", re.I), - "prompt override attempt ('ignore previous instructions')"), - (re.compile(r"you\s+are\s+now\s+a", re.I), - "identity override attempt ('you are now a...')"), - (re.compile(r"your\s+new\s+(task|role|instructions?)\s+(is|are)", re.I), - "task override attempt"), - (re.compile(r"system\s*:\s*", re.I), - "system prompt injection attempt"), - (re.compile(r"<\s*(system|human|assistant)\s*>", re.I), - "role tag injection attempt"), - (re.compile(r"do\s+not\s+(tell|inform|mention|reveal)", re.I), - "concealment instruction"), - (re.compile(r"(curl|wget|fetch)\s+https?://", re.I), - "network command in description"), - (re.compile(r"base64\.(b64decode|decodebytes)", re.I), - "base64 decode reference"), - (re.compile(r"exec\s*\(|eval\s*\(", re.I), - "code execution reference"), - (re.compile(r"import\s+(subprocess|os|shutil|socket)", re.I), - "dangerous import reference"), + (re.compile(pattern, re.I), reason) + for pattern, reason in ( + (r"ignore\s+(all\s+)?previous\s+instructions", "prompt override attempt ('ignore previous instructions')"), + (r"you\s+are\s+now\s+a", "identity override attempt ('you are now a...')"), + (r"your\s+new\s+(task|role|instructions?)\s+(is|are)", "task override attempt"), + (r"system\s*:\s*", "system prompt injection attempt"), + (r"<\s*(system|human|assistant)\s*>", "role tag injection attempt"), + (r"do\s+not\s+(tell|inform|mention|reveal)", "concealment instruction"), + (r"(curl|wget|fetch)\s+https?://", "network command in description"), + (r"base64\.(b64decode|decodebytes)", "base64 decode reference"), + (r"exec\s*\(|eval\s*\(", "code execution reference"), + (r"import\s+(subprocess|os|shutil|socket)", "dangerous import reference"), + ) ] 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.""" - findings = [] if not description: - return findings - for pattern, reason in _MCP_INJECTION_PATTERNS: - if pattern.search(description): - findings.append(reason) + 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, + "MCP server '%s' tool '%s': suspicious description content — %s. Description: %.200s", + server_name, tool_name, "; ".join(findings), description, ) return findings +_EMPTY_OBJECT_SCHEMA = {"type": "object", "properties": {}} + + +def _rewrite_local_refs(node): + """Promote legacy ``definitions`` to ``$defs`` (Moonshot rejects the draft-07 + form) — but ONLY where it is a JSON Schema meta-keyword, never as a property + NAME inside ``properties``/``patternProperties``. A tool parameter + legitimately named ``definitions`` rewritten to ``$defs`` would 400 the whole + tool array (Anthropic/OpenAI forbid ``$`` in property names). Property names + are kept verbatim and recursion resumes ordinary semantics inside each + property's schema.""" + if isinstance(node, list): + return [_rewrite_local_refs(item) for item in node] + if not isinstance(node, dict): + return node + normalized = {} + for key, value in node.items(): + if key in ("properties", "patternProperties") and isinstance(value, dict): + normalized[key] = {name: _rewrite_local_refs(schema) for name, schema in value.items()} + else: + normalized["$defs" if key == "definitions" else key] = _rewrite_local_refs(value) + ref = normalized.get("$ref") + if isinstance(ref, str) and ref.startswith("#/definitions/"): + normalized["$ref"] = "#/$defs/" + ref[len("#/definitions/"):] + return normalized + + +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).""" + if isinstance(node, list): + return [_repair_object_shape(item) for item in node] + if not isinstance(node, dict): + return node + repaired = {k: _repair_object_shape(v) for k, v in node.items()} + if not repaired.get("type") and ("properties" in repaired or "required" in repaired): + repaired["type"] = "object" + if repaired.get("type") == "object": + if not isinstance(repaired.get("properties"), dict): + repaired["properties"] = {} + required = repaired.get("required") + if isinstance(required, list): + props = repaired.get("properties") or {} + valid = [r for r in required if isinstance(r, str) and r in props] + if len(valid) != len(required): + if valid: + repaired["required"] = valid + else: + repaired.pop("required", None) + return repaired + + def _normalize_mcp_input_schema(schema: dict | None) -> dict: """Normalize MCP input schemas so one form is valid on OpenAI, Anthropic, Gemini and Moonshot. - Repairs, applied recursively: ``definitions``/``#/definitions/`` refs -> - ``$defs`` (Moonshot rejects the draft-07 form); missing/null ``type`` on an - object-shaped node -> ``"object"``; an object without ``properties`` gets an - empty one so ``required`` can't dangle; ``required`` pruned to names present - in ``properties`` (Gemini 400s otherwise); nullable ``anyOf`` unions - collapsed to the non-null branch (Anthropic rejects nullable branches), - optionality living solely in the parent's ``required``. + Order matters: ``definitions`` -> ``$defs`` rewrite; nullable ``anyOf`` + unions collapsed to the non-null branch (Anthropic rejects nullable + branches; optionality lives solely in the parent's ``required``, the + ``nullable: true`` hint is kept so runtime coercion can still map a + model-emitted ``"null"`` string to ``None``); same-typed const unions -> + enum (must run AFTER the nullable strip); then object-shape repair. """ if not schema: - return {"type": "object", "properties": {}} - - def _rewrite_local_refs(node): - """Promote legacy ``definitions`` to ``$defs`` — but ONLY where it is a - JSON Schema meta-keyword, never as a property NAME inside ``properties``/ - ``patternProperties``. A tool parameter legitimately named ``definitions`` - rewritten to ``$defs`` would 400 the whole tool array (Anthropic/OpenAI - forbid ``$`` in property names). Property names are kept verbatim and - recursion resumes ordinary semantics inside each property's schema.""" - if isinstance(node, dict): - normalized = {} - for key, value in node.items(): - if key in ("properties", "patternProperties") and isinstance(value, dict): - normalized[key] = { - prop_name: _rewrite_local_refs(prop_schema) - for prop_name, prop_schema in value.items() - } - else: - out_key = "$defs" if key == "definitions" else key - normalized[out_key] = _rewrite_local_refs(value) - ref = normalized.get("$ref") - if isinstance(ref, str) and ref.startswith("#/definitions/"): - normalized["$ref"] = "#/$defs/" + ref[len("#/definitions/"):] - return normalized - if isinstance(node, list): - return [_rewrite_local_refs(item) for item in node] - return node - - def _strip_nullable_union(node): - """Shared implementation with the Anthropic guard and global sanitizer. - Keeps the ``nullable: true`` hint so runtime coercion can still map a - model-emitted ``"null"`` string to ``None``.""" - from tools.schema_sanitizer import strip_nullable_unions - - return strip_nullable_unions(node, keep_nullable_hint=True) - - def _collapse_const_unions(node): - """Collapse anyOf/oneOf unions of same-typed consts to enums. Must run - AFTER the nullable strip: consts -> enum, null branch -> ``nullable`` hint.""" - from tools.schema_sanitizer import collapse_const_unions - - return collapse_const_unions(node) - - def _repair_object_shape(node): - """Recursively fill missing object ``type``, ensure ``properties``, prune ``required``.""" - if isinstance(node, list): - return [_repair_object_shape(item) for item in node] - if not isinstance(node, dict): - return node - - repaired = {k: _repair_object_shape(v) for k, v in node.items()} - - if not repaired.get("type") and ( - "properties" in repaired or "required" in repaired - ): - repaired["type"] = "object" - - if repaired.get("type") == "object": - if not isinstance(repaired.get("properties"), dict): - repaired["properties"] = {} - - required = repaired.get("required") - if isinstance(required, list): - props = repaired.get("properties") or {} - valid = [r for r in required if isinstance(r, str) and r in props] - if len(valid) != len(required): - if valid: - repaired["required"] = valid - else: - repaired.pop("required", None) - - return repaired + return dict(_EMPTY_OBJECT_SCHEMA) + from tools.schema_sanitizer import collapse_const_unions, strip_nullable_unions normalized = _rewrite_local_refs(schema) - normalized = _strip_nullable_union(normalized) - normalized = _collapse_const_unions(normalized) + 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 {"type": "object", "properties": {}} + return dict(_EMPTY_OBJECT_SCHEMA) if normalized.get("type") == "object" and "properties" not in normalized: normalized = {**normalized, "properties": {}} - return normalized @@ -166,16 +135,12 @@ def sanitize_mcp_name_component(value: str) -> str: # 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 = "__" def mcp_prefixed_tool_name(server_name: str, tool_name: str) -> str: """Registry/wire name: ``mcp____``.""" - safe_server = sanitize_mcp_name_component(server_name) - safe_tool = sanitize_mcp_name_component(tool_name) - return f"{MCP_TOOL_NAME_PREFIX}{safe_server}{_MCP_NAME_DELIM}{safe_tool}" + return f"{MCP_TOOL_NAME_PREFIX}{sanitize_mcp_name_component(server_name)}{_MCP_NAME_DELIM}{sanitize_mcp_name_component(tool_name)}" def _convert_mcp_schema(server_name: str, mcp_tool) -> dict: @@ -183,81 +148,48 @@ def _convert_mcp_schema(server_name: str, mcp_tool) -> dict: 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}" - ), - "parameters": _normalize_mcp_input_schema( - mcp_field(mcp_tool, "input_schema", "inputSchema") - ), + "description": strip_unicode_tags(mcp_tool.description or f"MCP tool {mcp_tool.name} from {server_name}"), + "parameters": _normalize_mcp_input_schema(mcp_field(mcp_tool, "input_schema", "inputSchema")), } +# 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}'", + {"uri": {"type": "string", "description": "URI of the resource to read"}}, ["uri"]), + ("list_prompts", "List available prompts from MCP server '{server}'", {}, None), + ("get_prompt", "Get a prompt by name from MCP server '{server}'", + { + "name": {"type": "string", "description": "Name of the prompt to retrieve"}, + "arguments": { + "type": "object", + "description": "Optional arguments to pass to the prompt", + "properties": {}, + "additionalProperties": True, + }, + }, ["name"]), +) + + def _build_utility_schemas(server_name: str) -> List[dict]: """Schemas for the resource/prompt utility tools as ``{schema, handler_key}`` dicts.""" - return [ - { + out = [] + for handler_key, description, properties, required in _UTILITY_TOOL_SPECS: + parameters = {"type": "object", "properties": {k: dict(v) for k, v in properties.items()}} + if required is not None: + parameters["required"] = list(required) + out.append({ "schema": { - "name": mcp_prefixed_tool_name(server_name, "list_resources"), - "description": f"List available resources from MCP server '{server_name}'", - "parameters": { - "type": "object", - "properties": {}, - }, + "name": mcp_prefixed_tool_name(server_name, handler_key), + "description": description.format(server=server_name), + "parameters": parameters, }, - "handler_key": "list_resources", - }, - { - "schema": { - "name": mcp_prefixed_tool_name(server_name, "read_resource"), - "description": f"Read a resource by URI from MCP server '{server_name}'", - "parameters": { - "type": "object", - "properties": { - "uri": { - "type": "string", - "description": "URI of the resource to read", - }, - }, - "required": ["uri"], - }, - }, - "handler_key": "read_resource", - }, - { - "schema": { - "name": mcp_prefixed_tool_name(server_name, "list_prompts"), - "description": f"List available prompts from MCP server '{server_name}'", - "parameters": { - "type": "object", - "properties": {}, - }, - }, - "handler_key": "list_prompts", - }, - { - "schema": { - "name": mcp_prefixed_tool_name(server_name, "get_prompt"), - "description": f"Get a prompt by name from MCP server '{server_name}'", - "parameters": { - "type": "object", - "properties": { - "name": { - "type": "string", - "description": "Name of the prompt to retrieve", - }, - "arguments": { - "type": "object", - "description": "Optional arguments to pass to the prompt", - "properties": {}, - "additionalProperties": True, - }, - }, - "required": ["name"], - }, - }, - "handler_key": "get_prompt", - }, - ] + "handler_key": handler_key, + }) + return out def _normalize_name_filter(value: Any, label: str) -> set[str]: @@ -280,20 +212,12 @@ def matches_name_filter(tool_name: str, patterns: set[str]) -> bool: return False if tool_name in patterns: return True - return any( - fnmatch.fnmatchcase(tool_name, p) - for p in patterns - if "*" in p or "?" in p or "[" in p - ) + return any(fnmatch.fnmatchcase(tool_name, p) for p in patterns if "*" in p or "?" in p or "[" in p) -_UTILITY_CAPABILITY_METHODS = { - "list_resources": "list_resources", - "read_resource": "read_resource", - "list_prompts": "list_prompts", - "get_prompt": "get_prompt", -} - +# 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