refactor(tools/mcp): split mcp_tool.py into transport/lifecycle/schema/handlers/... sibling modules; compact watchdog and schema cache

This commit is contained in:
Teknium
2026-09-02 13:55:31 -07:00
parent 606cb2de92
commit c1f8af1e86
17 changed files with 5832 additions and 7322 deletions
+1 -9
View File
@@ -43,31 +43,23 @@ class TestCacheRoundTrip:
assert entry is not None
assert msc.tools_from_cache_entry(entry) == tools
assert msc.utility_tools_from_cache_entry(entry) == []
assert msc.has_cached_entry("srv", "fp1")
def test_fingerprint_mismatch_returns_none(self, monkeypatch, tmp_path):
self._isolate(monkeypatch, tmp_path)
msc.write_cache_entry("srv", "fp1", tools=[], utility_tools=[])
assert msc.get_cached_entry("srv", "OTHER") is None
assert not msc.has_cached_entry("srv", "OTHER")
def test_missing_server_returns_none(self, monkeypatch, tmp_path):
self._isolate(monkeypatch, tmp_path)
assert msc.get_cached_entry("nope", "fp") is None
def test_clear_cache_entry(self, monkeypatch, tmp_path):
self._isolate(monkeypatch, tmp_path)
msc.write_cache_entry("srv", "fp1", tools=[], utility_tools=[])
msc.clear_cache_entry("srv")
assert msc.get_cached_entry("srv", "fp1") is None
def test_corrupt_cache_file_is_tolerated(self, monkeypatch, tmp_path):
self._isolate(monkeypatch, tmp_path)
(tmp_path / "cache.json").write_text("{not json", encoding="utf-8")
assert msc.get_cached_entry("srv", "fp") is None
# And writes recover the file.
msc.write_cache_entry("srv", "fp", tools=[], utility_tools=[])
assert msc.has_cached_entry("srv", "fp")
assert msc.get_cached_entry("srv", "fp") is not None
def test_malformed_entry_shapes_are_tolerated(self):
assert msc.tools_from_cache_entry({"tools": "nope"}) == []
+2 -2
View File
@@ -105,9 +105,9 @@ def test_watch_ok_probe_does_not_create_unawaited_coroutine():
"""
import inspect as _inspect
import tools.mcp_tool as mcp_mod
import tools.mcp_tool_handlers as handlers_mod
src = _inspect.getsource(mcp_mod)
src = _inspect.getsource(handlers_mod)
assert "isawaitable(_watch_children())" not in src
assert "iscoroutinefunction(_watch_children)" in src
+6 -15
View File
@@ -83,16 +83,15 @@ def get_cached_entry(server_name: str, fingerprint: str) -> Optional[dict]:
return None
ttl_ms = entry.get("ttl_ms")
written_at = entry.get("written_at")
if isinstance(ttl_ms, (int, float)) and isinstance(written_at, (int, float)):
if (time.time() - written_at) * 1000.0 >= float(ttl_ms):
return None
if (
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
def has_cached_entry(server_name: str, fingerprint: str) -> bool:
return get_cached_entry(server_name, fingerprint) is not None
def write_cache_entry(
server_name: str,
fingerprint: str,
@@ -132,14 +131,6 @@ def write_cache_entry(
_save_all(data)
def clear_cache_entry(server_name: str) -> None:
with _cache_lock:
data = _load_all()
if server_name in data:
del data[server_name]
_save_all(data)
def tools_from_cache_entry(entry: dict) -> List[dict]:
"""Return cached MCP tool dicts (name, description, inputSchema)."""
tools = entry.get("tools")
+21 -47
View File
@@ -1,38 +1,19 @@
#!/usr/bin/env python3
"""Parent-death watchdog supervisor for stdio MCP subprocesses.
Problem this fixes (#TBD): a stdio MCP server (e.g. ``npx -y mcp-remote
<url>``) is spawned as a direct child of the Hermes process. Hermes's own
teardown path (``MCPServerTask.shutdown()`` / ``_kill_orphaned_mcp_children``
at final exit) reaps it cleanly on a *graceful* exit. But if the spawning
Hermes process dies hard — ``kill -9``, an OS-level crash, a force-quit of
the TUI/desktop app — that teardown code never runs, and the child (plus any
of its own descendants, e.g. mcp-remote's spawned ``node`` process) is
orphaned. macOS has no direct equivalent of Linux's
``prctl(PR_SET_PDEATHSIG)`` to make the kernel auto-kill a child when its
parent dies, so nothing reaps these until the next Hermes startup's opt-in
``_kill_orphaned_mcp_children()`` sweep — which only runs if something calls
it. Repeated ungraceful session restarts can pile up N orphaned processes,
all racing to hold the same upstream SSE session, producing errors like
"Invalid request parameters" / "Received request before initialization was
complete" on the *legitimate* new connection.
If Hermes dies hard (kill -9, crash, force-quit), its graceful teardown never
runs and the stdio child plus its descendants are orphaned; macOS has no
``PR_SET_PDEATHSIG`` equivalent. Piled-up orphans then race the legitimate new
connection for the same upstream session. So the MCP command is spawned via
this supervisor, which:
1. runs the real command as its own child in a new process group (so the
whole tree can be killpg'd);
2. passes stdin/stdout/stderr straight through — the MCP stdio protocol
talks over those pipes, so this must be a no-op relay, not a proxy;
3. polls ``getppid()`` against the recorded parent PID and, when the parent
is gone, SIGTERMs the child's group, waits, then SIGKILLs.
Fix: don't spawn the MCP server command directly. Spawn this supervisor
instead, which:
1. execs the real command as its own child (own process group via
``start_new_session``, so it doesn't inherit the supervisor's
controlling terminal weirdly and so we can killpg it cleanly);
2. transparently passes stdin/stdout/stderr through — the MCP stdio
protocol talks directly over those pipes, so the supervisor must be a
no-op relay, not a bytes-in-the-middle proxy;
3. runs a background thread that polls the direct POSIX parent identity:
compare current ``getppid()`` against the parent PID recorded when the
wrapper was created;
4. the instant the original parent is gone, terminates the real child's
process group (SIGTERM, grace period, then SIGKILL) and exits.
This is intentionally a thin, standard-library-only script so it starts fast
and can't itself become a resource leak.
Standard-library only so it starts fast and cannot itself leak.
Usage (see ``tools/mcp_tool.py::_run_stdio``)::
@@ -60,13 +41,9 @@ def _is_orphaned(original_ppid: int, getppid=os.getppid) -> bool:
def _terminate_process_group(proc: subprocess.Popen) -> None:
"""Best-effort SIGTERM-then-SIGKILL of the child's process group.
This module only ever runs on POSIX (the wrap site in tools/mcp_tool.py
gates on ``os.name == "posix"``), but guard the POSIX-only primitives
anyway so an accidental Windows import/execute degrades to a plain
child kill instead of AttributeError.
"""
"""Best-effort SIGTERM-then-SIGKILL of the child's process group. Only runs
on POSIX in practice, but guards the POSIX-only primitives so an accidental
Windows run degrades to a plain child kill instead of AttributeError."""
killpg = getattr(os, "killpg", None)
if killpg is None: # windows-footgun: ok — non-POSIX fallback
try:
@@ -115,9 +92,8 @@ def main(argv: list[str] | None = None) -> int:
print("mcp_stdio_watchdog: no command given after '--'", file=sys.stderr)
return 2
# New process group so we can killpg() the whole tree the real command
# may spawn (e.g. mcp-remote's own child `node` process), without
# touching our own group or the (already-gone) original parent's.
# 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,
@@ -126,12 +102,10 @@ def main(argv: list[str] | None = None) -> int:
start_new_session=True,
)
# Because the real server lives in its OWN process group (above), the
# parent's graceful-shutdown killpg of *our* group no longer reaches it.
# Forward SIGTERM/SIGINT to the child's group so graceful teardown
# (`_kill_orphaned_mcp_children`, shutdown sweeps) still kills a wedged
# server that ignores stdin EOF — otherwise the watchdog wrap would
# invert the bug it fixes.
# 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
# group so graceful teardown still kills a wedged server that ignores stdin
# EOF — otherwise the wrap would invert the bug it fixes.
def _forward_shutdown(signum, frame): # noqa: ARG001
_terminate_process_group(proc)
sys.exit(128 + signum)
+593 -7249
View File
File diff suppressed because it is too large Load Diff
+291
View File
@@ -0,0 +1,291 @@
"""Live-agent tool-list maintenance after MCP (re)discovery: refreshing an
AIAgent's tools/tool names, preserving the cached tools[] prefix across rebuilds,
and re-injecting post-build tools."""
import logging
import json
import threading
from tools.mcp_tool_common import _core
logger = logging.getLogger("tools.mcp_tool")
# Serializes in-place swaps of ``agent.tools`` / ``agent.valid_tool_names`` by
# the reload RPC, gateway reload and late-binding refresh thread; the run loop
# reads them during tool iteration and must never see a half-updated pair.
_agent_tools_lock = threading.Lock()
def refresh_agent_mcp_tools(
agent,
*,
enabled_override=None,
disabled_override=None,
quiet_mode: bool = True,
content_aware: bool = False,
preserve_prefix: bool = False,
) -> set:
"""Re-derive an already-built agent's tool snapshot from the live registry.
The agent snapshots ``agent.tools`` once at build time; servers that connect
later (slow OAuth server, ``/reload-mcp``) are invisible until rebuilt. This
is the single shared rebuild for the TUI RPC, gateway reload, late-binding
thread and between-turns refresh. It respects the agent's own toolset
filter, diffs by tool NAME (a count compare misses an equal-size swap), and
is additive-preserving: memory-provider and context-engine (``lcm_*``) tools
that ``agent_init`` appends after ``get_tool_definitions`` are re-injected,
since a naive rebuild would silently delete them. ``(tools, valid_tool_names)``
are published together under ``_agent_tools_lock``.
``preserve_prefix`` is for rebuilds inside a live conversation, where the
tool array is a cached request prefix and any moved byte re-prefills the
whole history: existing tools keep their slot (schemas still refresh), a
still-registered tool whose ``check_fn`` merely flapped is carried forward,
a tool that left the registry is dropped, and new tools append at the tail.
Carrying an unavailable tool forward is safe: ``check_fn`` gates exposure,
never invocation, and every handler owns its own unavailability error.
Returns the newly-added tool names (empty when unchanged). The caller owns
the prompt-cache contract (turn-boundary policy differs per caller).
"""
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)
# 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
# newer set another caller already published.
snapshot_generation = registry._generation
# 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}
# Re-append the post-build families on LOCALS only; live agent attributes
# are untouched until the single atomic publish below.
staged_engine_names = _core._reinject_post_build_tools(agent, new_defs, new_names)
# Registry membership is read OUTSIDE ``_agent_tools_lock``: taking
# ``registry._lock`` under the tools lock would be the first nesting of the two.
registered_names: set = set()
if preserve_prefix:
try:
registered_names = {entry.name for entry in registry.get_all_entries()}
except Exception: # noqa: BLE001
preserve_prefix = False # fail open to the plain rebuild
# Single atomic read-diff-publish so ``added`` matches what was published
# and a stale (older-generation) rebuild can't overwrite a newer one.
with _agent_tools_lock:
# Tolerate an agent that never set the generation (or a non-int mock)
# rather than failing the whole refresh on the comparison.
published_gen_raw = getattr(agent, "_tool_snapshot_generation", -1)
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}
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()
agent.tools = new_defs
agent.valid_tool_names = new_names
# Publish context-engine routing names atomically with the snapshot.
engine_names = getattr(agent, "_context_engine_tool_names", None)
if isinstance(engine_names, set):
engine_names.clear()
engine_names.update(staged_engine_names)
agent._tool_snapshot_generation = max(published_gen, snapshot_generation)
added = new_names - current
# Re-pin the session's tool order so a rebuild-for-existing-session
# (gateway agent-cache eviction) restores exactly these names.
persist_agent_tool_names(agent)
return added
def reprobe_tool_availability() -> None:
"""Explicit ``/reload-mcp`` hatch out of the tools[] freeze: drop the
``check_fn`` verdict cache AND the ``get_tool_definitions`` memo (keyed on
registry generation, so it would otherwise replay the stale verdicts)."""
from model_tools import _clear_tool_defs_cache
from tools.registry import invalidate_check_fn_cache
invalidate_check_fn_cache()
_clear_tool_defs_cache()
def persist_agent_tool_names(agent) -> None:
"""Best-effort: write ``agent.tools`` names to the session row (freeze pin)."""
db = getattr(agent, "_session_db", None)
session_id = getattr(agent, "session_id", 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 [])],
)
except Exception: # noqa: BLE001
logger.debug("tool_names persist skipped", exc_info=True)
def restore_agent_tool_prefix(agent, saved_names: list) -> bool:
"""Fold a freshly built agent's ``tools`` onto the session's saved order.
The gateway rebuilds a NEW AIAgent for an existing session after agent-cache
eviction, with no predecessor to preserve; the saved name list stands in:
a saved tool still registered but failing its probe is carried forward from
the registry schema, a deregistered one is dropped, new tools append at the
tail (same rule as ``_merge_preserving_prefix``). Returns True if changed.
"""
if not saved_names:
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}
saved_defs = []
for name in saved_names:
entry_def = fresh.get(name)
if entry_def is None:
entry = registry.get_entry(name)
if entry is None:
continue
entry_def = {"type": "function", "function": {**entry.schema, "name": entry.name}}
saved_defs.append(entry_def)
registered_names = {entry.name for entry in registry.get_all_entries()}
merged, merged_names = _merge_preserving_prefix(saved_defs, fresh_defs, registered_names)
with _agent_tools_lock:
if merged == fresh_defs:
return False
agent.tools = merged
agent.valid_tool_names = merged_names
if [t["function"]["name"] 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]:
"""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
keeps its slot but takes the fresh schema; a name only in the live list is
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
merged = []
for entry in current_defs:
name = (entry.get("function") or {}).get("name", "")
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}
def _reinject_post_build_tools(agent, tools_list: list, name_set: set) -> set:
"""Append memory-provider and context-engine tools onto the caller's staged
``tools_list`` / ``name_set`` (never the live agent attributes), mirroring
``agent_init``'s post-build injection. Idempotent and fail-soft.
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", "")
if not name or name in name_set:
return False
tools_list.append({"type": "function", "function": schema})
name_set.add(name)
return True
try:
memory_manager = getattr(agent, "_memory_manager", None)
get_mem_schemas = getattr(memory_manager, "get_all_tool_schemas", None) if memory_manager else None
if callable(get_mem_schemas):
# 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,
):
for schema in get_mem_schemas():
if isinstance(schema, dict):
_add(schema)
except Exception:
logger.debug("Memory-provider tool re-injection skipped", exc_info=True)
# The `context_engine` toolset is intentionally empty, so lcm_* tools exist
# only via this append. Honor the enabled_toolsets gate agent_init uses, or a
# 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):
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)
except Exception:
logger.debug("Context-engine tool re-injection skipped", exc_info=True)
return staged_engine_names
+181
View File
@@ -0,0 +1,181 @@
"""Small pure helpers shared by the tools.mcp_tool_* modules: SDK 1.x/2.x field
access, error-text sanitising, numeric/bool coercion, timeouts and jitter. No
origin state."""
import logging
import math
import os
import random
import re
from typing import Any, Optional
logger = logging.getLogger("tools.mcp_tool")
class _OriginProxy:
"""Attribute proxy for ``tools.mcp_tool`` resolved at access time.
The split modules read origin state (``_servers``, ``_lock``, SDK symbols,
patchable helpers) through this so ``mock.patch("tools.mcp_tool.X")`` and
origin-side rebinds stay effective, and so no split module needs the origin
imported first (the origin imports them while it is still initialising).
"""
__slots__ = ()
def __getattr__(self, name: str):
from tools import mcp_tool
return getattr(mcp_tool, name)
_core = _OriginProxy()
_MISSING = object()
def mcp_field(obj, snake: str, camel: str, default=None):
"""Read an MCP model field across the 1.x -> 2.x rename to snake_case.
Pydantic aliases don't apply to attribute access, so ``getattr(result,
"isError", False)`` silently returns the default on 2.x — failed calls read
as successful, schemas as empty. Trying both spellings stays correct on
either SDK generation (``mcp`` is an optional extra at the user's version).
"""
value = getattr(obj, snake, _MISSING)
if value is not _MISSING:
return value
value = getattr(obj, camel, _MISSING)
return default if value is _MISSING else value
_DEFAULT_TOOL_TIMEOUT = 300 # seconds for tool calls
def _resolve_tool_timeout(config: dict) -> float:
"""Per-server tool-call timeout. Precedence: ``mcp_servers.<name>.timeout``
> ``timeouts.mcp.tool_call`` > the 300s default; values are platform-clamped
by ``resolve_timeout``."""
per_server = config.get("timeout")
if per_server is not None:
return per_server
try:
from agent.deadline import resolve_timeout
resolved = resolve_timeout("mcp.tool_call", default=_DEFAULT_TOOL_TIMEOUT)
if resolved is not None:
return resolved
except Exception:
logger.debug("mcp.tool_call timeout resolution failed", exc_info=True)
return _DEFAULT_TOOL_TIMEOUT
# Jitter on reconnect backoff so servers that lost the same backend don't
# retry in lockstep (thundering herd, synchronized log bursts).
_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))
# Credential patterns to strip from error messages.
_CREDENTIAL_PATTERN = re.compile(
r"(?:"
r"ghp_[A-Za-z0-9_]{1,255}" # GitHub PAT
r"|sk-[A-Za-z0-9_]{1,255}" # OpenAI-style key
r"|Bearer\s+\S+" # Bearer token
r"|token=[^\s&,;\"']{1,255}" # token=...
r"|key=[^\s&,;\"']{1,255}" # key=...
r"|API_KEY=[^\s&,;\"']{1,255}" # API_KEY=...
r"|password=[^\s&,;\"']{1,255}" # password=...
r"|secret=[^\s&,;\"']{1,255}" # secret=...
r")",
re.IGNORECASE,
)
def _env_ref_name(ref: str) -> str:
"""Bare env-var name from a ``${...}`` body; strips a Cursor-style ``env:`` prefix."""
ref = ref.strip()
if ref.startswith("env:"):
ref = ref[len("env:"):].strip()
return ref
def _sanitize_error(text: str) -> str:
"""Replace credential-like patterns with [REDACTED] before text reaches the LLM."""
return _CREDENTIAL_PATTERN.sub("[REDACTED]", text)
def _exc_str(exc: BaseException) -> str:
"""Non-empty string for *exc*: some exceptions (``anyio.ClosedResourceError``)
carry no message, so fall back to ``repr`` to keep diagnostics."""
text = str(exc).strip()
return text or repr(exc)
def _prepend_path(env: dict, directory: str) -> dict:
"""Prepend *directory* to env PATH if it is not already present."""
updated = dict(env or {})
if not directory:
return updated
existing = updated.get("PATH", "")
parts = [part for part in existing.split(os.pathsep) if part]
if directory not in parts:
parts = [directory, *parts]
updated["PATH"] = os.pathsep.join(parts) if parts else directory
return updated
def _safe_numeric(value, default, coerce=int, minimum=1):
"""Coerce a config value (YAML strings included) to a number, clamped to
*minimum*; *default* on failure or non-finite floats."""
try:
result = coerce(value)
if isinstance(result, float) and not math.isfinite(result):
return default
return max(result, minimum)
except (TypeError, ValueError, OverflowError):
return default
def _parse_boolish(value: Any, default: bool = True) -> bool:
"""Parse a bool-like config value with safe fallback."""
if value is None:
return default
if isinstance(value, bool):
return value
if isinstance(value, str):
lowered = value.strip().lower()
if lowered in {"true", "1", "yes", "on"}:
return True
if lowered in {"false", "0", "no", "off"}:
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."""
raw = config.get(key)
lifecycle = config.get("lifecycle")
if raw is None and isinstance(lifecycle, dict):
raw = lifecycle.get(key)
if raw is None:
return None
try:
seconds = float(raw)
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
+382
View File
@@ -0,0 +1,382 @@
"""MCP server config loading and stdio launch environment: ${VAR}/Cursor-style
interpolation, hidden-whitespace and suspicious-entry filtering, the filtered
subprocess env, command resolution, watchdog wrapping and the shared stderr log."""
import logging
import os
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
logger = logging.getLogger("tools.mcp_tool")
_mcp_stderr_log_fh: Optional[Any] = None
_mcp_stderr_log_lock = threading.Lock()
def _get_mcp_stderr_log() -> Any:
"""Shared append-mode handle for MCP subprocess stderr, opened once per process.
Must expose a real fd (``fileno()``) because asyncio wires the child's stderr
directly to it. Falls back to ``/dev/null``, then the real stderr.
"""
global _mcp_stderr_log_fh
with _mcp_stderr_log_lock:
if _mcp_stderr_log_fh is not None:
return _mcp_stderr_log_fh
try:
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.fileno() # confirm a real fd before committing
_mcp_stderr_log_fh = fh
except Exception as exc: # pragma: no cover — best-effort fallback
logger.debug("Failed to open MCP stderr log, using devnull: %s", exc)
try:
_mcp_stderr_log_fh = open(os.devnull, "w", encoding="utf-8")
except Exception:
_mcp_stderr_log_fh = sys.stderr
return _mcp_stderr_log_fh
def _write_stderr_log_header(server_name: str) -> None:
"""Write a session marker so operators can find each server's output in the
shared log without per-line prefixes (which would need a pipe + reader thread)."""
fh = _core._get_mcp_stderr_log()
try:
ts = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
fh.write(f"\n===== [{ts}] starting MCP server '{server_name}' =====\n")
fh.flush()
except Exception:
pass
# 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_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",
})
# ${VAR_NAME} interpolation; any non-} chars allowed so MY-VAR / my.var work.
_ENV_VAR_PATTERN = re.compile(r"\$\{([^}]+)\}")
def _workspace_folder() -> str:
"""Absolute workspace root for ``${workspaceFolder}``: the session's
authoritative root (terminal cwd / task override / $TERMINAL_CWD), else cwd."""
try:
from tools.file_tools import _authoritative_workspace_root
root = _authoritative_workspace_root()
if root:
return root
except Exception:
pass
return os.getcwd()
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
def _build_safe_env(user_env: Optional[dict]) -> dict:
"""Filtered env for stdio subprocesses so API keys/tokens don't leak.
Passes only the safe baseline keys, ``XDG_*``, vars injected by an external
secret source (users configured that backend precisely so subprocesses can
consume them), plus the server config's own ``env``.
"""
try:
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
if user_env:
env.update(user_env)
return env
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."""
resolved_command = os.path.expanduser(str(command).strip())
resolved_env = dict(env or {})
if os.sep not in resolved_command:
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
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
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."""
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
def _interpolate_env_vars(value):
"""Recursively resolve ``${VAR}`` / Cursor ``${env:VAR}`` placeholders plus the
Cursor context vars (see ``_context_var_value``).
Env refs resolve from the active profile's secret scope when multiplexing
(so ``${API_KEY}`` picks up the routed profile's value, not another
profile's in ``os.environ``). Unset vars keep the literal placeholder.
"""
from agent.secret_scope import get_secret as _get_secret
if isinstance(value, str):
def _replace(m):
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 _ENV_VAR_PATTERN.sub(_replace, value)
if isinstance(value, dict):
return {k: _interpolate_env_vars(v) for k, v in value.items()}
if isinstance(value, list):
return [_interpolate_env_vars(v) for v in value]
return value
# (server_name, dotted key path) pairs already warned about; config loads
# happen on every discovery pass, so warn once per process.
_whitespace_warned: Set[Tuple[str, str]] = set()
def _warn_hidden_whitespace(server_name: str, config: dict) -> List[str]:
"""Warn once per (server, key path) about string values with leading/trailing
whitespace — a pasted trailing newline or leading space causes opaque
auth/connect failures and is invisible in config.yaml.
Advisory only: values are never mutated (whitespace could be intentional)
and never logged (often secrets). Returns the flagged key paths.
"""
flagged: List[str] = []
def _walk(value: Any, path: str) -> None:
if isinstance(value, str):
if value != value.strip():
flagged.append(path)
elif isinstance(value, dict):
for k, v in value.items():
_walk(v, f"{path}.{k}" if path else str(k))
elif isinstance(value, list):
for i, v in enumerate(value):
_walk(v, f"{path}[{i}]")
_walk(config, "")
for key_path in flagged:
dedupe_key = (server_name, key_path)
if dedupe_key in _whitespace_warned:
continue
_whitespace_warned.add(dedupe_key)
logger.warning(
"MCP server '%s': config value '%s' has hidden leading or "
"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,
)
return flagged
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
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)
if issues:
logger.warning(
"Skipping suspicious MCP server '%s': %s",
name,
"; ".join(issues),
)
continue
safe_servers[name] = cfg
return safe_servers
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
``url``/``headers`` (HTTP) plus optional timeout/auth keys; ``${VAR}``
placeholders are interpolated after ``.env`` is loaded."""
try:
from hermes_cli.config import load_config
from utils import env_var_enabled as _env_enabled
if _env_enabled("HERMES_SAFE_MODE"):
return {}
config = load_config()
servers = config.get("mcp_servers")
if not isinstance(servers, dict):
servers = {}
# Ensure .env vars are available for interpolation
try:
from hermes_cli.env_loader import load_hermes_dotenv
load_hermes_dotenv()
except Exception:
pass
safe_servers: Dict[str, dict] = {}
for name, cfg in _core._filter_suspicious_mcp_servers(servers).items():
interpolated = _interpolate_env_vars(cfg)
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)
return safe_servers
except Exception as exc:
logger.debug("Failed to load MCP config: %s", exc)
return {}
+254
View File
@@ -0,0 +1,254 @@
"""Rendering of MCP tool-result content blocks into model-facing text: size
capping, _meta filtering, image/audio caching to MEDIA tags, resource links and
embedded resources."""
import logging
from tools.ansi_strip import strip_unicode_tags
from tools.mcp_tool_common import mcp_field, _core
from tools.mcp_tool_schema import mcp_prefixed_tool_name
logger = logging.getLogger("tools.mcp_tool")
# Hard allocation ceiling for one MCP text payload (chars): the first line of
# defense against a multi-megabyte flood being JSON-encoded and handed
# downstream. Deliberately far ABOVE the budget layer's 50K spillover threshold
# so ordinary large results reach spillover intact; only pathological floods
# are lossy-truncated here.
_MCP_HARD_RESULT_CAP_CHARS = 2_000_000
def _truncate_mcp_text_result(text: str, max_chars: int = _MCP_HARD_RESULT_CAP_CHARS) -> str:
"""Pass text at or under ``max_chars`` unchanged; otherwise keep a 40% head /
60% tail split with an omission notice between."""
if len(text) <= max_chars:
return text
head_chars = int(max_chars * 0.4)
tail_chars = max_chars - head_chars
omitted = len(text) - head_chars - tail_chars
return (
text[:head_chars]
+ f"\n\n... [MCP RESULT TRUNCATED - {omitted:,} chars omitted "
f"out of {len(text):,} total] ...\n\n"
+ text[-tail_chars:]
)
def _is_reserved_mcp_meta_key(key: str) -> bool:
"""True if an MCP ``_meta`` key uses a protocol-reserved prefix: a
``modelcontextprotocol`` or ``mcp`` label followed by at least one more
label. A trailing one (``com.example.mcp/...``) is a vendor namespace."""
slash = key.find("/")
if slash <= 0:
return False
labels = key[:slash].split(".")
return any(
label in ("modelcontextprotocol", "mcp") and i < len(labels) - 1
for i, label in enumerate(labels)
)
def _strip_reserved_meta_keys(meta) -> "Optional[Dict[str, Any]]":
"""Drop protocol-reserved keys from ``_meta``; None if nothing model-facing
remains or the input wasn't a mapping."""
if not isinstance(meta, dict):
return None
out = {k: v for k, v in meta.items()
if isinstance(k, str) and not _is_reserved_mcp_meta_key(k)}
return out or None
def _mcp_image_extension_for_mime_type(mime_type: str) -> str:
"""File extension for an MCP image MIME type (``.png`` fallback)."""
import mimetypes
normalized = (mime_type or "").split(";", 1)[0].strip().lower()
if normalized in {"image/jpeg", "image/jpg"}:
return ".jpg"
return mimetypes.guess_extension(normalized) or ".png"
def _cache_mcp_image_block(block) -> str:
"""Cache an ``ImageContent`` block and return a ``MEDIA:<path>`` tag.
Returns "" (logging, not raising) when the block isn't an image, the base64
is malformed, or the cache rejects the bytes: one bad block must not kill
the tool result, and the caller falls through to any text blocks.
"""
import base64
data = getattr(block, "data", None)
mime_type = mcp_field(block, "mime_type", "mimeType")
normalized_mime = str(mime_type or "").split(";", 1)[0].strip().lower()
if data is None or not normalized_mime.startswith("image/"):
return ""
try:
raw_bytes = base64.b64decode(data)
except (TypeError, ValueError) as exc:
logger.warning("MCP image block decode failed (%s): %s", normalized_mime, exc)
return ""
try:
from gateway.platforms.base import cache_image_from_bytes
image_path = cache_image_from_bytes(
raw_bytes,
ext=_mcp_image_extension_for_mime_type(normalized_mime),
)
except ImportError:
# gateway.platforms.base unavailable (e.g. cron without gateway deps):
# drop silently, callers get any text blocks that parsed.
logger.debug("MCP image caching skipped — gateway.platforms.base unavailable")
return ""
except Exception as exc:
logger.warning("MCP image block cache failed: %s", exc)
return ""
return f"MEDIA:{image_path}"
# Hard cap on decoded resource bytes from one block, so a misbehaving server
# can't fill the cache disk.
_MCP_RESOURCE_MAX_BYTES = 50 * 1024 * 1024
# Base64 expands ~4/3; reject oversized payloads BEFORE decoding so a multi-GB
# blob string is never transiently doubled in memory.
_MCP_RESOURCE_MAX_B64_CHARS = _MCP_RESOURCE_MAX_BYTES * 4 // 3 + 4
def _mcp_resource_filename(uri: str, mime_type: str) -> str:
"""Safe display filename from the URI's last path segment, used only as a
name hint: ``cache_document_from_bytes`` re-sanitizes and prefixes it, so
remote path components can't steer the cache location."""
import mimetypes
import re as _re
from pathlib import Path
from urllib.parse import urlparse, unquote
name = ""
if uri:
try:
name = Path(unquote(urlparse(str(uri)).path or "")).name
except (ValueError, TypeError):
name = ""
# Strip control chars (hostile URIs could inject newlines/ANSI into the
# filename and transcript marker) and cap length, preserving the extension.
name = _re.sub(r"[\x00-\x1f\x7f]", "", name).strip()
if len(name) > 150:
stem, dot, ext = name.rpartition(".")
if dot and 0 < len(ext) <= 12:
name = stem[: 150 - len(ext) - 1] + "." + ext
else:
name = name[:150]
if not name or name in {".", ".."}:
normalized = (mime_type or "").split(";", 1)[0].strip().lower()
ext = mimetypes.guess_extension(normalized) or ".bin"
name = f"resource{ext}"
return name
def _cache_mcp_audio_block(block) -> str:
"""Cache an ``AudioContent`` block and return a ``MEDIA:`` tag; "" when not
audio or on any failure (same fail-open contract as the image path)."""
import base64
data = getattr(block, "data", None)
mime_type = str(mcp_field(block, "mime_type", "mimeType") or "").split(";", 1)[0].strip().lower()
if data is None or not mime_type.startswith("audio/"):
return ""
if len(data) > _core._MCP_RESOURCE_MAX_B64_CHARS:
return f"[MCP audio resource too large to cache: ~{len(data) * 3 // 4} bytes]"
try:
raw_bytes = base64.b64decode(data)
except (TypeError, ValueError) as exc:
logger.warning("MCP audio block decode failed (%s): %s", mime_type, exc)
return ""
if len(raw_bytes) > _core._MCP_RESOURCE_MAX_BYTES:
return f"[MCP audio resource too large to cache: {len(raw_bytes)} bytes]"
try:
from gateway.platforms.base import cache_audio_from_bytes
import mimetypes
ext = (
{"audio/wav": ".wav", "audio/x-wav": ".wav", "audio/wave": ".wav"}.get(mime_type)
or mimetypes.guess_extension(mime_type)
or ".ogg"
)
audio_path = cache_audio_from_bytes(raw_bytes, ext=ext)
except ImportError:
logger.debug("MCP audio caching skipped — gateway.platforms.base unavailable")
return ""
except Exception as exc:
logger.warning("MCP audio block cache failed: %s", exc)
return ""
return f"MEDIA:{audio_path}"
def _render_mcp_resource_block(block, server_name: str = "") -> str:
"""Render a ``ResourceLink`` or ``EmbeddedResource`` block as text.
Embedded text → the text; embedded blob → decoded (size-capped) into the
document cache with a path marker; link → the URI plus a pointer at the
server's read_resource tool (no fetch here — links are only readable via
the originating session). "" for non-resource blocks; failures are
reported inline rather than silently dropped.
"""
block_type = getattr(block, "type", "")
if block_type == "resource_link" or (
hasattr(block, "uri") and not hasattr(block, "resource") and block_type != "text"
):
uri = getattr(block, "uri", None)
if not uri:
return ""
name = getattr(block, "name", "") or ""
mime = mcp_field(block, "mime_type", "mimeType", "") or ""
details = f"uri={uri}"
if name:
details += f", name={name}"
if mime:
details += f", mimeType={mime}"
reader = (
mcp_prefixed_tool_name(server_name, "read_resource")
if server_name
else "the MCP server's read_resource tool"
)
return f"[MCP resource link: {details} — fetch it with {reader}]"
resource = getattr(block, "resource", None)
if resource is None:
return ""
text = getattr(resource, "text", None)
if text is not None:
return strip_unicode_tags(str(text))
blob = getattr(resource, "blob", None)
if blob is None:
return ""
import base64
uri = str(getattr(resource, "uri", "") or "")
mime = str(mcp_field(resource, "mime_type", "mimeType", "") or "")
if len(blob) > _core._MCP_RESOURCE_MAX_B64_CHARS:
return f"[MCP embedded resource too large to cache: ~{len(blob) * 3 // 4} bytes, uri={uri}]"
try:
raw_bytes = base64.b64decode(blob)
except (TypeError, ValueError) as exc:
logger.warning("MCP embedded resource decode failed (%s): %s", mime or uri, exc)
return f"[MCP embedded resource could not be decoded: {mime or uri}]"
if len(raw_bytes) > _core._MCP_RESOURCE_MAX_BYTES:
return f"[MCP embedded resource too large to cache: {len(raw_bytes)} bytes, uri={uri}]"
try:
from gateway.platforms.base import cache_document_from_bytes
path = cache_document_from_bytes(raw_bytes, _mcp_resource_filename(uri, mime))
except ImportError:
logger.debug("MCP resource caching skipped — gateway.platforms.base unavailable")
return f"[MCP embedded resource received ({len(raw_bytes)} bytes, {mime or 'unknown type'}) but document cache unavailable in this process]"
except Exception as exc:
logger.warning("MCP embedded resource cache failed: %s", exc)
return f"[MCP embedded resource could not be cached: {mime or uri}]"
return f"[MCP resource saved to {path} ({mime or 'unknown type'}, {len(raw_bytes)} bytes) — read it with read_file or terminal tools]"
+541
View File
@@ -0,0 +1,541 @@
"""MCP connection/transport error classification: URL validation, TLS client certs, identity headers, redirect header stripping, exception-group unwrapping, auth/session-expired/method-not-found detection and connect-error formatting. Split from tools/mcp_tool.py."""
import logging
import asyncio
import errno
import os
import re
from typing import Any, List, Optional
from urllib.parse import urlparse
from tools.mcp_tool_common import _sanitize_error, _core
logger = logging.getLogger("tools.mcp_tool")
# Stateless (2026-07-28) servers reject a legacy ``initialize`` with
# UnsupportedProtocolVersion (-32022) or plain method-not-found.
_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION = -32022
def _handshake_rejected_as_modern(exc: BaseException) -> bool:
"""True when a failed ``initialize`` signals a stateless-only (2026-07-28) server.
Structural code check first, then substring fallback — never ``isinstance`` on
SDK exception types (they arrive wrapped in ExceptionGroups and drift across generations).
"""
err = getattr(exc, "error", None)
code = getattr(err, "code", None) or getattr(exc, "code", None)
if code in (_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, _core._JSONRPC_METHOD_NOT_FOUND):
return True
msg = str(exc).lower()
if not msg:
return False
return (
"unsupported protocol version" in msg
or str(_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION) in msg
or _is_method_not_found_error(exc)
)
def _is_method_not_found_error(exc: BaseException) -> bool:
"""True if *exc* is a JSON-RPC ``method not found`` (-32601).
``ping`` is optional in MCP; servers lacking it answer -32601. Structural
``MCPError.error.code`` check first, then substring fallback — including
"Unknown method: <name>", which some servers use; without it the
ping→list_tools keepalive fallback never latches and reconnect-loops.
"""
err = getattr(exc, "error", None)
code = getattr(err, "code", None)
if code == _core._JSONRPC_METHOD_NOT_FOUND:
return True
msg = str(exc).lower()
if not msg:
return False
return (
str(_core._JSONRPC_METHOD_NOT_FOUND) in msg
or "method not found" in msg
or "unknown method" in msg
or "not found: ping" in msg
)
class InvalidMcpUrlError(ValueError):
"""A remote MCP server's ``url`` is not parseable http(s)://.
Validated once at startup so we fail fast instead of burning the
reconnect-backoff loop on every attempt.
"""
class NonMcpEndpointError(ConnectionError):
"""An HTTP MCP URL served a non-MCP 2xx response (e.g. ``text/html``).
Real Streamable-HTTP endpoints answer ``application/json`` or
``text/event-stream``. Non-retryable: every attempt gets the same page, so
the backoff loop is skipped and the server is failed immediately.
Subclasses ConnectionError so broad catches still see a connection problem.
"""
def _unwrap_exception_group(exc: BaseException) -> BaseException:
"""Extract the root-cause leaf from anyio ``(Base)ExceptionGroup`` wrappers.
Group ``str()`` is opaque ("unhandled errors in a TaskGroup"), so log sites
must unwrap. Two rules: a ``KeyboardInterrupt``/``SystemExit`` leaf anywhere
is re-raised (never flattened into a loggable error); a non-cancellation
leaf is preferred over the ``CancelledError`` noise anyio sprays on siblings.
"""
while isinstance(exc, BaseExceptionGroup) and exc.exceptions:
fatal, _rest = exc.split((KeyboardInterrupt, SystemExit))
if fatal is not None:
leaf: BaseException = fatal
while isinstance(leaf, BaseExceptionGroup) and leaf.exceptions:
leaf = leaf.exceptions[0]
raise leaf
chosen = exc.exceptions[0]
for sub in exc.exceptions:
if not _contains_only_cancellation(sub):
chosen = sub
break
exc = chosen
return exc
def _contains_only_cancellation(exc: BaseException) -> bool:
"""True if ``exc`` is (or a group containing only) CancelledError."""
if isinstance(exc, BaseExceptionGroup):
return all(_contains_only_cancellation(sub) for sub in exc.exceptions)
return isinstance(exc, asyncio.CancelledError)
def _classify_mcp_failure(exc: BaseException) -> str:
"""Classify a connection failure as ``'permanent'`` or ``'transient'``.
Permanent (deterministic — ``run()`` parks immediately instead of burning the
retry ladder): auth 401/403, NonMcpEndpointError, InvalidMcpUrlError, missing
stdio command (FileNotFoundError / ENOENT). Everything else keeps backoff retry.
"""
root = _unwrap_exception_group(exc)
if _core._is_auth_error(root):
return "permanent"
if isinstance(root, (NonMcpEndpointError, InvalidMcpUrlError)):
return "permanent"
if isinstance(root, FileNotFoundError):
return "permanent"
if isinstance(root, OSError) and getattr(root, "errno", None) == errno.ENOENT:
return "permanent"
# 401/403 HTTPStatusError that _is_auth_error's type-gate missed
# (auth types not importable in this environment).
status = getattr(getattr(root, "response", None), "status_code", None)
if status in (401, 403):
return "permanent"
return "transient"
def _validate_remote_mcp_url(server_name: str, url: Any) -> str:
"""Return the stripped URL if it is a valid http(s) remote MCP URL.
Raises InvalidMcpUrlError naming the server for non-strings, missing/other
schemes (stdio servers use ``command``, not ``url``), and empty hosts.
"""
if not isinstance(url, str):
raise InvalidMcpUrlError(
f"Invalid MCP URL for '{server_name}': expected a string, got "
f"{type(url).__name__}"
)
stripped = url.strip()
if not stripped:
raise InvalidMcpUrlError(
f"Invalid MCP URL for '{server_name}': empty url"
)
try:
parsed = urlparse(stripped)
except Exception as exc: # urlparse is very permissive — belt and braces
raise InvalidMcpUrlError(
f"Invalid MCP URL for '{server_name}': {stripped!r} ({exc})"
) from exc
if parsed.scheme.lower() not in {"http", "https"}:
raise InvalidMcpUrlError(
f"Invalid MCP URL for '{server_name}': scheme must be http or "
f"https, got {parsed.scheme!r} ({stripped!r})"
)
if not parsed.netloc:
raise InvalidMcpUrlError(
f"Invalid MCP URL for '{server_name}': missing host ({stripped!r})"
)
# ``urlparse`` accepts ``http://:8080`` (empty host, explicit port) — reject it.
if not parsed.hostname:
raise InvalidMcpUrlError(
f"Invalid MCP URL for '{server_name}': missing hostname "
f"({stripped!r})"
)
return stripped
def _resolve_client_cert(server_name: str, config: dict):
"""Resolve ``client_cert`` / ``client_key`` into httpx's ``cert=`` shape.
None when neither is set; a single path for a combined PEM; ``(cert, key)``
or ``(cert, key, password)`` for the pair/list forms. ``~`` is expanded and
missing files raise a server-scoped FileNotFoundError instead of an opaque
TLS handshake error.
"""
raw_cert = config.get("client_cert")
raw_key = config.get("client_key")
if raw_cert is None and raw_key is None:
return None
def _expand(path: Any, label: str) -> str:
if not isinstance(path, str) or not path.strip():
raise ValueError(
f"MCP server '{server_name}': {label} must be a non-empty "
f"string path (got {type(path).__name__})"
)
expanded = os.path.expanduser(path.strip())
if not os.path.isfile(expanded):
raise FileNotFoundError(
f"MCP server '{server_name}': {label} not found at "
f"{expanded!r}"
)
return expanded
if isinstance(raw_cert, (list, tuple)):
if raw_key is not None:
raise ValueError(
f"MCP server '{server_name}': specify either client_cert as "
f"a list [cert, key] OR client_cert + client_key, not both"
)
if len(raw_cert) == 2:
return (_expand(raw_cert[0], "client_cert[0]"), _expand(raw_cert[1], "client_cert[1]"))
if len(raw_cert) == 3:
cert_path = _expand(raw_cert[0], "client_cert[0]")
key_path = _expand(raw_cert[1], "client_cert[1]")
password = raw_cert[2]
if not isinstance(password, str):
raise ValueError(
f"MCP server '{server_name}': client_cert[2] (key "
f"passphrase) must be a string"
)
return (cert_path, key_path, password)
raise ValueError(
f"MCP server '{server_name}': client_cert list form must have 2 "
f"or 3 elements (got {len(raw_cert)})"
)
cert_path = _expand(raw_cert, "client_cert")
if raw_key is not None:
return (cert_path, _expand(raw_key, "client_key"))
return cert_path # single combined PEM (cert + key)
def _resolve_identity_header(server_name: str, config: dict):
"""Resolve the optional per-server ``identity_header`` config.
Shape: ``{name: "X-User-Id", value_from: "static"|"profile", value: "..."}``
(``value`` required for static). Returns ``(name, value)`` or None. Invalid
configs warn and are ignored — an identity header must never break the
connection. ``profile`` resolves once at connect time; no per-call mutation.
"""
raw = config.get("identity_header")
if raw is None:
return None
if not isinstance(raw, dict):
logger.warning(
"MCP server '%s': identity_header must be a mapping with "
"'name' and 'value'/'value_from' keys (got %s) — ignoring",
server_name, type(raw).__name__,
)
return None
name = raw.get("name")
if not isinstance(name, str) or not name.strip():
logger.warning(
"MCP server '%s': identity_header requires a non-empty "
"'name' — ignoring", server_name,
)
return None
value_from = (raw.get("value_from") or "static").strip().lower()
if value_from == "static":
value = raw.get("value")
if not isinstance(value, str) or not value.strip():
logger.warning(
"MCP server '%s': identity_header with value_from: static "
"requires a non-empty string 'value' — ignoring",
server_name,
)
return None
return (name.strip(), value)
if value_from == "profile":
from hermes_cli.profiles import get_active_profile_name
return (name.strip(), get_active_profile_name())
logger.warning(
"MCP server '%s': identity_header value_from must be 'static' or "
"'profile' (got %r) — ignoring", server_name, value_from,
)
return None
def _apply_identity_header(server_name: str, config: dict, headers: dict) -> dict:
"""Merge the resolved identity header into ``headers`` in place.
An explicit per-server ``headers`` entry with the same name (any casing)
wins — the identity header never silently overrides user config.
"""
resolved = _resolve_identity_header(server_name, config)
if resolved is None:
return headers
name, value = resolved
if any(key.lower() == name.lower() for key in headers):
logger.debug(
"MCP server '%s': identity_header '%s' already set via explicit "
"headers config — keeping the explicit value", server_name, name,
)
return headers
headers[name] = value
return headers
def _make_redirect_header_stripper(
original_url,
*,
strict: bool = False,
configured_header_names: "set[str] | frozenset[str]" = frozenset(),
):
"""Build an httpx response hook that guards cross-origin redirects.
Always strips ``Authorization`` when a redirect leaves the original origin.
With *strict* (Agent Plugins v1 ``strict_redirect_headers``) every configured
header (lowercase names in *configured_header_names*) is stripped too — the
v1 spec forbids forwarding package-configured headers cross-origin.
"""
async def _strip_on_cross_origin_redirect(response):
if response.is_redirect and response.next_request:
target = response.next_request.url
if (target.scheme, target.host, target.port) != (
original_url.scheme, original_url.host, original_url.port,
):
response.next_request.headers.pop("authorization", None)
response.next_request.headers.pop("Authorization", None)
if strict:
for _name in configured_header_names:
while _name in response.next_request.headers:
del response.next_request.headers[_name]
return _strip_on_cross_origin_redirect
def _format_connect_error(exc: BaseException) -> str:
"""Render nested MCP connection errors into an actionable short message."""
def _find_missing(current: BaseException) -> Optional[str]:
nested = getattr(current, "exceptions", None)
if nested:
for child in nested:
missing = _find_missing(child)
if missing:
return missing
return None
if isinstance(current, FileNotFoundError):
if getattr(current, "filename", None):
return str(current.filename)
match = re.search(r"No such file or directory: '([^']+)'", str(current))
if match:
return match.group(1)
for attr in ("__cause__", "__context__"):
nested_exc = getattr(current, attr, None)
if isinstance(nested_exc, BaseException):
missing = _find_missing(nested_exc)
if missing:
return missing
return None
def _flatten_messages(current: BaseException) -> List[str]:
nested = getattr(current, "exceptions", None)
if nested:
flattened: List[str] = []
for child in nested:
flattened.extend(_flatten_messages(child))
return flattened
messages = []
text = str(current).strip()
if text:
messages.append(text)
for attr in ("__cause__", "__context__"):
nested_exc = getattr(current, attr, None)
if isinstance(nested_exc, BaseException):
messages.extend(_flatten_messages(nested_exc))
return messages or [current.__class__.__name__]
missing = _find_missing(exc)
if missing:
message = f"missing executable '{missing}'"
if os.path.basename(missing) in {"npx", "npm", "node"}:
message += (
" (ensure Node.js is installed and PATH includes its bin directory, "
"or set mcp_servers.<name>.command to an absolute path and include "
"that directory in mcp_servers.<name>.env.PATH)"
)
return _sanitize_error(message)
deduped: List[str] = []
for item in _flatten_messages(exc):
if item not in deduped:
deduped.append(item)
return _sanitize_error("; ".join(deduped[:3]))
# Lazily-built caches so this module imports even without the SDK OAuth module.
_AUTH_ERROR_TYPES: tuple = ()
_HTTP_STATUS_ERROR_TYPES: Optional[tuple] = None
def _http_status_error_types() -> tuple:
"""``HTTPStatusError`` classes from both httpx flavours.
A 401 may come from the SDK's own stack (``httpx2`` on mcp >= 2.0) or from
Hermes' pinned ``httpx``; the classes are unrelated, so both go in the tuple.
"""
global _HTTP_STATUS_ERROR_TYPES
if _HTTP_STATUS_ERROR_TYPES is not None:
return _HTTP_STATUS_ERROR_TYPES
found: list = []
sdk_mod = _core.sdk_httpx()
if sdk_mod is not None:
found.append(sdk_mod.HTTPStatusError)
try:
import httpx
if httpx.HTTPStatusError not in found:
found.append(httpx.HTTPStatusError)
except ImportError:
pass
_HTTP_STATUS_ERROR_TYPES = tuple(found)
return _HTTP_STATUS_ERROR_TYPES
def _get_auth_error_types() -> tuple:
"""Cached tuple of exception types indicating MCP OAuth failure.
SDK ``OAuthFlowError``/``OAuthTokenError`` (+ legacy ``UnauthorizedError``),
our ``OAuthNonInteractiveError``, and ``HTTPStatusError`` from both httpx
flavours — the latter needs the 401 status check in :func:`_is_auth_error`.
"""
global _AUTH_ERROR_TYPES
if _AUTH_ERROR_TYPES:
return _AUTH_ERROR_TYPES
types: list = []
try:
from mcp.client.auth import OAuthFlowError, OAuthTokenError
types.extend([OAuthFlowError, OAuthTokenError])
except ImportError:
pass
try:
from mcp.client.auth import UnauthorizedError # type: ignore # older SDKs
types.append(UnauthorizedError)
except ImportError:
pass
try:
from tools.mcp_oauth import OAuthNonInteractiveError
types.append(OAuthNonInteractiveError)
except ImportError:
pass
types.extend(_http_status_error_types())
_AUTH_ERROR_TYPES = tuple(types)
return _AUTH_ERROR_TYPES
def _is_auth_error(exc: BaseException) -> bool:
"""True if ``exc`` indicates an MCP OAuth failure.
``HTTPStatusError`` counts only with status 401; other HTTP errors fall
through to the generic error path.
"""
types = _get_auth_error_types()
if not types or not isinstance(exc, types):
return False
status_error_types = _http_status_error_types()
if status_error_types and isinstance(exc, status_error_types):
return getattr(exc.response, "status_code", None) == 401
return True
# Lower-cased substrings meaning the server-side transport session expired /
# was GC'd. The OAuth token is still valid — only the transport needs rebuilding.
_SESSION_EXPIRED_MARKERS: tuple = (
"invalid or expired session",
"expired session",
"session expired",
"session not found",
"unknown session",
"session terminated",
"closedresourceerror",
"closed resource",
"transport is closed",
"connection closed",
"broken pipe",
"end of file",
)
# Node budget for ``_is_session_expired_error``. The visited set breaks cycles;
# the budget bounds pathological acyclic graphs. Kept well above
# ``sys.getrecursionlimit()`` so deep task-group nesting is still fully scanned.
_EXC_TRAVERSAL_MAX_NODES = 10_000
def _is_session_expired_error(exc: BaseException) -> bool:
"""True if ``exc`` looks like an MCP transport session expiry.
Streamable-HTTP servers GC session state (idle TTL, restart, pod rotation)
while the OAuth token stays valid, so unlike :func:`_is_auth_error` the fix
is a transport reconnect (``_reconnect_event``), not an OAuth refresh.
"""
# AnyIO stream exceptions are often message-less (``str(ClosedResourceError()) == ""``),
# so type checks are needed in addition to marker matching.
try:
from anyio import BrokenResourceError, ClosedResourceError, EndOfStream
transport_error_types = (
BrokenResourceError,
ClosedResourceError,
EndOfStream,
)
except ImportError: # pragma: no cover - AnyIO is supplied by the MCP SDK
transport_error_types = ()
# Iterative traversal over ``exceptions`` / ``__cause__`` / ``__context__``
# with an identity-visited set AND a node budget (graphs can be deep or
# cyclic). Every reachable node is inspected so an InterruptedError anywhere
# overrides transport markers; the chain walk matters because SDK wrappers
# often raise a generic RuntimeError *from* the message-less ClosedResourceError.
stack: "list[BaseException | None]" = [exc]
seen: set[int] = set()
transport_error_found = False
budget = _EXC_TRAVERSAL_MAX_NODES
while stack and budget > 0:
current = stack.pop()
if current is None:
continue
identity = id(current)
if identity in seen:
continue
seen.add(identity)
budget -= 1
if isinstance(current, InterruptedError):
return False
if isinstance(current, transport_error_types):
transport_error_found = True
# Messages vary across SDK versions and servers: match a narrow
# allow-list of stable substrings, not exception type, to avoid false positives.
msg = str(current).lower()
if msg and any(marker in msg for marker in _SESSION_EXPIRED_MARKERS):
transport_error_found = True
stack.extend(getattr(current, "exceptions", ()))
stack.append(getattr(current, "__cause__", None))
stack.append(getattr(current, "__context__", None))
return transport_error_found
+785
View File
@@ -0,0 +1,785 @@
"""Registry-facing sync handlers for MCP tools and utility tools (resources/prompts), plus the per-call recovery ladder: trust gating, circuit breaker, auth (401) refresh, session-expired reconnect and dead-stdio respawn retry. Split from tools/mcp_tool.py."""
import logging
import asyncio
import contextvars
import inspect
import json
import time
from contextlib import asynccontextmanager
from types import SimpleNamespace
from typing import Any, Dict, List, Optional
from tools.registry import tool_error
from tools.ansi_strip import strip_unicode_tags
from tools.mcp_tool_common import _exc_str, _sanitize_error, mcp_field, _core
from tools.mcp_tool_content import _MCP_HARD_RESULT_CAP_CHARS, _cache_mcp_audio_block, _cache_mcp_image_block, _render_mcp_resource_block, _strip_reserved_meta_keys, _truncate_mcp_text_result
from tools.mcp_tool_errors import _is_session_expired_error
logger = logging.getLogger("tools.mcp_tool")
def _trust_gate_check(server_name: str, tool_name: str) -> Optional[str]:
"""Approval gate for write-capable tools on ``trust: untrusted`` servers.
Returns None to proceed, or a ``tool_error`` string when blocked.
Fail-closed: approval-system errors block the call.
"""
trust = _core._server_trust_levels.get(server_name, _core._TRUST_FULL)
if trust != _core._TRUST_UNTRUSTED:
return None
if _core._tool_read_only_hints.get(server_name, {}).get(tool_name) is True:
return None
# Lazy import: tools.approval routes the prompt to whichever surface owns
# the session (CLI, TUI, Telegram, ...) and normalizes the answer.
try:
from tools.approval import request_elicitation_consent
answer = request_elicitation_consent(
(
f"MCP tool '{tool_name}' on UNTRUSTED server "
f"'{server_name}' wants to run. This tool is write-capable "
f"(no readOnlyHint=true annotation) and may modify external "
f"state."
),
(
f"Server '{server_name}' is configured 'trust: untrusted'. "
f"Approve to run '{tool_name}' once, or deny to block it."
),
surface=f"mcp-trust/{server_name}",
)
except Exception as exc:
logger.error(
"MCP trust gate: approval check failed for %s.%s: %s",
server_name, tool_name, exc, exc_info=True,
)
return tool_error(
f"MCP tool '{tool_name}' on untrusted server '{server_name}' "
f"was blocked: the approval system was unavailable "
f"(fail-closed)."
)
if answer == "accept":
return None
logger.info(
"MCP trust gate: user %s '%s' on untrusted server '%s'",
"cancelled" if answer == "cancel" else "denied",
tool_name, server_name,
)
return tool_error(
f"The user did not approve running write-capable MCP tool "
f"'{tool_name}' on untrusted server '{server_name}'. The command "
f"was NOT run. Do not retry without explicit user direction."
)
def _result_is_error(result) -> bool:
"""True only for a JSON payload carrying an ``error`` key (non-JSON = success)."""
try:
return "error" in json.loads(result)
except (json.JSONDecodeError, TypeError):
return False
def _retry_once(server_name: str, retry_call, op_description: str, what: str):
"""Re-run ``retry_call`` after a recovery step.
Returns the result (and closes the circuit breaker) when it is not an error
payload; None when the retry raised or errored, so the caller falls through.
"""
try:
result = retry_call()
except Exception as retry_exc:
logger.warning(
"MCP %s/%s retry after %s failed: %s",
server_name, op_description, what, retry_exc,
)
return None
if _result_is_error(result):
return None
_core._reset_server_error(server_name)
return result
def _handle_auth_error_and_retry(
server_name: str,
exc: BaseException,
retry_call,
op_description: str,
):
"""Attempt OAuth recovery and one retry; return None to fall through.
Non-auth exceptions return None. Otherwise: ask ``MCPOAuthManager.handle_401``
whether recovery is viable; if so, signal ``_reconnect_event`` so the server
task rebuilds the session with fresh credentials, wait for ready, retry once.
Any failure returns the structured ``needs_reauth`` error so the model stops
trying to refresh manually.
"""
if not _core._is_auth_error(exc):
return None
from tools.mcp_oauth_manager import get_manager
manager = get_manager()
async def _recover():
return await manager.handle_401(server_name, None)
try:
recovered = _core._run_on_mcp_loop(_recover, timeout=10)
except Exception as rec_exc:
logger.warning(
"MCP OAuth '%s': recovery attempt failed: %s",
server_name, rec_exc,
)
recovered = False
if recovered:
with _core._lock:
srv = _core._servers.get(server_name)
reconnected = False
if srv is not None and hasattr(srv, "_reconnect_event"):
reconnected = _core._signal_reconnect_and_wait(
server_name,
srv,
op_description=f"{op_description} after OAuth recovery",
timeout=15,
)
# OAuth recovery + reconnect is independent evidence the server is viable,
# so close the breaker here, not only on retry success — otherwise a failing
# retry would leave it pinned open forever. A broken server re-trips it via
# _bump_server_error on the retry.
if reconnected:
_core._reset_server_error(server_name)
result = _retry_once(server_name, retry_call, op_description, "auth recovery")
if result is not None:
return result
# No recovery, or retry failed: structured needs_reauth error + breaker strike.
_core._bump_server_error(server_name)
return tool_error(
f"MCP server '{server_name}' requires re-authentication. "
f"Run `hermes mcp login {server_name}` (or delete the tokens "
f"file under ~/.hermes/mcp-tokens/ and restart). Do NOT retry "
f"this tool — ask the user to re-authenticate.",
needs_reauth=True,
server=server_name,
)
def _handle_session_expired_and_retry(
server_name: str,
exc: BaseException,
retry_call,
op_description: str,
):
"""Trigger a transport reconnect and retry once on session expiry.
Unlike :func:`_handle_auth_error_and_retry` this skips ``handle_401`` — the
token is still valid, only the server-side session is stale. Returns None to
fall through (not session-expired, no server record / loop, reconnect did not
ready in time, or the retry also failed).
"""
if not _is_session_expired_error(exc):
return None
with _core._lock:
srv = _core._servers.get(server_name)
if srv is None or not hasattr(srv, "_reconnect_event"):
return None
loop = _core._mcp_loop
if loop is None or not loop.is_running():
return None
logger.info(
"MCP server '%s': %s failed with session-expired error (%s); "
"signalling transport reconnect and retrying once.",
server_name, op_description, exc,
)
if not _core._signal_reconnect_and_wait(
server_name,
srv,
op_description=op_description,
timeout=15,
):
logger.warning(
"MCP server '%s': reconnect did not ready within 15s after "
"session-expired error; falling through to error response.",
server_name,
)
return None
return _retry_once(server_name, retry_call, op_description, "session reconnect")
class _StdioChildExited(RuntimeError):
"""A server's stdio subprocess was gone when (or while) a call ran.
Deliberately NOT a TimeoutError: nothing timed out — the child was already
dead (typically a gateway restart killed it under a live agent session).
Handled by :func:`_handle_stdio_child_exited_and_retry`.
"""
def _handle_stdio_child_exited_and_retry(
server_name: str,
exc: Exception,
retry_call,
op_description: str,
):
"""Respawn a dead stdio child and retry the call once; None if not our error.
Cannot hot-cycle respawns: this never spawns anything — it sets
``_reconnect_event`` once and waits for the server task to publish a fresh
session, so spawn frequency stays governed by ``run()``'s rapid-drop budget.
Single-shot: a child that dies again immediately reports and stops.
"""
if not isinstance(exc, _StdioChildExited):
return None
with _core._lock:
srv = _core._servers.get(server_name)
reconnected = False
if srv is not None and hasattr(srv, "_reconnect_event"):
logger.info(
"MCP server '%s': %s found the stdio subprocess dead (%s); "
"respawning and retrying once.",
server_name, op_description, exc,
)
loop = _core._mcp_loop
if loop is not None and loop.is_running():
reconnected = _core._signal_reconnect_and_wait(
server_name,
srv,
op_description=op_description,
timeout=_core._STDIO_RESPAWN_WAIT_SEC,
)
else:
# No MCP loop to wait on (non-async adapters, tests) — still request
# the respawn so the next call lands on a live transport.
_core._signal_reconnect(srv)
if reconnected:
try:
result = retry_call()
except _StdioChildExited as retry_exc:
# Died again right after respawn: broken server, not a restart
# artifact. Stop here — run()'s budget takes it to the park.
logger.warning(
"MCP server '%s': %s stdio subprocess exited again right "
"after respawn (%s); not retrying further.",
server_name, op_description, retry_exc,
)
_core._bump_server_error(server_name)
return tool_error(
f"MCP server '{server_name}' respawned its stdio subprocess "
f"and it exited again immediately. The server is not "
f"starting cleanly — do NOT retry this tool; ask the user to "
f"check the server's command and its stderr log."
)
except Exception as retry_exc:
logger.warning(
"MCP %s/%s retry after stdio respawn failed: %s",
server_name, op_description, retry_exc,
)
_core._bump_server_error(server_name)
return tool_error(_sanitize_error(
f"MCP call failed after respawning the stdio subprocess for "
f"'{server_name}': {type(retry_exc).__name__}: "
f"{_exc_str(retry_exc)}"
))
if _result_is_error(result):
_core._bump_server_error(server_name)
else:
_core._reset_server_error(server_name)
return result
_core._bump_server_error(server_name)
return tool_error(
f"MCP server '{server_name}' stdio subprocess had exited (this is "
f"not a timeout — the call never reached the server). A respawn was "
f"requested but no fresh session came back within "
f"{_core._STDIO_RESPAWN_WAIT_SEC:.0f}s. Wait a few seconds before retrying; "
f"if it keeps failing the server is not starting and needs the user."
)
def _interrupted_call_result() -> str:
"""Standardized JSON error for a user-interrupted MCP tool call."""
return tool_error("MCP call interrupted: user sent a new message")
def _mark_server_call_started(server: Any) -> None:
"""Record a user-visible MCP operation when the server supports it."""
mark_tool_call = getattr(server, "mark_tool_call", None)
if callable(mark_tool_call):
mark_tool_call()
@asynccontextmanager
async def _track_inflight_rpc(server: Any, server_name: str, op: str):
"""Register the running RPC on the server so teardown can fail it fast.
If a deliberate reconnect/shutdown teardown cancels the task
(``_fail_inflight_calls`` sets ``_reconnecting`` first) the cancel becomes a
clean retryable RuntimeError; external cancels (caller timeout, user
interrupt) propagate unchanged. Test doubles without ``_inflight_tasks``
simply skip tracking.
"""
inflight = getattr(server, "_inflight_tasks", None)
task = asyncio.current_task()
if task is not None and inflight is not None:
inflight.add(task)
try:
yield
except asyncio.CancelledError:
if getattr(server, "_reconnecting", False):
raise RuntimeError(
f"MCP {op} on '{server_name}' was aborted by a reconnect "
f"teardown; retry the request on the rebuilt session"
) from None
raise
finally:
if task is not None and inflight is not None:
inflight.discard(task)
def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float):
"""Return a sync registry handler (``handler(args_dict, **kwargs) -> str``)
that calls an MCP tool via the background loop."""
def _handler(args: dict, **kwargs) -> str:
# Security boundary: untrusted-server write tools need approval before
# ANY transport work, including the lazy first-use spawn below.
gate_error = _trust_gate_check(server_name, tool_name)
if gate_error is not None:
return gate_error
# Circuit breaker. After the cooldown the breaker is half-open: the next
# call goes through as a probe; success resets it, failure re-bumps (which
# re-stamps the open-time and re-arms the cooldown).
if _core._server_error_counts.get(server_name, 0) >= _core._CIRCUIT_BREAKER_THRESHOLD:
opened_at = _core._server_breaker_opened_at.get(server_name, 0.0)
age = time.monotonic() - opened_at
if age < _core._CIRCUIT_BREAKER_COOLDOWN_SEC:
remaining = max(1, int(_core._CIRCUIT_BREAKER_COOLDOWN_SEC - age))
return tool_error(
f"MCP server '{server_name}' is unreachable after "
f"{_core._server_error_counts[server_name]} consecutive "
f"failures. Auto-retry available in ~{remaining}s. "
f"Do NOT retry this tool yet — use alternative "
f"approaches or ask the user to check the MCP server."
)
server = _core._get_connected_server_for_call(server_name)
if not server:
_core._bump_server_error(server_name)
return tool_error(f"MCP server '{server_name}' is not connected")
# No session: a reconnect may be completing (fresh session swaps in
# asynchronously), so wait briefly before charging a breaker strike.
if not server.session and not _core._wait_for_server_session_ready(
server, timeout=min(5.0, float(tool_timeout or 5.0)),
):
# Still down — reconnecting or parked (e.g. dead stdio child).
# Probing a dead transport would re-arm the breaker forever, so
# ask the server task to rebuild (respawns stdio) and return a
# clean "reconnecting" error; the breaker resets once the fresh
# session initializes.
_core._bump_server_error(server_name)
if _core._signal_reconnect(server):
return tool_error(
f"MCP server '{server_name}' transport is down; "
f"reconnect requested. Do NOT retry this tool "
f"immediately — give it a few seconds to come back."
)
return tool_error(f"MCP server '{server_name}' is not connected")
async def _call():
_mark_server_call_started(server)
async with server._rpc_lock, _track_inflight_rpc(
server, server_name, f"tools/call {tool_name}"
):
# Snapshot contextvars so an elicitation callback (fired on the
# MCP recv loop, which doesn't inherit them) can replay them
# for gateway platform / session routing.
server._pending_call_context = contextvars.copy_context()
try:
# Fast-fail: an already-dead stdio child must not hold this
# slot for the full tool timeout. callable() + real-bool check
# because MagicMock attributes return truthy Mocks.
_stdio_dead = getattr(server, "_stdio_children_dead", None)
if (
callable(_stdio_dead)
and isinstance(_stdio_dead_result := _stdio_dead(), bool)
and _stdio_dead_result
):
# server.session is stale so the transport-down path above
# never fired; hand this to the respawn-and-retry path.
raise _StdioChildExited(
f"MCP stdio subprocess for '{server_name}' had "
f"already exited when the call was dispatched"
)
_call_coro = server.session.call_tool(tool_name, arguments=args)
_watch_children = getattr(server, "_watch_stdio_children", None)
_watch_ok = (
_watch_children is not None
and inspect.iscoroutinefunction(_watch_children)
and asyncio.iscoroutine(_call_coro)
)
if not _watch_ok:
# Stubbed sessions return a non-awaitable, or there is no
# child-watcher to race: plain await.
result = (
await _call_coro
if asyncio.iscoroutine(_call_coro)
else _call_coro
)
else:
# Race the RPC against the stdio-children watcher so a
# mid-call death fails immediately.
rpc_task = asyncio.ensure_future(_call_coro)
watch_task = asyncio.ensure_future(_watch_children())
try:
done, _pending = await asyncio.wait(
{rpc_task, watch_task},
return_when=asyncio.FIRST_COMPLETED,
)
if watch_task in done and not rpc_task.done():
rpc_task.cancel()
# Nothing clears server.session on a mid-call
# death; the respawn-and-retry path owns the
# reconnect signal.
raise _StdioChildExited(
f"MCP stdio subprocess for "
f"'{server_name}' exited mid-call"
)
result = await rpc_task
finally:
watch_task.cancel()
if not rpc_task.done():
rpc_task.cancel()
await asyncio.gather(
rpc_task, watch_task, return_exceptions=True
)
finally:
server._pending_call_context = None
# Round-trip completed: transport is healthy even if the tool
# returned isError. Clear the rapid-drop budget.
_mark_proven = getattr(server, "_mark_session_proven", None)
if _mark_proven is not None:
_mark_proven()
# CallToolResult: .content (blocks) and .is_error (.isError before mcp 2.0).
if mcp_field(result, "is_error", "isError", False):
error_text = ""
for block in (result.content or []):
if getattr(block, "text", None):
error_text += block.text
continue
# EmbeddedResource error payloads carry text under .resource.text.
res_text = getattr(getattr(block, "resource", None), "text", None)
if res_text:
error_text += str(res_text)
return tool_error(_sanitize_error(
_truncate_mcp_text_result(
error_text or "MCP tool returned an error"
)
))
# Text blocks pass through; image/audio blocks are cached via the
# gateway image-cache so they flow out as MEDIA: tags; resource blocks
# (PDFs, docs, ...) are materialized rather than silently dropped.
parts: List[str] = []
for block in (result.content or []):
if hasattr(block, "text") and block.text:
parts.append(strip_unicode_tags(block.text))
continue
image_tag = _cache_mcp_image_block(block)
if image_tag:
parts.append(image_tag)
continue
audio_tag = _cache_mcp_audio_block(block)
if audio_tag:
parts.append(audio_tag)
continue
resource_text = _render_mcp_resource_block(block, server_name)
if resource_text:
parts.append(resource_text)
continue
# Benign empty renders log at debug; warn only for unknown shapes.
block_type = getattr(block, "type", None) or type(block).__name__
if block_type in {"text", "resource", "audio", "image"}:
logger.debug(
"MCP %s: content block type %r rendered empty",
server_name, block_type,
)
else:
logger.warning(
"MCP %s: dropping unsupported content block type %r",
server_name, block_type,
)
text_result = "\n".join(parts) if parts else ""
# Hard-cap pathological payloads; ordinary large results pass to spillover.
text_result = _truncate_mcp_text_result(text_result)
# content is the primary (model-oriented) payload; structuredContent
# supplements it. Server-level `_meta` is surfaced too, minus
# protocol-reserved keys (`modelcontextprotocol`/`mcp` label followed
# by another label); vendor-namespaced keys pass through.
structured = mcp_field(result, "structured_content", "structuredContent")
# Cap structuredContent too (multi-MB JSON flood); over the hard cap it
# degrades to the head+tail truncated string.
if structured is not None:
try:
_structured_json = json.dumps(structured, ensure_ascii=False, default=str)
except (TypeError, ValueError):
_structured_json = None
if _structured_json is not None and len(_structured_json) > _MCP_HARD_RESULT_CAP_CHARS:
structured = _truncate_mcp_text_result(_structured_json)
meta = _strip_reserved_meta_keys(mcp_field(result, "meta", "meta"))
if structured is not None or meta is not None:
payload: Dict[str, Any] = {}
if text_result:
payload["result"] = text_result
if structured is not None:
if text_result:
payload["structuredContent"] = structured
else:
payload["result"] = structured
if meta is not None:
payload["_meta"] = meta
if "result" not in payload:
payload["result"] = text_result
try:
return json.dumps(payload, ensure_ascii=False)
except (TypeError, ValueError):
# Non-serializable metadata: drop the extras, keep the call.
return json.dumps({"result": text_result}, ensure_ascii=False)
return json.dumps({"result": text_result}, ensure_ascii=False)
def _call_once():
return _core._run_on_mcp_loop(_call, timeout=tool_timeout)
try:
result = _call_once()
# An error payload from the tool itself still counts as a strike.
if _result_is_error(result):
_core._bump_server_error(server_name)
else:
_core._reset_server_error(server_name)
return result
except InterruptedError:
return _interrupted_call_result()
except Exception as exc:
# Recovery ladder, in order: dead stdio child (respawn + retry),
# auth (OAuth recovery + retry), session expiry (reconnect + retry).
# Each returns None when the exception is not its kind.
for recover in (
_handle_stdio_child_exited_and_retry,
_handle_auth_error_and_retry,
_handle_session_expired_and_retry,
):
recovered = recover(server_name, exc, _call_once, f"tools/call {tool_name}")
if recovered is not None:
return recovered
_core._bump_server_error(server_name)
logger.error(
"MCP tool %s/%s call failed: %s",
server_name, tool_name, exc,
)
return tool_error(_sanitize_error(
f"MCP call failed: {type(exc).__name__}: {_exc_str(exc)}"
))
return _handler
def _make_utility_handler(server_name: str, tool_timeout: float, op: str,
log_label: str, build_call):
"""Shared shape of the four utility handlers (resources/prompts).
``build_call(server, args)`` returns an error string (validation failed) or a
zero-arg coroutine function doing the RPC under ``_rpc_lock``. The wrapper owns
the connected check and the auth / session-expired recovery ladder.
"""
def _handler(args: dict, **kwargs) -> str:
server = _core._get_connected_server_for_call(server_name)
if not server or not server.session:
return tool_error(f"MCP server '{server_name}' is not connected")
call = build_call(server, args)
if isinstance(call, str):
return call
def _call_once():
return _core._run_on_mcp_loop(call, timeout=tool_timeout)
try:
return _call_once()
except InterruptedError:
return _interrupted_call_result()
except Exception as exc:
for recover in (_handle_auth_error_and_retry, _handle_session_expired_and_retry):
recovered = recover(server_name, exc, _call_once, op)
if recovered is not None:
return recovered
logger.error("MCP %s/%s failed: %s", server_name, log_label, exc)
return tool_error(_sanitize_error(
f"MCP call failed: {type(exc).__name__}: {_exc_str(exc)}"
))
return _handler
def _make_list_resources_handler(server_name: str, tool_timeout: float):
"""Return a sync handler that lists resources from an MCP server."""
def _build(server, args):
async def _call():
_mark_server_call_started(server)
async with server._rpc_lock:
all_resources = await _core._paginate_full_list(
server.session.list_resources, "resources", server_name
)
resources = []
for r in all_resources:
entry = {}
if hasattr(r, "uri"):
entry["uri"] = str(r.uri)
if hasattr(r, "name"):
entry["name"] = r.name
if hasattr(r, "description") and r.description:
entry["description"] = r.description
# Key stays camelCase — this is the tool's own JSON output shape.
_mime = mcp_field(r, "mime_type", "mimeType")
if _mime:
entry["mimeType"] = _mime
resources.append(entry)
return json.dumps({"resources": resources}, ensure_ascii=False)
return _call
return _make_utility_handler(server_name, tool_timeout, "resources/list", "list_resources", _build)
def _make_read_resource_handler(server_name: str, tool_timeout: float):
"""Return a sync handler that reads a resource by URI from an MCP server."""
def _build(server, args):
uri = args.get("uri")
if not uri:
return tool_error("Missing required parameter 'uri'")
async def _call():
_mark_server_call_started(server)
async with server._rpc_lock:
result = await server.session.read_resource(uri)
parts: List[str] = []
contents = result.contents if hasattr(result, "contents") else []
for block in contents:
if getattr(block, "text", None) is not None:
parts.append(strip_unicode_tags(block.text))
elif getattr(block, "blob", None) is not None:
# Materialize binary contents into the document cache
# (same contract as EmbeddedResource blocks in tool results).
rendered = _render_mcp_resource_block(
SimpleNamespace(type="resource", resource=block),
server_name,
)
parts.append(rendered or f"[binary data, {len(block.blob)} bytes]")
return json.dumps({"result": "\n".join(parts) if parts else ""}, ensure_ascii=False)
return _call
return _make_utility_handler(server_name, tool_timeout, "resources/read", "read_resource", _build)
def _make_list_prompts_handler(server_name: str, tool_timeout: float):
"""Return a sync handler that lists prompts from an MCP server."""
def _build(server, args):
async def _call():
_mark_server_call_started(server)
async with server._rpc_lock:
all_prompts = await _core._paginate_full_list(
server.session.list_prompts, "prompts", server_name
)
prompts = []
for p in all_prompts:
entry = {}
if hasattr(p, "name"):
entry["name"] = p.name
if hasattr(p, "description") and p.description:
entry["description"] = p.description
if hasattr(p, "arguments") and p.arguments:
entry["arguments"] = [
{
"name": a.name,
**({"description": a.description} if hasattr(a, "description") and a.description else {}),
**({"required": a.required} if hasattr(a, "required") else {}),
}
for a in p.arguments
]
prompts.append(entry)
return json.dumps({"prompts": prompts}, ensure_ascii=False)
return _call
return _make_utility_handler(server_name, tool_timeout, "prompts/list", "list_prompts", _build)
def _make_get_prompt_handler(server_name: str, tool_timeout: float):
"""Return a sync handler that gets a prompt by name from an MCP server."""
def _build(server, args):
name = args.get("name")
if not name:
return tool_error("Missing required parameter 'name'")
arguments = args.get("arguments", {})
async def _call():
_mark_server_call_started(server)
async with server._rpc_lock:
result = await server.session.get_prompt(name, arguments=arguments)
messages = []
for msg in (result.messages if hasattr(result, "messages") else []):
entry = {}
if hasattr(msg, "role"):
entry["role"] = msg.role
if hasattr(msg, "content"):
content = msg.content
if hasattr(content, "text"):
entry["content"] = strip_unicode_tags(content.text)
elif isinstance(content, str):
entry["content"] = strip_unicode_tags(content)
else:
entry["content"] = strip_unicode_tags(str(content))
messages.append(entry)
resp = {"messages": messages}
if hasattr(result, "description") and result.description:
resp["description"] = result.description
return json.dumps(resp, ensure_ascii=False)
return _call
return _make_utility_handler(server_name, tool_timeout, "prompts/get", "get_prompt", _build)
def _make_check_fn(server_name: str):
"""Return a check function that verifies the MCP connection is alive."""
def _check() -> bool:
with _core._lock:
server = _core._servers.get(server_name)
if server is not None and (
server.session is not None or server._is_recycled_stdio()
):
return True
# Lazy (schema-cache registered) servers count as available: the
# first real call spawns/connects them.
return server_name in _core._lazy_server_configs
return _check
+392
View File
@@ -0,0 +1,392 @@
"""Session health for MCPServerTask: dynamic tool refresh on list_changed notifications, server log forwarding, keepalive probes, suspect-mark / lazy-verify, in-flight call fail-fast, stdio child liveness and stdio idle/lifetime recycling. Split from tools/mcp_tool.py."""
import asyncio
import json
import logging
import time
from typing import Optional
from tools.mcp_tool_errors import _is_method_not_found_error, _unwrap_exception_group
from tools.mcp_tool_schema import mcp_prefixed_tool_name
from tools.mcp_tool_registration import _forget_mcp_tool_server
from tools.mcp_tool_common import _core
logger = logging.getLogger("tools.mcp_tool")
class MCPServerHealthMixin:
"""Methods of :class:`tools.mcp_tool.MCPServerTask` (mixed in; relies on its attributes)."""
__slots__ = ()
def _is_http(self) -> bool:
"""Check if this server uses HTTP transport."""
return "url" in self._config
def _is_recycled_stdio(self) -> bool:
"""Return True when a stdio server was intentionally recycled."""
return not self._is_http() and self._recycled_reason is not None
def mark_tool_call(self) -> None:
"""Record that a user-visible MCP operation is starting."""
self._last_tool_call_at = time.monotonic()
def _mark_lifecycle_started(self) -> None:
now = time.monotonic()
self._lifecycle_started_at = now
self._last_tool_call_at = now
self._recycled_reason = None
def _stdio_recycle_reason(self, now: Optional[float] = None) -> Optional[str]:
"""Return the stdio recycle reason if idle/age limits have elapsed."""
if self._is_http() or self._rpc_lock.locked():
return None
now = time.monotonic() if now is None else now
if (
self._max_lifetime_seconds is not None
and now - self._lifecycle_started_at >= self._max_lifetime_seconds
):
return "max_lifetime_seconds"
if (
self._idle_timeout_seconds is not None
and now - self._last_tool_call_at >= self._idle_timeout_seconds
):
return "idle_timeout_seconds"
return None
def _next_stdio_recycle_deadline(self) -> Optional[float]:
"""Return the next monotonic recycle deadline for stdio, if any."""
if self._is_http() or self._rpc_lock.locked():
return None
deadlines = []
if self._max_lifetime_seconds is not None:
deadlines.append(self._lifecycle_started_at + self._max_lifetime_seconds)
if self._idle_timeout_seconds is not None:
deadlines.append(self._last_tool_call_at + self._idle_timeout_seconds)
return min(deadlines) if deadlines else None
def _mark_stdio_recycled(self, reason: str) -> None:
"""Mark a stdio session dormant before its transport finishes closing."""
self._recycled_reason = reason
self.session = None
async def _refresh_tools_task(self):
"""Run a dynamic tool refresh and log failures from background tasks."""
try:
await self._refresh_tools()
except Exception:
logger.exception("MCP server '%s': dynamic tool refresh failed", self.name)
def _schedule_tools_refresh(self) -> asyncio.Task:
"""Schedule a background tool refresh and keep it strongly referenced."""
task = asyncio.create_task(self._refresh_tools_task())
self._pending_refresh_tasks.add(task)
task.add_done_callback(self._pending_refresh_tasks.discard)
return task
def _make_logging_callback(self):
"""Build a ``logging_callback`` that forwards server ``notifications/message``
into Hermes logging tagged with the server name (the SDK default drops them)."""
async def _on_log(params):
try:
level = _core._MCP_LOG_LEVEL_MAP.get(
str(getattr(params, "level", "info")).lower(), logging.INFO,
)
data = getattr(params, "data", None)
if not isinstance(data, str):
try:
data = json.dumps(data, ensure_ascii=False, default=str)
except (TypeError, ValueError):
data = str(data)
# Cap payloads so a chatty server can't flood agent.log.
if len(data) > 2000:
data = data[:2000] + "... [truncated]"
logger_name = getattr(params, "logger", None)
origin = f"{self.name}/{logger_name}" if logger_name else self.name
logger.log(level, "MCP server log [%s]: %s", origin, data)
except Exception:
logger.debug(
"Failed to handle MCP log notification from '%s'",
self.name, exc_info=True,
)
return _on_log
def _make_message_handler(self):
"""Build a ``message_handler`` for ``ClientSession``: only
``ToolListChangedNotification`` triggers a refresh; prompt/resource changes are logged."""
async def _handler(message):
try:
if isinstance(message, Exception):
logger.debug("MCP message handler (%s): exception: %s", self.name, message)
return
if _core._MCP_NOTIFICATION_TYPES and isinstance(message, _core.ServerNotification):
# mcp 2.0 made ServerNotification a plain union (payload IS the
# message) instead of a RootModel (payload under ``.root``).
# ``isinstance`` accepts both; only the unwrap differs — without
# it ``.root`` raises into the catch-all and refreshes stop.
match getattr(message, "root", message):
case _core.ToolListChangedNotification():
logger.info(
"MCP server '%s': received tools/list_changed notification",
self.name,
)
# Refresh in a separate task: some servers emit
# list_changed right after initialize while another
# request is in flight, and refreshing synchronously
# inside the handler can wedge the stdio JSON-RPC stream.
self._schedule_tools_refresh()
# Yield one tick so short-lived notification contexts
# (and tests) can observe the scheduled refresh.
await asyncio.sleep(0)
case _core.PromptListChangedNotification():
logger.debug("MCP server '%s': prompts/list_changed (ignored)", self.name)
case _core.ResourceListChangedNotification():
logger.debug("MCP server '%s': resources/list_changed (ignored)", self.name)
case _:
pass
except Exception:
logger.exception("Error in MCP message handler for '%s'", self.name)
return _handler
async def _refresh_tools(self):
"""Re-fetch tools on ``tools/list_changed`` and update the registry.
The lock serializes rapid-fire notifications. After the list_tools
``await``, all mutations are synchronous — atomic on the event loop.
"""
from tools.registry import registry
if not self._advertises_tools():
# Shouldn't happen, but tools/list would raise MCPError(-32601).
return
async with self._refresh_lock:
old_tool_names = set(self._registered_tool_names)
# 1. Fetch the current tool list (follow nextCursor).
async with self._rpc_lock:
new_mcp_tools = await _core._paginate_full_list(
self.session.list_tools, "tools", self.name
)
# 2. Remove only stale names first — no nuke-and-repave: live agent
# turns may hold tool-call IDs pointing at existing handlers, and
# in-place replacement avoids transient "tool not connected" races.
toolset_name = f"mcp-{self.name}"
stale_tool_names = old_tool_names - {
mcp_prefixed_tool_name(self.name, tool.name)
for tool in new_mcp_tools
}
for tool_name in stale_tool_names:
# Never remove a colliding name currently owned by another server.
if registry.get_toolset_for_tool(tool_name) != toolset_name:
continue
registry.deregister(tool_name, scope=_core._server_registry_scope(self.name))
_forget_mcp_tool_server(tool_name)
# 3. Re-register; the helper may skip names ambiguous after normalization.
self._tools = new_mcp_tools
registered_names = _core._register_server_tools(
self.name, self, self._config
)
# A raw name can become ambiguous without changing its normalized
# name, so the pre-pass misses it: drop any old entry the final
# collision-checked registration no longer owns.
registered_name_set = set(registered_names)
for tool_name in old_tool_names - registered_name_set:
if registry.get_toolset_for_tool(tool_name) != toolset_name:
continue
registry.deregister(tool_name, scope=_core._server_registry_scope(self.name))
_forget_mcp_tool_server(tool_name)
self._registered_tool_names = registered_names
# 4. Log what changed (user-visible).
new_tool_names = set(self._registered_tool_names)
added = new_tool_names - old_tool_names
removed = old_tool_names - new_tool_names
changes = []
if added:
changes.append(f"added: {', '.join(sorted(added))}")
if removed:
changes.append(f"removed: {', '.join(sorted(removed))}")
if changes:
logger.warning(
"MCP server '%s': tools changed dynamically — %s. "
"Verify these changes are expected.",
self.name, "; ".join(changes),
)
else:
logger.info(
"MCP server '%s': dynamically refreshed %d tool(s) (no changes)",
self.name, len(self._registered_tool_names),
)
async def _keepalive_probe(self) -> None:
"""Exercise the session; raise on a genuine connection failure.
``ping`` first (cheap, OPTIONAL utility). On -32601 latch
``_ping_unsupported`` and fall back to ``list_tools`` when the server
advertises tools; otherwise the -32601 propagates (no liveness primitive
left). The latch resets on each fresh transport connection.
"""
if not self._ping_unsupported:
try:
await asyncio.wait_for(self.session.send_ping(), timeout=30.0)
return
except Exception as exc:
if _is_method_not_found_error(exc):
# Ping is definitively unsupported.
if not self._advertises_tools():
raise
self._ping_unsupported = True
logger.info(
"MCP server '%s': does not implement the optional "
"'ping' utility (-32601); using 'list_tools' for "
"keepalive on this connection.",
self.name,
)
elif isinstance(exc, (TimeoutError, asyncio.TimeoutError)) and self._advertises_tools():
# A server that silently drops ping looks like a dead transport.
# Confirm with list_tools before declaring it dead; if that
# also fails, propagate the original failure.
try:
await asyncio.wait_for(self.session.list_tools(), timeout=30.0)
except Exception:
raise exc from None
# Transport alive; latch so later keepalives skip the 30s wait.
self._ping_unsupported = True
logger.info(
"MCP server '%s': ping timed out but list_tools "
"succeeded — server silently drops ping; using "
"'list_tools' for keepalive on this connection.",
self.name,
)
return
else:
# Closed transport, expired session, etc. — real failure.
raise
# Fallback probe for servers without ping support.
await asyncio.wait_for(self.session.list_tools(), timeout=30.0)
def _mark_session_proven(self) -> None:
"""Record that the session demonstrated real health (keepalive or tool-call success).
Only then is the reconnect budget cleared: a handshake that drops moments
later must keep consuming ``_reconnect_retries`` so a flapping transport
still reaches the park instead of respawning forever.
"""
if not self._session_proven:
self._session_proven = True
self._reconnect_retries = 0
if self._was_parked:
self._was_parked = False
logger.warning(
"MCP server '%s': revived — session healthy again after "
"parking (state: parked → connected)",
self.name,
)
# A proven fresh transport clears the one-time permanent-failure
# grace and any race bookkeeping.
self._permanent_grace_used = False
self._teardown_race = False
def mark_suspect(self, reason: str) -> None:
"""Latch a suspicion (no I/O). The NEXT call verifies via
:meth:`ensure_healthy` and recycles the transport if the probe fails."""
if self._suspect_reason is None and reason:
logger.warning(
"MCP server '%s': connection marked suspect (%s); next call "
"will health-check it",
self.name, reason,
)
self._suspect_reason = reason or None
async def ensure_healthy(self, timeout: float = 5.0) -> bool:
"""Verify a suspect connection before reuse; recycle if dead.
True when healthy (suspicion cleared). On failure requests a reconnect,
drops the stale session so the caller's no-session path takes over, and
returns False. Never raises.
"""
reason = self._suspect_reason
if not reason:
return True
if self.session is None:
# Nothing to verify — the reconnect path owns recovery now.
self._suspect_reason = None
self._reconnect_event.set()
return False
try:
await asyncio.wait_for(self._keepalive_probe(), timeout=timeout)
except Exception as exc:
root = _unwrap_exception_group(exc)
logger.warning(
"MCP server '%s': suspect connection (%s) failed health "
"check (%s: %s) — requesting reconnect (state: suspect → "
"degraded)",
self.name, reason, type(root).__name__, root,
)
self._suspect_reason = None
self.mark_suspect(f"health check failed after {reason}")
self.session = None
self._ready.clear()
self._reconnect_event.set()
return False
logger.info(
"MCP server '%s': suspect connection passed health check "
"(%s) — clearing suspicion",
self.name, reason,
)
self._suspect_reason = None
self._mark_session_proven()
return True
def _fail_inflight_calls(self, reason: str) -> None:
"""Cancel every in-flight RPC on this connection.
Called from lifecycle exits BEFORE the transport unwinds: the SDK does
not always fail pending requests when streams close, so a call would
otherwise wait out the full tool timeout. Cancelling anything flags
``_teardown_race`` so run() treats the next reconnect as recovery
rather than charging the rapid-drop budget.
"""
victims = [t for t in self._inflight_tasks if not t.done()]
if not victims:
return
self._reconnecting = True
self._teardown_race = True
self.mark_suspect(f"{reason} tore down {len(victims)} in-flight call(s)")
for task in victims:
task.cancel()
def _stdio_children_dead(self) -> bool:
"""True when every stdio child we spawned has exited.
Best-effort: False (unknown → don't fail fast) for HTTP, no captured
PIDs, missing psutil, or a failed probe.
"""
pids = getattr(self, "_stdio_child_pids", None)
if not pids or self._is_http():
return False
try:
import psutil
except ImportError:
return False
for pid in pids:
# pid_exists handles Windows without signal-permission noise.
try:
alive = psutil.pid_exists(pid)
except Exception:
return False
if alive:
return False
return True
async def _watch_stdio_children(self) -> None:
"""Poll child liveness while a stdio RPC is in flight; resolves when a
tracked child dies so the caller cancels the RPC instead of waiting out the timeout."""
while True:
if self._stdio_children_dead():
return
await asyncio.sleep(0.25)
+342
View File
@@ -0,0 +1,342 @@
"""MCP process lifecycle: stdio child PID tracking and orphan cleanup, graceful
server shutdown and draining of the background MCP loop."""
import logging
import asyncio
import os
import time
from typing import Dict, Optional
from tools.mcp_tool_common import _core
logger = logging.getLogger("tools.mcp_tool")
# Live stdio MCP children (pid -> server_name), added after connection and
# 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
# separately from _stdio_pids so the PGID survives the child's removal.
# Empty on Windows (os.getpgid is POSIX-only).
_stdio_pgids: Dict[int, int] = {}
def _snapshot_child_pids() -> set:
"""Current direct-child PIDs: /proc on Linux, else psutil, else empty set."""
my_pid = os.getpid()
# /proc/<pid>/task/<tid>/children is per-THREAD, and stdio_client() spawns
# from the MCP loop thread, so union every task's children — reading only
# the main thread's file returns an empty set on every Linux install.
try:
task_dir = f"/proc/{my_pid}/task"
tids = os.listdir(task_dir)
found: set = set()
for tid in tids:
try:
with open(f"{task_dir}/{tid}/children", encoding="utf-8") as f:
found.update(int(p) for p in f.read().split() if p.strip())
except (FileNotFoundError, OSError, ValueError):
continue # thread exited between listdir and open
return found
except (FileNotFoundError, OSError, ValueError):
pass
try:
import psutil
return {c.pid for c in psutil.Process(my_pid).children()}
except Exception:
pass
return set()
# argv markers of non-MCP gateway children that can race into the snapshot
# delta during an MCP spawn (defense-in-depth; LSP/slash_worker already use
# start_new_session). Matched against argv[1:] because Python/Java children
# start with the interpreter path.
_NON_MCP_CHILD_CMDLINE_MARKERS: tuple[str, ...] = (
"tui_gateway.slash_worker",
"tui_gateway.entry",
"-dorg.eclipse.equinox.launcher", # jdtls (legacy arg style)
"eclipse.jdt.ls",
"org.eclipse.equinox.launcher_",
)
def _filter_mcp_children(pids: set) -> set:
"""Drop non-MCP children from a PID snapshot delta.
Tracking a stray child in _stdio_pgids is catastrophic if it lacks
start_new_session: its pgid can be the TUI parent's, so the shutdown
killpg() would kill the TUI itself.
"""
if not pids:
return pids
try:
import psutil
except ImportError:
return pids # keep all PIDs (prior behavior)
filtered: set = set()
for pid in pids:
try:
argv = psutil.Process(pid).cmdline()
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
):
continue
filtered.add(pid)
return filtered
def shutdown_mcp_servers(*, scope: Optional[str] = None):
"""Close MCP server connections (in parallel) and stop the background loop.
Each server Task is signalled to exit its own ``async with`` so the anyio
cancel-scope cleanup runs in the Task that opened it. ``scope`` restricts
teardown to one multiplexed profile's servers (its ``/reload-mcp`` must not
kill other profiles') and leaves the shared loop running if anything else
is still connected.
"""
with _core._lock:
selected = [
name for name in _core._servers
if scope is None or _core._server_scope_keys.get(name) == scope
]
servers_snapshot = [_core._servers[name] for name in selected]
# Fast path: nothing to shut down. Still clear the connect-cooldown maps —
# a server that failed to connect is never in ``_servers``, so this is the
# 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()
_core._stop_mcp_loop(only_if_idle=scope is not None)
return
async def _shutdown():
results = await asyncio.gather(
*(server.shutdown() for server in servers_snapshot),
return_exceptions=True,
)
for server, result in zip(servers_snapshot, results):
if isinstance(result, Exception):
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()
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",
)
if future is not None:
try:
future.result(timeout=15)
except BaseException as exc:
logger.debug("Error during MCP shutdown: %s", exc)
# 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()
_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:
"""Best-effort reap of stdio MCP subprocesses: SIGTERM, wait 2s, SIGKILL survivors.
By default only ``_orphan_stdio_pids`` (PIDs that outlived their session
context) are reaped so concurrent cron jobs / live sessions are untouched;
``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)
# 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.
try:
_my_pgid = os.getpgrp()
except (AttributeError, OSError):
_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)
time.sleep(2)
_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():
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,
)
def _stop_mcp_loop_if_idle() -> bool:
"""Stop the MCP loop only when no registered server still owns it.
Probe paths create temporary MCPServerTasks not placed in ``_servers``;
they may clean up an idle loop but must not tear down the process-global
loop under live agent tools, or later calls fail with
``MCP event loop is not running``.
"""
return _core._stop_mcp_loop(only_if_idle=True)
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
cycle before the loop goes away; wait for them here, on their owning loop,
bounded so a task that suppresses cancellation cannot hang process exit.
"""
if timeout is None:
timeout = _core._MCP_LOOP_DRAIN_TIMEOUT
current = asyncio.current_task()
pending = [t for t in asyncio.all_tasks() if t is not current and not t.done()]
if not pending:
return
logger.debug("Draining %d pending task(s) from the MCP loop", len(pending))
for task in pending:
task.cancel()
done, still_pending = await asyncio.wait(pending, timeout=timeout)
for task in done:
if task.cancelled():
continue
try:
task.exception()
except asyncio.CancelledError:
pass
except Exception as exc:
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,
)
async def _drain_and_stop_mcp_loop() -> None:
"""Drain pending tasks, then stop the loop from its owning thread.
Both must run as one loop-owned sequence: a ``loop.stop`` queued separately
by a timed-out caller can overtake the scheduled drain, leaving the drain
coroutine itself pending when the loop is closed.
"""
loop = asyncio.get_running_loop()
try:
await _drain_mcp_loop_tasks(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT)
finally:
loop.call_soon(loop.stop)
+537
View File
@@ -0,0 +1,537 @@
"""Registering a connected (or schema-cached) MCP server's tools into the tool
registry: include/exclude filtering, trust-tier metadata capture, utility-tool
selection, name-collision resolution and the schema-cache write-through."""
import logging
from types import SimpleNamespace
from typing import TYPE_CHECKING, Any, Callable, Dict, List
from tools.mcp_tool_common import _parse_boolish, _core
from tools.mcp_tool_handlers import _make_check_fn, _make_get_prompt_handler, _make_list_prompts_handler, _make_list_resources_handler, _make_read_resource_handler
from tools.mcp_tool_common import _resolve_tool_timeout
from tools.mcp_tool_schema import _UTILITY_CAPABILITY_ATTRS, _UTILITY_CAPABILITY_METHODS, _build_utility_schemas, _normalize_name_filter, matches_name_filter
if TYPE_CHECKING: # pragma: no cover
from tools.mcp_tool import MCPServerTask
logger = logging.getLogger("tools.mcp_tool")
def _normalize_server_trust(value: Any) -> str:
"""Normalize a config ``trust`` value: None -> ``full`` (backward-compatible
default); any unrecognized string -> ``untrusted``, so a misspelled tier
fails closed rather than silently disabling gating."""
if value is None:
return _core._TRUST_FULL
text = str(value).strip().lower()
if text == _core._TRUST_FULL:
return _core._TRUST_FULL
if text == _core._TRUST_UNTRUSTED:
return _core._TRUST_UNTRUSTED
logger.warning(
"MCP trust: unrecognized trust value %r — treating as 'untrusted' "
"(valid values: full, untrusted)", value,
)
return _core._TRUST_UNTRUSTED
def _annotation_read_only_hint(mcp_tool: Any) -> bool:
"""True only when annotations (SDK object or schema-cache dict) carry
``readOnlyHint is True``; anything else is False — unknown metadata means
the tool must be treated as write-capable."""
annotations = getattr(mcp_tool, "annotations", None)
if annotations is None:
return False
if isinstance(annotations, dict):
hint = annotations.get("readOnlyHint")
else:
hint = getattr(annotations, "readOnlyHint", None)
return hint is True
def _record_tool_trust_metadata(
server_name: str, config: dict, tools: List[Any]
) -> None:
"""Capture per-server trust and per-tool readOnlyHint at discovery."""
with _core._lock:
_core._server_trust_levels[server_name] = _normalize_server_trust(
(config or {}).get("trust")
)
hints = _core._tool_read_only_hints.setdefault(server_name, {})
for tool in tools:
name = getattr(tool, "name", None)
if name:
hints[name] = _annotation_read_only_hint(tool)
def _track_mcp_tool_server(tool_name: str, server_name: str) -> None:
"""Remember the exact raw MCP server that registered *tool_name*."""
with _core._lock:
_core._mcp_tool_server_names[tool_name] = server_name
def _forget_mcp_tool_server(tool_name: str) -> None:
"""Forget MCP server provenance for a deregistered tool."""
with _core._lock:
_core._mcp_tool_server_names.pop(tool_name, None)
def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: dict) -> List[dict]:
"""Select utility schemas based on config and server capabilities."""
tools_filter = config.get("tools") or {}
resources_enabled = _parse_boolish(tools_filter.get("resources"), default=True)
prompts_enabled = _parse_boolish(tools_filter.get("prompts"), default=True)
# ``initialize_result.capabilities`` is the source of truth: its sub-objects
# are non-None iff the server advertises that request family. The old
# ``hasattr(server.session, ...)`` gate never filtered anything because
# ClientSession defines all four methods on the class.
advertised_caps = None
init_result = getattr(server, "initialize_result", None)
if init_result is not None:
advertised_caps = getattr(init_result, "capabilities", None)
selected: List[dict] = []
for entry in _build_utility_schemas(server_name):
handler_key = entry["handler_key"]
if handler_key in {"list_resources", "read_resource"} and not resources_enabled:
logger.debug("MCP server '%s': skipping utility '%s' (resources disabled)", server_name, handler_key)
continue
if handler_key in {"list_prompts", "get_prompt"} and not prompts_enabled:
logger.debug("MCP server '%s': skipping utility '%s' (prompts disabled)", server_name, handler_key)
continue
if advertised_caps is not None:
cap_attr = _UTILITY_CAPABILITY_ATTRS[handler_key]
if getattr(advertised_caps, cap_attr, None) is None:
logger.debug(
"MCP server '%s': skipping utility '%s' "
"(server does not advertise '%s' capability)",
server_name,
handler_key,
cap_attr,
)
continue
else:
# Legacy fallback when initialize_result wasn't captured (test
# fixtures, older paths): register every stub, as before.
required_method = _UTILITY_CAPABILITY_METHODS[handler_key]
if not hasattr(server.session, required_method):
logger.debug(
"MCP server '%s': skipping utility '%s' (session lacks %s)",
server_name,
handler_key,
required_method,
)
continue
selected.append(entry)
return selected
def _existing_tool_names() -> List[str]:
"""Return tool names for all currently connected servers."""
names: List[str] = []
for _sname, server in _core._servers.items():
if hasattr(server, "_registered_tool_names"):
names.extend(server._registered_tool_names)
continue
for mcp_tool in server._tools:
schema = _core._convert_mcp_schema(server.name, mcp_tool)
names.append(schema["name"])
# Lazy servers registered from the schema cache have no MCPServerTask yet —
# their tools live only in the registry.
with _core._lock:
lazy_names = [
n
for sname, tool_names in _core._lazy_server_tool_names.items()
if sname not in _core._servers
for n in tool_names
]
names.extend(lazy_names)
return names
# Utility tool key -> handler factory; each takes (server_name, tool_timeout).
_UTILITY_HANDLER_FACTORIES = {
"list_resources": _make_list_resources_handler,
"read_resource": _make_read_resource_handler,
"list_prompts": _make_list_prompts_handler,
"get_prompt": _make_get_prompt_handler,
}
def _make_tool_filter(name: str, config: dict) -> Callable[[str], bool]:
"""Build the include/exclude predicate for a server's tool names.
Rules: ``tools.include`` is a whitelist, ``tools.exclude`` a blacklist;
entries may be exact names or fnmatch globs; include wins over exclude;
``include: []`` is an explicit empty whitelist (register nothing); neither
set registers everything.
"""
tools_filter = config.get("tools") or {}
include_raw = tools_filter.get("include")
include_set = _normalize_name_filter(include_raw, f"mcp_servers.{name}.tools.include")
include_active = isinstance(include_raw, (str, list, tuple, set))
exclude_set = _normalize_name_filter(
tools_filter.get("exclude"), f"mcp_servers.{name}.tools.exclude"
)
def _should_register(tool_name: str) -> bool:
if include_active:
return matches_name_filter(tool_name, include_set)
if exclude_set:
return not matches_name_filter(tool_name, exclude_set)
return True
return _should_register
def _resolve_name_collisions(name: str, candidates: List[dict]):
"""Preflight registry-name collisions among one server's candidates.
Returns ``(unique_candidates, ambiguous_names, shadowed_utilities)``. Exact
duplicate rows (same name + origin) are dropped silently; a generated
utility that normalizes onto a server-native tool's name is shadowed (the
native tool wins); any other multi-origin collision is ambiguous and every
colliding entry is skipped (fail closed).
"""
unique_candidates: List[dict] = []
seen_candidates: set[tuple[str, str]] = set()
origins_by_name: Dict[str, set[str]] = {}
for candidate in candidates:
key = (candidate["registry_name"], candidate["origin"])
if key in seen_candidates:
logger.debug(
"MCP server '%s': duplicate registration candidate %s for '%s'; "
"keeping one",
name,
candidate["origin"],
candidate["registry_name"],
)
continue
seen_candidates.add(key)
unique_candidates.append(candidate)
origins_by_name.setdefault(candidate["registry_name"], set()).add(
candidate["origin"]
)
ambiguous_names: Dict[str, List[str]] = {}
shadowed_utilities: set[tuple[str, str]] = set()
for registry_name, origins in origins_by_name.items():
if len(origins) <= 1:
continue
utility_origins = sorted(
o for o in origins if o.startswith("generated utility ")
)
native_origins = sorted(origins - set(utility_origins))
if len(native_origins) == 1 and utility_origins:
for util_origin in utility_origins:
shadowed_utilities.add((registry_name, util_origin))
logger.info(
"MCP server '%s': generated utility %s normalizes onto "
"server-native %s — keeping the native tool and dropping the "
"utility (the utility only applies when the server has no such "
"tool of its own)",
name,
", ".join(utility_origins),
native_origins[0],
)
continue
ambiguous_names[registry_name] = sorted(origins)
for registry_name, origins in sorted(ambiguous_names.items()):
logger.error(
"MCP server '%s': name normalization collision for '%s' from %s; "
"skipping every colliding entry instead of choosing an arbitrary "
"handler",
name,
registry_name,
", ".join(origins),
)
return unique_candidates, ambiguous_names, shadowed_utilities
def _write_schema_cache(name: str, server: "MCPServerTask", config: dict, should_register) -> None:
"""Write-through: persist the manifest so the next startup can register this
server lazily without spawning it. Never raises."""
try:
from tools.mcp_schema_cache import config_fingerprint, write_cache_entry
tools_payload: List[dict] = []
for mcp_tool in server._tools:
if not should_register(mcp_tool.name):
continue
schema_obj = getattr(mcp_tool, "inputSchema", None)
tools_payload.append({
"name": mcp_tool.name,
"description": mcp_tool.description or "",
"inputSchema": schema_obj if isinstance(schema_obj, dict) else {},
# Persisted so the lazy path trust-gates identically next startup.
"annotations": {
"readOnlyHint": _annotation_read_only_hint(mcp_tool),
},
})
utility_payload = [
{"schema": entry["schema"], "handler_key": entry["handler_key"]}
for entry in _select_utility_schemas(name, server, config)
]
cache_meta = getattr(server, "_list_cache_meta", None) or {}
write_cache_entry(
name,
config_fingerprint(config),
tools=tools_payload,
utility_tools=utility_payload,
ttl_ms=cache_meta.get("ttl_ms"),
cache_scope=cache_meta.get("cache_scope"),
)
except Exception as exc:
logger.debug("MCP schema cache write failed for '%s': %s", name, exc)
def _register_server_tools(name: str, server: "MCPServerTask", config: dict) -> List[str]:
"""Register an already-connected server's tools (plus utility tools) into
the registry; used by initial discovery and list_changed refresh.
Toolset resolution for ``mcp-{server}`` / raw-name aliases derives from the
live registry rather than mutating ``toolsets.TOOLSETS``. Lossy name
normalization can map distinct raw names (``read-file``/``read_file``) to
one registry name; such collisions fail closed — every ambiguous entry is
skipped. Returns the registered prefixed names.
"""
from tools.registry import registry
registered_names: List[str] = []
toolset_name = f"mcp-{name}"
_should_register = _make_tool_filter(name, config)
check_fn = _make_check_fn(name)
candidates: List[dict] = []
# Security boundary: capture trust tier and readOnlyHint NOW, at discovery,
# so the call-time gate classifies from data we control, not re-read
# server-supplied state.
_record_tool_trust_metadata(name, config, server._tools)
for mcp_tool in server._tools:
if not _should_register(mcp_tool.name):
logger.debug(
"MCP server '%s': skipping tool '%s' (filtered by config)",
name,
mcp_tool.name,
)
continue
_core._scan_mcp_description(name, mcp_tool.name, mcp_tool.description or "")
schema = _core._convert_mcp_schema(name, mcp_tool)
candidates.append(
{
"registry_name": schema["name"],
"origin": f"tool {mcp_tool.name!r}",
"schema": schema,
"handler": _core._make_tool_handler(
name, mcp_tool.name, server.tool_timeout
),
"check_fn": check_fn,
}
)
# Generated resource/prompt utility tools share the same namespace as raw
# MCP tools, so they must participate in the same collision preflight.
for entry in _select_utility_schemas(name, server, config):
schema = entry["schema"]
handler_key = entry["handler_key"]
candidates.append(
{
"registry_name": schema["name"],
"origin": f"generated utility {handler_key!r}",
"schema": schema,
"handler": _UTILITY_HANDLER_FACTORIES[handler_key](
name, server.tool_timeout
),
"check_fn": check_fn,
}
)
unique_candidates, ambiguous_names, shadowed_utilities = _resolve_name_collisions(name, candidates)
for candidate in unique_candidates:
registry_name = candidate["registry_name"]
if registry_name in ambiguous_names:
continue
if (registry_name, candidate["origin"]) in shadowed_utilities:
continue
existing_toolset = registry.get_toolset_for_tool(registry_name)
if existing_toolset and existing_toolset != toolset_name:
if existing_toolset.startswith("mcp-"):
logger.error(
"MCP server '%s': %s normalizes to '%s', already owned by "
"MCP toolset '%s' — skipping to preserve the existing owner",
name,
candidate["origin"],
registry_name,
existing_toolset,
)
else:
logger.warning(
"MCP server '%s': %s (→ '%s') collides with built-in tool "
"in toolset '%s' — skipping to preserve built-in",
name,
candidate["origin"],
registry_name,
existing_toolset,
)
continue
registry.register(
name=registry_name,
toolset=toolset_name,
schema=candidate["schema"],
handler=candidate["handler"],
check_fn=candidate["check_fn"],
is_async=False,
description=candidate["schema"]["description"],
scope=_core._server_registry_scope(name),
)
# The pre-check above is advisory only. Multiple servers connect in
# parallel, so ToolRegistry.register() is the atomic ownership gate.
if registry.get_toolset_for_tool(registry_name) != toolset_name:
logger.error(
"MCP server '%s': registration of %s as '%s' was rejected by "
"the registry; skipping provenance/count updates",
name,
candidate["origin"],
registry_name,
)
continue
_core._track_mcp_tool_server(registry_name, name)
registered_names.append(registry_name)
if registered_names:
registry.register_toolset_alias(name, toolset_name)
_write_schema_cache(name, server, config, _should_register)
return registered_names
class _CachedMCPTool:
"""Minimal stand-in for MCP Tool objects loaded from the schema cache."""
__slots__ = ("name", "description", "inputSchema")
def __init__(self, name: str, description: str, inputSchema: dict):
self.name = name
self.description = description
self.inputSchema = inputSchema or {}
def _register_from_cache_sync(name: str, config: dict, entry: dict) -> List[str]:
"""Lazy startup: register a server's tools from a cached manifest with no
child process. The first real call routes through
``_get_connected_server_for_call`` -> ``_ensure_lazy_server_connected``."""
from tools.registry import registry
from tools.mcp_schema_cache import (
config_fingerprint,
tools_from_cache_entry,
utility_tools_from_cache_entry,
)
registered_names: List[str] = []
toolset_name = f"mcp-{name}"
fingerprint = config_fingerprint(config)
tool_timeout = _resolve_tool_timeout(config)
_should_register = _make_tool_filter(name, config)
check_fn = _make_check_fn(name)
# Record trust metadata before registration so the call-time gate is
# identical whether the server was spawned live or registered from cache.
# Missing "annotations" in older cache files fails closed to write-capable.
cached_tool_objs = [
SimpleNamespace(
name=raw.get("name"),
annotations=raw.get("annotations")
if isinstance(raw.get("annotations"), dict) else None,
)
for raw in tools_from_cache_entry(entry)
if isinstance(raw, dict) and raw.get("name")
]
_record_tool_trust_metadata(name, config, cached_tool_objs)
for raw in tools_from_cache_entry(entry):
if not isinstance(raw, dict):
continue
raw_name = raw.get("name")
if not raw_name or not _should_register(raw_name):
continue
raw_schema = raw.get("inputSchema")
mcp_tool = _CachedMCPTool(
raw_name,
raw.get("description") or "",
raw_schema if isinstance(raw_schema, dict) else {},
)
# Defense-in-depth: the cache file is user-writable JSON, so apply the
# same injection scan as eager discovery.
_core._scan_mcp_description(name, mcp_tool.name, mcp_tool.description or "")
schema = _core._convert_mcp_schema(name, mcp_tool)
registry_name = schema["name"]
existing_toolset = registry.get_toolset_for_tool(registry_name)
if existing_toolset and existing_toolset != toolset_name:
logger.warning(
"MCP server '%s' (lazy): cached tool '%s' collides with "
"toolset '%s' — skipping",
name, registry_name, existing_toolset,
)
continue
registry.register(
name=registry_name,
toolset=toolset_name,
schema=schema,
handler=_core._make_tool_handler(name, raw_name, tool_timeout),
check_fn=check_fn,
is_async=False,
description=schema["description"],
scope=_core._mcp_registry_scope(),
)
if registry.get_toolset_for_tool(registry_name) != toolset_name:
continue
_core._track_mcp_tool_server(registry_name, name)
registered_names.append(registry_name)
for raw in utility_tools_from_cache_entry(entry):
if not isinstance(raw, dict):
continue
schema = raw.get("schema")
handler_key = raw.get("handler_key")
if not isinstance(schema, dict) or handler_key not in _UTILITY_HANDLER_FACTORIES:
continue
util_name = schema.get("name") or ""
if not util_name:
continue
existing_toolset = registry.get_toolset_for_tool(util_name)
if existing_toolset and existing_toolset != toolset_name:
continue
registry.register(
name=util_name,
toolset=toolset_name,
schema=schema,
handler=_UTILITY_HANDLER_FACTORIES[handler_key](name, tool_timeout),
check_fn=check_fn,
is_async=False,
description=schema.get("description") or "",
scope=_core._mcp_registry_scope(),
)
if registry.get_toolset_for_tool(util_name) != toolset_name:
continue
_core._track_mcp_tool_server(util_name, name)
registered_names.append(util_name)
if registered_names:
registry.register_toolset_alias(name, toolset_name)
with _core._lock:
_core._lazy_server_configs[name] = dict(config)
_core._lazy_server_fingerprints[name] = fingerprint
_core._lazy_server_tool_names[name] = list(registered_names)
logger.info(
"MCP server '%s' (lazy): registered %d tool(s) from schema cache",
name, len(registered_names),
)
return registered_names
+521
View File
@@ -0,0 +1,521 @@
"""MCP client-side handlers for server-initiated requests: sampling (sampling/createMessage, text and tool-use results) and elicitation. Split from tools/mcp_tool.py."""
import asyncio
import json
import logging
import time
from typing import List, Optional
from tools.mcp_tool_common import _MISSING, _exc_str, _safe_numeric, _sanitize_error, mcp_field, _core
from tools.mcp_tool_schema import _normalize_mcp_input_schema
logger = logging.getLogger("tools.mcp_tool")
class SamplingHandler:
"""Handles sampling/createMessage requests for one MCP server.
Deprecated upstream (MCP 2026-07-28, SEP-2577, 12-month window): stays fully
functional because handshake-era servers still issue it, but do NOT grow new
capability here — modern servers use MRTR, handled by the SDK session layer.
Callable; passed to ``ClientSession`` as ``sampling_callback``. All state
(rate-limit timestamps, metrics, tool-loop counter) is per instance. Runs on
the MCP background loop; the sync LLM call is offloaded via ``asyncio.to_thread``.
"""
_STOP_REASON_MAP = {"stop": "endTurn", "length": "maxTokens", "tool_calls": "toolUse"}
def __init__(self, server_name: str, config: dict):
self.server_name = server_name
self.max_rpm = _safe_numeric(config.get("max_rpm", 10), 10, int)
self.timeout = _safe_numeric(config.get("timeout", 30), 30, float)
self.max_tokens_cap = _safe_numeric(config.get("max_tokens_cap", 4096), 4096, int)
self.max_tool_rounds = _safe_numeric(
config.get("max_tool_rounds", 5), 5, int, minimum=0,
)
self.model_override = config.get("model")
self.allowed_models = config.get("allowed_models", [])
_log_levels = {"debug": logging.DEBUG, "info": logging.INFO, "warning": logging.WARNING}
self.audit_level = _log_levels.get(
str(config.get("log_level", "info")).lower(), logging.INFO,
)
self._rate_timestamps: List[float] = []
self._tool_loop_count = 0
self.metrics = {"requests": 0, "errors": 0, "tokens_used": 0, "tool_use_count": 0}
def _check_rate_limit(self) -> bool:
"""Sliding-window (60s) limiter; True if the request is allowed."""
now = time.time()
window = now - 60
self._rate_timestamps[:] = [t for t in self._rate_timestamps if t > window]
if len(self._rate_timestamps) >= self.max_rpm:
return False
self._rate_timestamps.append(now)
return True
def _resolve_model(self, preferences) -> Optional[str]:
"""Config override > server hint > None (use default)."""
if self.model_override:
return self.model_override
if preferences and hasattr(preferences, "hints") and preferences.hints:
for hint in preferences.hints:
if hasattr(hint, "name") and hint.name:
return hint.name
return None
@staticmethod
def _extract_tool_result_text(block) -> str:
"""Extract text from a ToolResultContent block."""
if not hasattr(block, "content") or block.content is None:
return ""
items = block.content if isinstance(block.content, list) else [block.content]
return "\n".join(item.text for item in items if hasattr(item, "text"))
def _convert_messages(self, params) -> List[dict]:
"""Convert MCP SamplingMessages to OpenAI format.
Uses ``msg.content_as_list`` when the SDK provides it; dispatches per
block by duck-typing.
"""
# A tool-use id is the discriminator for a tool *result* block; it must be
# read under both spellings (mcp_field) — on mcp 2.x a bare
# ``hasattr(b, "toolUseId")`` is False, silently dropping tool results.
def _tool_use_id(block):
return mcp_field(block, "tool_use_id", "toolUseId", _MISSING)
def _is_tool_use(block):
return hasattr(block, "name") and hasattr(block, "input")
messages: List[dict] = []
for msg in params.messages:
blocks = msg.content_as_list if hasattr(msg, "content_as_list") else (
msg.content if isinstance(msg.content, list) else [msg.content]
)
tool_results = [b for b in blocks if _tool_use_id(b) is not _MISSING]
tool_uses = [
b for b in blocks
if _is_tool_use(b) and _tool_use_id(b) is _MISSING
]
content_blocks = [
b for b in blocks
if _tool_use_id(b) is _MISSING and not _is_tool_use(b)
]
for tr in tool_results:
messages.append({
"role": "tool",
"tool_call_id": _tool_use_id(tr),
"content": self._extract_tool_result_text(tr),
})
if tool_uses:
tc_list = []
for tu in tool_uses:
tc_list.append({
"id": getattr(tu, "id", f"call_{len(tc_list)}"),
"type": "function",
"function": {
"name": tu.name,
"arguments": json.dumps(tu.input, ensure_ascii=False) if isinstance(tu.input, dict) else str(tu.input),
},
})
msg_dict: dict = {"role": msg.role, "tool_calls": tc_list}
text_parts = [b.text for b in content_blocks if hasattr(b, "text")]
if text_parts:
msg_dict["content"] = "\n".join(text_parts)
messages.append(msg_dict)
elif content_blocks:
# Pure text/image content.
if len(content_blocks) == 1 and hasattr(content_blocks[0], "text"):
messages.append({"role": msg.role, "content": content_blocks[0].text})
else:
parts = []
for block in content_blocks:
block_mime = mcp_field(
block, "mime_type", "mimeType", _MISSING
)
if hasattr(block, "text"):
parts.append({"type": "text", "text": block.text})
elif hasattr(block, "data") and block_mime is not _MISSING:
parts.append({
"type": "image_url",
"image_url": {"url": f"data:{block_mime};base64,{block.data}"},
})
else:
logger.warning(
"Unsupported sampling content block type: %s (skipped)",
type(block).__name__,
)
if parts:
messages.append({"role": msg.role, "content": parts})
return messages
@staticmethod
def _error(message: str, code: int = -1):
"""Return ErrorData (MCP spec) or raise as fallback."""
if _core._MCP_SAMPLING_TYPES:
return _core.ErrorData(code=code, message=message)
raise Exception(message)
def _build_tool_use_result(self, choice, response):
"""Build a CreateMessageResultWithTools from an LLM tool_calls response."""
self.metrics["tool_use_count"] += 1
# Tool-loop governance.
if self.max_tool_rounds == 0:
self._tool_loop_count = 0
return self._error(
f"Tool loops disabled for server '{self.server_name}' (max_tool_rounds=0)"
)
self._tool_loop_count += 1
if self._tool_loop_count > self.max_tool_rounds:
self._tool_loop_count = 0
return self._error(
f"Tool loop limit exceeded for server '{self.server_name}' "
f"(max {self.max_tool_rounds} rounds)"
)
content_blocks = []
for tc in choice.message.tool_calls:
args = tc.function.arguments
if isinstance(args, str):
try:
parsed = json.loads(args)
except (json.JSONDecodeError, ValueError):
logger.warning(
"MCP server '%s': malformed tool_calls arguments "
"from LLM (wrapping as raw): %.100s",
self.server_name, args,
)
parsed = {"_raw": args}
else:
parsed = args if isinstance(args, dict) else {"_raw": str(args)}
content_blocks.append(_core.ToolUseContent(
type="tool_use",
id=tc.id,
name=tc.function.name,
input=parsed,
))
logger.log(
self.audit_level,
"MCP server '%s' sampling response: model=%s, tokens=%s, tool_calls=%d",
self.server_name, response.model,
getattr(getattr(response, "usage", None), "total_tokens", "?"),
len(content_blocks),
)
return _core.CreateMessageResultWithTools(
role="assistant",
content=content_blocks,
model=response.model,
stopReason="toolUse",
)
def _build_text_result(self, choice, response):
"""Build a CreateMessageResult from a normal text response (resets the tool loop)."""
self._tool_loop_count = 0
response_text = choice.message.content or ""
logger.log(
self.audit_level,
"MCP server '%s' sampling response: model=%s, tokens=%s",
self.server_name, response.model,
getattr(getattr(response, "usage", None), "total_tokens", "?"),
)
return _core.CreateMessageResult(
role="assistant",
content=_core.TextContent(type="text", text=_sanitize_error(response_text)),
model=response.model,
stopReason=self._STOP_REASON_MAP.get(choice.finish_reason, "endTurn"),
)
def session_kwargs(self) -> dict:
"""Kwargs to pass to ClientSession for sampling support."""
return {
"sampling_callback": self,
"sampling_capabilities": _core.SamplingCapability(
tools=_core.SamplingToolsCapability(),
),
}
async def __call__(self, context, params):
"""SDK sampling callback (``SamplingFnT``). Returns CreateMessageResult,
CreateMessageResultWithTools, or ErrorData."""
if not self._check_rate_limit():
logger.warning(
"MCP server '%s' sampling rate limit exceeded (%d/min)",
self.server_name, self.max_rpm,
)
self.metrics["errors"] += 1
return self._error(
f"Sampling rate limit exceeded for server '{self.server_name}' "
f"({self.max_rpm} requests/minute)"
)
model = self._resolve_model(
mcp_field(params, "model_preferences", "modelPreferences")
)
from agent.auxiliary_client import call_llm
resolved_model = model or self.model_override or ""
if self.allowed_models and resolved_model and resolved_model not in self.allowed_models:
logger.warning(
"MCP server '%s' requested model '%s' not in allowed_models",
self.server_name, resolved_model,
)
self.metrics["errors"] += 1
return self._error(
f"Model '{resolved_model}' not allowed for server "
f"'{self.server_name}'. Allowed: {', '.join(self.allowed_models)}"
)
messages = self._convert_messages(params)
system_prompt = mcp_field(params, "system_prompt", "systemPrompt")
if system_prompt:
messages.insert(0, {"role": "system", "content": system_prompt})
max_tokens = min(
mcp_field(params, "max_tokens", "maxTokens", self.max_tokens_cap),
self.max_tokens_cap,
)
call_temperature = None
if hasattr(params, "temperature") and params.temperature is not None:
call_temperature = params.temperature
# Forward server-provided tools.
call_tools = None
server_tools = getattr(params, "tools", None)
if server_tools:
call_tools = [
{
"type": "function",
"function": {
"name": getattr(t, "name", ""),
"description": getattr(t, "description", "") or "",
"parameters": _normalize_mcp_input_schema(
mcp_field(t, "input_schema", "inputSchema")
),
},
}
for t in server_tools
]
logger.log(
self.audit_level,
"MCP server '%s' sampling request: model=%s, max_tokens=%d, messages=%d",
self.server_name, resolved_model, max_tokens, len(messages),
)
# Offload the sync LLM call so the MCP loop is not blocked.
def _sync_call():
return call_llm(
task="mcp",
model=resolved_model or None,
messages=messages,
temperature=call_temperature,
max_tokens=max_tokens,
tools=call_tools,
timeout=self.timeout,
)
try:
response = await asyncio.wait_for(
asyncio.to_thread(_sync_call), timeout=self.timeout,
)
except asyncio.TimeoutError:
self.metrics["errors"] += 1
return self._error(
f"Sampling LLM call timed out after {self.timeout}s "
f"for server '{self.server_name}'"
)
except Exception as exc:
self.metrics["errors"] += 1
return self._error(
f"Sampling LLM call failed: {_sanitize_error(_exc_str(exc))}"
)
# Empty choices happen on content filtering / provider errors.
if not getattr(response, "choices", None):
self.metrics["errors"] += 1
return self._error(
f"LLM returned empty response (no choices) for server "
f"'{self.server_name}'"
)
choice = response.choices[0]
self.metrics["requests"] += 1
total_tokens = getattr(getattr(response, "usage", None), "total_tokens", 0)
if isinstance(total_tokens, int):
self.metrics["tokens_used"] += total_tokens
if (
choice.finish_reason == "tool_calls"
and hasattr(choice.message, "tool_calls")
and choice.message.tool_calls
):
return self._build_tool_use_result(choice, response)
return self._build_text_result(choice, response)
def _format_elicitation_schema_summary(schema: dict, server_name: str) -> str:
"""Render a flat-object requested_schema as a human-readable field list
(names, types, descriptions) so the user knows what they're approving."""
props = schema.get("properties") if isinstance(schema, dict) else None
if not isinstance(props, dict) or not props:
return f"Approval requested by MCP server '{server_name}'."
lines = [f"Fields requested by MCP server '{server_name}':"]
for field_name, field_spec in props.items():
field_type = ""
field_desc = ""
if isinstance(field_spec, dict):
field_type = str(field_spec.get("type", "") or "")
field_desc = str(field_spec.get("description", "") or "")
suffix = f" ({field_type})" if field_type else ""
if field_desc:
lines.append(f" - {field_name}{suffix}: {field_desc}")
else:
lines.append(f" - {field_name}{suffix}")
return "\n".join(lines)
class ElicitationHandler:
"""Handles ``elicitation/create`` requests for one MCP server.
Callable; passed to ``ClientSession`` as ``elicitation_callback``. Form-mode
requests route through Hermes' approval system (CLI, TUI, Telegram, ...);
URL-mode is declined as unsupported. Fail-closed: any timeout, exception,
or unexpected state returns decline/cancel, never a silent accept.
"""
# asyncio-side safety net over the approval's own input() timeout so the
# MCP loop never blocks indefinitely if the inner timeout is bypassed.
_OUTER_TIMEOUT_GRACE_SECONDS = 5
def __init__(self, server_name: str, config: dict, owner: Optional["MCPServerTask"] = None):
self.server_name = server_name
# Default 5 min mirrors the gateway approval default so async surfaces
# (Telegram, Slack) have time to respond.
self.timeout = _safe_numeric(config.get("timeout", 300), 300, float)
# Back-reference for the agent's contextvars snapshot; optional so the
# handler stays unit-testable in isolation.
self.owner = owner
self.metrics = {
"requests": 0,
"accepted": 0,
"declined": 0,
"errors": 0,
}
def session_kwargs(self) -> dict:
"""Kwargs to pass to ClientSession for elicitation support."""
return {"elicitation_callback": self}
async def __call__(self, context, params):
"""SDK elicitation callback (``ElicitationFnT``). Returns ElicitResult or ErrorData."""
self.metrics["requests"] += 1
# URL-mode (OAuth, payment) would need a browser + waiting for
# notifications/elicitation/complete — not implemented; decline cleanly.
mode = getattr(params, "mode", "form")
if mode == "url":
logger.info(
"MCP server '%s' requested URL-mode elicitation; "
"declining (URL-mode elicitation not implemented)",
self.server_name,
)
self.metrics["declined"] += 1
return _core.ElicitResult(action="decline")
message = getattr(params, "message", "") or (
f"MCP server '{self.server_name}' is requesting your approval"
)
# ``requestedSchema`` on mcp 1.x, ``requested_schema`` on 2.0 (pydantic
# aliases don't apply to attribute access) — read both or the user is
# asked to approve without seeing which fields the server wants.
schema = (
getattr(params, "requestedSchema", None)
or getattr(params, "requested_schema", None)
or {}
)
description = _format_elicitation_schema_summary(schema, self.server_name)
logger.info(
"MCP server '%s' elicitation request: %s",
self.server_name, _sanitize_error(message)[:200],
)
# Lazy import avoids import-order coupling with early-bootstrap tools.approval.
try:
from tools.approval import request_elicitation_consent
except Exception as exc: # pragma: no cover -- defensive
logger.error(
"MCP server '%s' elicitation: approval system unavailable: %s",
self.server_name, exc,
)
self.metrics["errors"] += 1
return _core.ElicitResult(action="decline")
# Offload the sync consent flow to a thread — inline it would freeze the
# MCP loop and every other RPC on this session. The recv-loop task does
# NOT inherit the agent's contextvars, so replay the snapshot captured on
# owner._pending_call_context for gateway-platform detection.
captured = getattr(self.owner, "_pending_call_context", None) if self.owner else None
def _invoke_consent() -> str:
if captured is None:
return request_elicitation_consent(
message,
description,
timeout_seconds=int(self.timeout),
surface=f"mcp-elicitation/{self.server_name}",
)
# Context.run executes a context once — copy so multiple
# elicitations within one tool call work.
return captured.copy().run(
request_elicitation_consent,
message,
description,
timeout_seconds=int(self.timeout),
surface=f"mcp-elicitation/{self.server_name}",
)
try:
answer = await asyncio.wait_for(
asyncio.to_thread(_invoke_consent),
timeout=self.timeout + self._OUTER_TIMEOUT_GRACE_SECONDS,
)
except asyncio.TimeoutError:
logger.warning(
"MCP server '%s' elicitation timed out after %ds",
self.server_name, int(self.timeout),
)
self.metrics["errors"] += 1
return _core.ElicitResult(action="cancel")
except Exception as exc:
logger.error(
"MCP server '%s' elicitation failed: %s",
self.server_name, exc, exc_info=True,
)
self.metrics["errors"] += 1
return _core.ElicitResult(action="decline")
if answer == "accept":
self.metrics["accepted"] += 1
return _core.ElicitResult(action="accept", content={})
if answer == "cancel":
self.metrics["errors"] += 1
return _core.ElicitResult(action="cancel")
self.metrics["declined"] += 1
return _core.ElicitResult(action="decline")
+307
View File
@@ -0,0 +1,307 @@
"""MCP tool schema conversion and naming: JSON-schema normalisation for provider
compatibility, mcp__server__tool naming, utility-tool schemas, include/exclude
filters and description injection scanning."""
import logging
import fnmatch
import re
from typing import Any, List
from tools.ansi_strip import strip_unicode_tags
from tools.mcp_tool_common import mcp_field
logger = logging.getLogger("tools.mcp_tool")
# Prompt-injection indicators in MCP tool descriptions. WARNING-level only:
# log but never block, since false positives would break legitimate servers.
_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"),
]
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)
if findings:
logger.warning(
"MCP server '%s' tool '%s': suspicious description content — %s. "
"Description: %.200s",
server_name, tool_name, "; ".join(findings),
description,
)
return findings
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``.
"""
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
normalized = _rewrite_local_refs(schema)
normalized = _strip_nullable_union(normalized)
normalized = _collapse_const_unions(normalized)
normalized = _repair_object_shape(normalized)
if not isinstance(normalized, dict):
return {"type": "object", "properties": {}}
if normalized.get("type") == "object" and "properties" not in normalized:
normalized = {**normalized, "properties": {}}
return normalized
def sanitize_mcp_name_component(value: str) -> str:
"""Replace every char outside ``[A-Za-z0-9_]`` with ``_`` (hyphens included,
the historical behavior) so generated names pass provider validation."""
return re.sub(r"[^A-Za-z0-9_]", "_", str(value or ""))
# ``mcp__<server>__<tool>``: the convention shared by Claude Code, Codex and
# OpenCode. The double underscore disambiguates the server/tool boundary even
# when either contains underscores, and matches the Anthropic-OAuth wire form.
MCP_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}"
def _convert_mcp_schema(server_name: str, mcp_tool) -> dict:
"""Convert an MCP ``Tool`` (``.input_schema``, or ``.inputSchema`` before
mcp 2.0) to a ``registry.register(schema=...)`` dict."""
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")
),
}
def _build_utility_schemas(server_name: str) -> List[dict]:
"""Schemas for the resource/prompt utility tools as ``{schema, handler_key}`` dicts."""
return [
{
"schema": {
"name": mcp_prefixed_tool_name(server_name, "list_resources"),
"description": f"List available resources from MCP server '{server_name}'",
"parameters": {
"type": "object",
"properties": {},
},
},
"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",
},
]
def _normalize_name_filter(value: Any, label: str) -> set[str]:
"""Normalize include/exclude config to a set of exact names or fnmatch globs."""
if value is None:
return set()
if isinstance(value, str):
return {value}
if isinstance(value, (list, tuple, set)):
return {str(item) for item in value}
logger.warning("MCP config %s must be a string or list of strings; ignoring %r", label, value)
return set()
def matches_name_filter(tool_name: str, patterns: set[str]) -> bool:
"""True if ``tool_name`` matches any entry: exact names literally, entries
with ``*``/``?``/``[`` as case-sensitive globs (same semantics as
``approvals.deny``). Exact membership is checked first so big lists stay O(1)."""
if not patterns:
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
)
_UTILITY_CAPABILITY_METHODS = {
"list_resources": "list_resources",
"read_resource": "read_resource",
"list_prompts": "list_prompts",
"get_prompt": "get_prompt",
}
# Utility handler -> capability key that must be non-None on the server's
# ``initialize`` response for the handler to be registered. Without this gate a
# tools-only server got all four stubs and every call returned JSON-RPC -32601,
# making the model conclude the server was broken.
_UTILITY_CAPABILITY_ATTRS = {
"list_resources": "resources",
"read_resource": "resources",
"list_prompts": "prompts",
"get_prompt": "prompts",
}
+676
View File
@@ -0,0 +1,676 @@
"""Transport bring-up for MCPServerTask: stdio spawn (OSV preflight, watchdog wrap, child PID ledger), Streamable HTTP / SSE connect (preflight, identity header, client certs, OAuth), protocol negotiation and initial tool discovery. Split from tools/mcp_tool.py."""
import logging
import asyncio
import os
from typing import Dict, Optional
from tools.mcp_tool_config import _wrap_command_with_watchdog
from tools.mcp_tool_errors import NonMcpEndpointError, _apply_identity_header, _handshake_rejected_as_modern, _make_redirect_header_stripper, _resolve_client_cert
from tools.mcp_tool_lifecycle import _filter_mcp_children, _orphan_stdio_pid_servers, _orphan_stdio_pids, _stdio_pgids, _stdio_pids
from tools.mcp_tool_common import _core
logger = logging.getLogger("tools.mcp_tool")
class MCPServerTransportMixin:
"""Methods of :class:`tools.mcp_tool.MCPServerTask` (mixed in; relies on its attributes)."""
__slots__ = ()
def _advertises_tools(self) -> bool:
"""Whether the server advertises the ``tools`` capability.
Prompt-/resource-only servers omit it, and ``tools/list`` against them
raises ``MCPError(-32601)``. True when no capability info was captured
(legacy fallback: keep the old always-call-list_tools behavior).
"""
init_result = self.initialize_result
caps = getattr(init_result, "capabilities", None) if init_result is not None else None
if caps is None:
return True
return getattr(caps, "tools", None) is not None
async def _negotiate_session(self, session, connect_timeout: float):
"""Negotiate the protocol era (``initialize`` vs ``server/discover``) and return its result.
Per-server ``protocol`` key: ``auto`` (default) tries the legacy handshake
FIRST and falls back to ``server/discover`` only when the server signals
modern-only (-32022 / initialize -32601) — the reverse of the SDK's
discover-first mode, on purpose: zero extra round-trips for the handshake-era
servers that dominate today. ``stateless`` probes discover first (one legacy
retry on error); ``legacy`` is handshake only, no fallback. Both result
types expose ``.capabilities``, so downstream gates work on either.
"""
mode = str((self._config or {}).get("protocol", "auto")).lower().strip()
if mode in ("stateless", "modern", "2026-07-28"):
try:
return await asyncio.wait_for(
session.discover(), timeout=connect_timeout
)
except asyncio.TimeoutError:
raise
except Exception as exc:
logger.info(
"MCP server '%s': server/discover rejected (%s) despite "
"protocol=%s — falling back to the legacy handshake",
self.name, exc, mode,
)
return await asyncio.wait_for(
session.initialize(), timeout=connect_timeout
)
if mode in ("legacy", "handshake"):
return await asyncio.wait_for(
session.initialize(), timeout=connect_timeout
)
if mode != "auto":
logger.warning(
"MCP server '%s': unknown protocol=%r — treating as 'auto' "
"(valid: auto, stateless, legacy)", self.name, mode,
)
try:
return await asyncio.wait_for(
session.initialize(), timeout=connect_timeout
)
except asyncio.TimeoutError:
raise
except Exception as exc:
if not _handshake_rejected_as_modern(exc):
raise
if not hasattr(session, "discover"):
# mcp 1.x has no server/discover client — nothing to fall back to.
raise
logger.info(
"MCP server '%s': legacy handshake rejected (%s) — "
"retrying via server/discover (2026-07-28 stateless server)",
self.name, exc,
)
return await asyncio.wait_for(
session.discover(), timeout=connect_timeout
)
async def _serve_session(self, session, connect_timeout: float,
label: str = "", mark_lifecycle: bool = False) -> str:
"""Handshake, discover, publish readiness, then serve until a lifecycle event.
Clears stale breaker state from a prior outage, but leaves the session
UNPROVEN: a completed handshake is not proof of health (flapping
transports handshake fine and drop moments later); only keepalive or
tool-call success clears the reconnect budget.
"""
self.initialize_result = await self._negotiate_session(session, connect_timeout)
self.session = session
if mark_lifecycle:
self._mark_lifecycle_started()
await self._discover_tools()
self._ready.set()
self._ever_connected = True
_core._reset_server_error(self.name)
self._session_proven = False
reason = await self._wait_for_lifecycle_event()
if label and reason == "reconnect":
logger.info(
"MCP server '%s': reconnect requested — tearing down %s session",
self.name, label,
)
return reason
async def _run_stdio(self, config: dict):
"""Run the server using stdio transport."""
if config.get("identity_header") is not None:
# No headers on stdio — warn so a copy-pasted HTTP block doesn't mislead.
logger.warning(
"MCP server '%s': identity_header is only supported on "
"HTTP/SSE transports — ignored for stdio servers", self.name,
)
if not _core._ensure_mcp_sdk():
raise ImportError(
f"MCP server '{self.name}' requires the 'mcp' Python SDK, but "
"it is not installed. Run `hermes setup` to install MCP support, "
"then retry."
)
command = config.get("command")
args = config.get("args", [])
user_env = config.get("env")
if not command:
raise ValueError(
f"MCP server '{self.name}' has no 'command' in config"
)
safe_env = _core._build_safe_env(user_env)
command, safe_env = _core._resolve_stdio_command(command, safe_env)
# OSV malware preflight: off-loop (blocking HTTPS) with a wall-clock bound
# so a stalled handshake can't freeze discovery; fail-open on timeout.
# Must run against the REAL command/args — the watchdog wrap below rewrites
# argv to the supervisor, which would turn the check into a no-op.
from tools.osv_check import check_package_for_malware
try:
malware_error = await asyncio.wait_for(
asyncio.to_thread(check_package_for_malware, command, args),
timeout=_core._OSV_MALWARE_CHECK_TIMEOUT_S,
)
except asyncio.TimeoutError:
logger.warning(
"MCP server '%s': OSV malware preflight timed out after %.0fs "
"(network slow/unreachable) — proceeding without the check.",
self.name, _core._OSV_MALWARE_CHECK_TIMEOUT_S,
)
malware_error = None
if malware_error:
raise ValueError(
f"MCP server '{self.name}': {malware_error}"
)
# Parent-death watchdog: an ungraceful Hermes exit (kill -9, crash) can't
# leave the child and its descendants running. Clean-exit reaping is
# unchanged. POSIX-only (process groups); no-op elsewhere. AFTER the OSV
# preflight so the check inspects the real package.
command, args = _wrap_command_with_watchdog(command, args)
server_params = _core.StdioServerParameters(
command=command,
args=args,
env=safe_env if safe_env else None,
cwd=config.get("cwd"),
# Windows pipes can split non-UTF-8 bytes at chunk boundaries;
# substitute U+FFFD instead of raising UnicodeDecodeError.
encoding_error_handler="replace",
)
sampling_kwargs = self._sampling.session_kwargs() if self._sampling else {}
if self._elicitation:
sampling_kwargs.update(self._elicitation.session_kwargs())
if _core._MCP_NOTIFICATION_TYPES and _core._MCP_MESSAGE_HANDLER_SUPPORTED:
sampling_kwargs["message_handler"] = self._make_message_handler()
if _core._MCP_LOGGING_CALLBACK_SUPPORTED:
sampling_kwargs["logging_callback"] = self._make_logging_callback()
# Reap orphans from prior failed attempts before spawning, else each
# reconnect retry piles up zombie pairs. Unscoped on purpose (also reaps
# orphans of servers that never reconnect). Worker thread: the reaper
# blocks up to 2s (SIGTERM → wait → SIGKILL) and would stall the loop.
await asyncio.to_thread(_core._kill_orphaned_mcp_children)
# Snapshot child PIDs before spawning so the new one can be identified.
pids_before = _core._snapshot_child_pids()
new_pids: set = set()
# Route subprocess stderr to ~/.hermes/logs/mcp-stderr.log so server
# banners don't land on the user's TTY and corrupt the TUI.
_core._write_stderr_log_header(self.name)
_errlog = _core._get_mcp_stderr_log()
try:
async with _core.stdio_client(server_params, errlog=_errlog) as (
read_stream,
write_stream,
):
# Capture the new PID for force-kill cleanup, filtering non-MCP
# children (slash_worker, LSP servers) that race into the snapshot
# window: they share the TUI parent's pgid, so leaking them into
# _stdio_pgids makes the shutdown killpg() kill the TUI itself.
new_pids = _filter_mcp_children(
_core._snapshot_child_pids() - pids_before
)
if new_pids:
# Capture pgid while the child is alive — getpgid fails once it
# exits, and the sweep needs it to reach reparented descendants.
new_pgids: Dict[int, int] = {}
for _pid in new_pids:
try:
new_pgids[_pid] = os.getpgid(_pid)
except (AttributeError, ProcessLookupError, OSError):
# AttributeError: Windows; ProcessLookupError: already exited.
pass
with _core._lock:
for _pid in new_pids:
_stdio_pids[_pid] = self.name
_stdio_pgids.update(new_pgids)
# Machine spawn ledger so startup sweeps can reap orphans after
# an unclean parent exit. Best-effort — never break startup.
for _pid in new_pids:
try:
from hermes_cli.process_identity import register_child
register_child(_pid, "mcp-helper")
except Exception:
logger.debug(
"spawn-ledger register_child failed for MCP "
"helper pid %s",
_pid,
exc_info=True,
)
# Tracked on the connection so in-flight calls fail fast when
# the subprocess dies.
self._stdio_child_pids = set(new_pids)
async with _core.ClientSession(
read_stream, write_stream, **sampling_kwargs
) as session:
# Bound the handshake: ``connect_timeout`` only bounds the
# caller's ``.result()`` wait, not this coroutine. A server that
# never answers ``initialize`` would otherwise hang here forever,
# the ``finally`` below would never run, and the child + pipes
# would leak on every retry until EMFILE.
connect_timeout = float(
config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT)
)
return await self._serve_session(
session, connect_timeout, mark_lifecycle=True
)
finally:
# Runs on clean exit, exceptions AND cancellation. Any spawned PID
# still alive means SDK teardown failed (common on cancel mid-way on
# Linux, where setsid() children escape the cgroup): mark it orphaned
# for the next cleanup sweep.
if new_pids:
from gateway.status import _pid_exists
_killpg = getattr(os, "killpg", None)
with _core._lock:
for _pid in new_pids:
_stdio_pids.pop(_pid, None)
for pid in new_pids:
# ``os.kill(pid, 0)`` is NOT a no-op on Windows; use the
# cross-platform check.
pid_alive = _pid_exists(pid)
pgroup_alive = False
pgid = _stdio_pgids.get(pid)
if not pid_alive and pgid is not None and _killpg is not None:
# Child exited but descendants may remain in its pgroup;
# signal 0 succeeds iff any member is alive.
try:
_killpg(pgid, 0)
pgroup_alive = True
except (ProcessLookupError, PermissionError, OSError):
pgroup_alive = False
if pid_alive or pgroup_alive:
_orphan_stdio_pids.add(pid)
_orphan_stdio_pid_servers[pid] = self.name
else:
# Nothing to reap — drop the pgid so PID reuse can't
# surface stale pgroup state later.
_stdio_pgids.pop(pid, None)
async def _preflight_content_type(
self,
url: str,
*,
headers: Optional[dict] = None,
ssl_verify: bool = True,
client_cert=None,
timeout: float = 5.0,
) -> None:
"""Probe *url* for an MCP-shaped response before the SDK connects.
A URL pointing at a plain web page makes the SDK sit out the full
``connect_timeout`` before an opaque ``CancelledError``; this raises
:class:`NonMcpEndpointError` within ``timeout`` instead. Allow-list based:
only a 2xx with a definite non-MCP content type is rejected, and only
after a JSON-RPC ``initialize`` POST also fails to look like MCP (some
servers serve a UI on GET but speak Streamable HTTP via POST). Missing
content type, non-2xx, or transport errors pass silently — the real
handshake stays the source of truth. Uses its own httpx client OUTSIDE
the SDK's anyio task group so the error isn't wrapped in an ExceptionGroup.
"""
try:
import httpx as _httpx
except ImportError:
return # No httpx → skip probe; SDK import would have failed first.
client_kwargs: dict = {
"verify": ssl_verify,
"follow_redirects": True,
"timeout": _httpx.Timeout(timeout),
}
if client_cert is not None:
client_kwargs["cert"] = client_cert
probe_headers = dict(headers) if headers else {}
try:
async with _httpx.AsyncClient(**client_kwargs) as client:
# HEAD is cheapest; fall back to GET on 405/501.
resp = await client.head(url, headers=probe_headers)
if resp.status_code in (405, 501):
resp = await client.get(url, headers=probe_headers)
# Non-MCP content type on HEAD/GET: try a JSON-RPC POST before
# rejecting, so POST-only servers aren't false positives.
ct = (
resp.headers.get("content-type", "")
.split(";")[0]
.strip()
.lower()
)
if (
ct
and ct not in self._MCP_CONTENT_TYPES
and 200 <= resp.status_code < 300
):
post_resp = await client.post(
url,
headers={
**probe_headers,
"Content-Type": "application/json",
"Accept": "application/json, text/event-stream",
},
content=(
'{"jsonrpc":"2.0","id":"_probe",'
'"method":"initialize",'
'"params":{"protocolVersion":"2025-03-26",'
'"capabilities":{},'
'"clientInfo":{"name":"hermes-probe",'
'"version":"0.1"}}}'
),
)
if 200 <= post_resp.status_code < 300:
post_ct = (
post_resp.headers.get("content-type", "")
.split(";")[0]
.strip()
.lower()
)
if post_ct in self._MCP_CONTENT_TYPES:
resp = post_resp
except _httpx.HTTPError:
return # DNS/connect/timeout/transport error — let the SDK try.
# Only judge 2xx: a 4xx/5xx may be an auth challenge or transient error
# the real handshake handles correctly.
if not (200 <= resp.status_code < 300):
return
ct_base = resp.headers.get("content-type", "").split(";")[0].strip().lower()
if not ct_base:
return # No content type advertised — don't second-guess the SDK.
if ct_base in self._MCP_CONTENT_TYPES:
return # Looks like a real MCP endpoint.
raise NonMcpEndpointError(
f"MCP server '{self.name}' at {url} returned Content-Type "
f"'{ct_base}', not an MCP response (expected one of: "
f"{', '.join(self._MCP_CONTENT_TYPES)}). The URL most likely "
"points at a web page rather than an MCP endpoint — check it "
"resolves to a Streamable HTTP / SSE endpoint "
"(e.g. https://host/mcp, not https://host/)."
)
def _reconnect_or_reraise_group(self, eg: BaseExceptionGroup) -> str:
"""Map an SDK transport TaskGroup failure to a clean ``"reconnect"``.
HTTP/SSE stream pumps run in an anyio TaskGroup, so a transient stream
drop escapes as a ``BaseExceptionGroup``. Unmapped, ``run()`` would back
off and eventually park the server for 300s (deregistering its tools)
over a sub-second glitch; ``"reconnect"`` rebuilds the session at once.
Re-raise when it is not a transient drop: shutdown in progress
(``_shutdown_event`` is set before the task is cancelled), the group
carries KeyboardInterrupt/SystemExit or a real CancelledError (must
propagate), or no live session was reached this attempt (``_ready``
unset — connect failures must go through backoff, not hot-loop).
"""
if self._shutdown_event.is_set():
raise eg
fatal, _rest = eg.split((KeyboardInterrupt, SystemExit))
if fatal is not None:
raise eg
cancelled, _rest = eg.split(asyncio.CancelledError)
if cancelled is not None:
raise eg
if not self._ready.is_set():
raise eg
logger.debug(
"MCP server '%s': transport TaskGroup exited after a live session "
"(%r) — reconnecting immediately instead of backing off",
self.name, eg,
)
return "reconnect"
async def _run_http(self, config: dict):
"""Run the server using HTTP/StreamableHTTP transport."""
_core._ensure_mcp_sdk()
if not _core._MCP_HTTP_AVAILABLE:
raise ImportError(
f"MCP server '{self.name}' requires HTTP transport but "
"mcp.client.streamable_http is not available. "
"Upgrade the mcp package to get HTTP support."
)
url = config["url"]
headers = dict(config.get("headers") or {})
# Portable Agent Plugins v1 (strict_redirect_headers): configured
# headers MUST NOT follow a redirect to a different origin. Capture the
# configured names BEFORE client-generated headers are merged in.
_strict_cfg_headers = bool(config.get("strict_redirect_headers"))
_configured_header_names = {key.lower() for key in headers}
# Optional per-user identity header; explicit headers of the same name win.
headers = _apply_identity_header(self.name, config, headers)
# Some servers require MCP-Protocol-Version on the initial request; seed
# it (case-insensitive user override wins). Seeded from the HANDSHAKE
# version, not the latest: the body sent by ``initialize()`` speaks the
# handshake era, and a 2026-07-28 header would route the request onto
# the server's per-request-envelope ladder, which rejects that body.
# The header must agree with what the body actually speaks.
if not any(key.lower() == "mcp-protocol-version" for key in headers):
headers["mcp-protocol-version"] = _core.LATEST_HANDSHAKE_VERSION
connect_timeout = config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT)
ssl_verify = config.get("ssl_verify", True)
client_cert = _resolve_client_cert(self.name, config)
# OAuth 2.1 PKCE via the central MCPOAuthManager so one provider is
# reused across reconnects and shared with config-time CLI paths. On
# setup failure (e.g. non-interactive without cached tokens) re-raise so
# only this server is reported failed.
_oauth_auth = None
if self._auth_type == "oauth":
try:
from tools.mcp_oauth_manager import get_manager
_oauth_auth = get_manager().get_or_build_provider(
self.name, url, config.get("oauth"),
)
except Exception as exc:
logger.warning("MCP OAuth setup failed for '%s': %s", self.name, exc)
raise
sampling_kwargs = self._sampling.session_kwargs() if self._sampling else {}
if self._elicitation:
sampling_kwargs.update(self._elicitation.session_kwargs())
if _core._MCP_NOTIFICATION_TYPES and _core._MCP_MESSAGE_HANDLER_SUPPORTED:
sampling_kwargs["message_handler"] = self._make_message_handler()
if _core._MCP_LOGGING_CALLBACK_SUPPORTED:
sampling_kwargs["logging_callback"] = self._make_logging_callback()
# SSE transport (``transport: sse`` in the mcp_servers entry).
if config.get("transport") == "sse":
if _strict_cfg_headers:
# Fail closed: SSE cannot enforce the redirect boundary.
raise ValueError(
f"MCP server '{self.name}': strict_redirect_headers is "
"not supported on the SSE transport."
)
if _core.sse_client is None:
raise ImportError(
f"MCP server '{self.name}' requires SSE transport but "
"mcp.client.sse.sse_client is not available. "
"Upgrade the mcp package to get SSE support."
)
# sse_read_timeout bounds the gap between SSE events. SSE servers
# commonly idle for minutes, so tool_timeout (60s) would drop the
# stream; 300s matches the Streamable HTTP read timeout below.
_sse_kwargs: dict = {
"url": url,
"headers": headers or None,
"timeout": float(connect_timeout),
"sse_read_timeout": 300.0,
}
if _oauth_auth is not None:
# Forward OAuth to sse_client, else OAuth SSE servers 401 silently.
_sse_kwargs["auth"] = _oauth_auth
if client_cert is not None or ssl_verify is not True:
# sse_client has no verify/cert kwargs: wrap the SDK defaults
# (follow_redirects=True) in an httpx_client_factory, forwarding
# the SDK's (headers, auth, timeout) and layering TLS on top. The
# client MUST come from the SDK's own httpx module (httpx2 on
# mcp >= 2.0) — see sdk_httpx().
_httpx_mod = _core.sdk_httpx()
_cert_for_factory = client_cert
_verify_for_factory = ssl_verify
def _mcp_http_client_factory(
headers=None, timeout=None, auth=None,
):
kwargs: dict = {
"follow_redirects": True,
"verify": _verify_for_factory,
}
if timeout is not None:
kwargs["timeout"] = timeout
else:
kwargs["timeout"] = _httpx_mod.Timeout(30.0, read=300.0)
if headers is not None:
kwargs["headers"] = headers
if auth is not None:
kwargs["auth"] = auth
if _cert_for_factory is not None:
kwargs["cert"] = _cert_for_factory
return _httpx_mod.AsyncClient(**kwargs)
_sse_kwargs["httpx_client_factory"] = _mcp_http_client_factory
try:
async with _core.sse_client(**_sse_kwargs) as (read_stream, write_stream):
async with _core.ClientSession(
read_stream, write_stream, **sampling_kwargs
) as session:
reason = await self._serve_session(
session, float(connect_timeout), "SSE"
)
except BaseExceptionGroup as _eg:
# Transport TaskGroup dropped: reconnect instead of backoff/park.
reason = self._reconnect_or_reraise_group(_eg)
return reason
if _core._MCP_NEW_HTTP:
# mcp >= 1.24.0: build an explicit AsyncClient matching the SDK's
# create_mcp_http_client defaults. It MUST come from the SDK's httpx
# module (httpx2 on mcp >= 2.0) since the SDK sends its own Request
# objects through it — see sdk_httpx().
httpx = _core.sdk_httpx()
_original_url = httpx.URL(url)
_strip_auth_on_cross_origin_redirect = _make_redirect_header_stripper(
_original_url,
strict=_strict_cfg_headers,
configured_header_names=_configured_header_names,
)
client_kwargs: dict = {
"follow_redirects": True,
"timeout": httpx.Timeout(float(connect_timeout), read=300.0),
"verify": ssl_verify,
"event_hooks": {"response": [_strip_auth_on_cross_origin_redirect]},
}
if headers:
client_kwargs["headers"] = headers
if _oauth_auth is not None:
client_kwargs["auth"] = _oauth_auth
if client_cert is not None:
client_kwargs["cert"] = client_cert
# Caller owns the client lifecycle — the SDK skips cleanup when
# http_client is provided.
try:
async with httpx.AsyncClient(**client_kwargs) as http_client:
# Unpacked positionally: mcp 1.x yields (read, write,
# get_session_id), 2.x yields (read, write).
async with _core.streamable_http_client(url, http_client=http_client) as _streams:
read_stream, write_stream = _streams[0], _streams[1]
async with _core.ClientSession(read_stream, write_stream, **sampling_kwargs) as session:
reason = await self._serve_session(
session, float(connect_timeout), "HTTP"
)
except BaseExceptionGroup as _eg:
# Transport TaskGroup dropped: reconnect instead of backoff/park.
reason = self._reconnect_or_reraise_group(_eg)
return reason
else:
# Deprecated API (mcp < 1.24.0): the SDK owns the httpx client.
if _strict_cfg_headers:
# Fail closed: without an owned client we cannot hook redirects,
# so the cross-origin header boundary cannot be enforced.
raise ImportError(
f"MCP server '{self.name}' requires mcp >= 1.24.0 to "
"enforce the portable redirect-header boundary "
"(strict_redirect_headers). Upgrade the mcp package."
)
_http_kwargs: dict = {
"headers": headers,
"timeout": float(connect_timeout),
"verify": ssl_verify,
}
if _oauth_auth is not None:
_http_kwargs["auth"] = _oauth_auth
try:
async with _core.streamablehttp_client(url, **_http_kwargs) as (
read_stream, write_stream, _get_session_id,
):
async with _core.ClientSession(read_stream, write_stream, **sampling_kwargs) as session:
reason = await self._serve_session(
session, float(connect_timeout), "legacy HTTP"
)
except BaseExceptionGroup as _eg:
# Transport TaskGroup dropped: reconnect instead of backoff/park.
reason = self._reconnect_or_reraise_group(_eg)
return reason
async def _discover_tools(self):
"""Discover tools from the connected session.
Capability-gated: prompt-/resource-only servers raise ``MCPError(-32601)``
on ``tools/list``, which would abort the connection — skip the call when
``tools`` isn't advertised.
"""
# Fresh transport: re-probe with cheap ``ping`` in case the server gained
# support across the reconnect.
self._ping_unsupported = False
if self.session is None:
return
if not self._advertises_tools():
logger.info(
"MCP server '%s': does not advertise 'tools' capability — "
"skipping tools/list (prompts/resources remain available)",
self.name,
)
self._tools = []
self._register_discovered_tools_if_needed()
return
async with self._rpc_lock:
self._list_cache_meta = {}
self._tools = await _core._paginate_full_list(
self.session.list_tools, "tools", self.name,
cache_meta_out=self._list_cache_meta,
)
self._register_discovered_tools_if_needed()
def _register_discovered_tools_if_needed(self) -> None:
"""Publish freshly discovered tools for a registry-owned server if none are registered.
Initial registration normally happens in ``_discover_and_register_server``
after ``start()``. On reconnect, outage handling may clear ``_ready`` and
deregister stale tools; ownership via ``_servers`` authorizes publishing
before readiness is restored so a revival never comes back with zero
tools. A server retained after a recoverable initial failure is likewise
owned before its first session, which authorizes its first publication.
"""
if self._registered_tool_names:
return
if not self._ready.is_set():
with _core._lock:
if _core._servers.get(self.name) is not self:
return
self._registered_tool_names = _core._register_server_tools(
self.name, self, self._config
)
# A retained initial-failure server that just published tools has
# recovered: drop its stale connect error from status surfaces.
with _core._lock:
if _core._servers.get(self.name) is self:
_core._server_connect_errors.pop(self.name, None)