refactor(tools/mcp): split mcp_tool.py into transport/lifecycle/schema/handlers/... sibling modules; compact watchdog and schema cache
This commit is contained in:
@@ -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"}) == []
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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
File diff suppressed because it is too large
Load Diff
@@ -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
|
||||
@@ -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
|
||||
@@ -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 {}
|
||||
@@ -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]"
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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")
|
||||
@@ -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",
|
||||
}
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user