refactor(mcp): split orphan reaper/config/schema/agent-refresh helpers; table-build utility schemas and injection patterns

This commit is contained in:
Teknium
2026-09-02 15:48:34 -07:00
parent d6ec8c3018
commit faceb45337
7 changed files with 397 additions and 573 deletions

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@@ -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__<sanitizedServer>__<sanitizedTool>``."""
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