refactor(tools): MCP docstring/comment compaction, status predicate, renderer tidy

This commit is contained in:
Teknium
2026-09-02 23:36:18 -07:00
parent 7f7fd4a533
commit 49c38ef10f
5 changed files with 229 additions and 360 deletions
+78 -115
View File
@@ -77,9 +77,8 @@ from tools.mcp_tool_discovery import ( # noqa: F401
has_registered_mcp_tools, get_registered_mcp_server_names)
# Wall-clock bound on the (fail-open) OSV malware preflight run off the loop before a
# stdio spawn. Kept just ABOVE osv_check._TIMEOUT (10s) so the inner socket timeout
# normally fires first; this only bites when a stalled SSL handshake defeats it.
# Wall-clock bound on the fail-open OSV malware preflight before a stdio spawn; just ABOVE
# osv_check._TIMEOUT (10s) so it only bites when a stalled SSL handshake defeats that.
_OSV_MALWARE_CHECK_TIMEOUT_S = 12.0
@@ -91,17 +90,15 @@ _MCP_AVAILABLE = _MCP_HTTP_AVAILABLE = _MCP_NEW_HTTP = _MCP_LEGACY_HTTP = False
_MCP_SAMPLING_TYPES = _MCP_NOTIFICATION_TYPES = _MCP_ELICITATION_TYPES = False
_MCP_MESSAGE_HANDLER_SUPPORTED = _MCP_LOGGING_CALLBACK_SUPPORTED = False
sse_client = None
# Fallback for SDKs that don't export LATEST_PROTOCOL_VERSION (Streamable HTTP arrived
# with 2025-03-26, so this stays valid for the HTTP path).
# Fallback for SDKs without LATEST_PROTOCOL_VERSION (Streamable HTTP arrived with 2025-03-26).
LATEST_PROTOCOL_VERSION = "2025-03-26"
# Newest revision ``ClientSession.initialize()`` actually speaks. From 2026-07-28 the
# handshake is replaced by a per-request envelope, so this can be OLDER than
# LATEST_PROTOCOL_VERSION; the MCP-Protocol-Version header must be seeded from this one.
# Newest revision ``ClientSession.initialize()`` speaks; from 2026-07-28 the handshake is a
# per-request envelope so this can be OLDER than LATEST_PROTOCOL_VERSION, and the
# MCP-Protocol-Version header must be seeded from THIS one.
LATEST_HANDSHAKE_VERSION = LATEST_PROTOCOL_VERSION
# Importing ``mcp`` costs ~260ms, so it is deferred to first real use (_ensure_mcp_sdk).
# Availability is decided here with a metadata-only find_spec probe so every
# ``if not _MCP_AVAILABLE`` gate / test patch / skipif keeps its exact semantics.
# Importing ``mcp`` costs ~260ms, so it is deferred to first use (_ensure_mcp_sdk); availability
# is decided now via find_spec so every ``if not _MCP_AVAILABLE`` gate / patch / skipif holds.
try:
_MCP_AVAILABLE = importlib.util.find_spec("mcp") is not None
except Exception:
@@ -113,9 +110,9 @@ ClientSession: Any = None
_MCP_SDK_IMPORT_ATTEMPTED = False
_MCP_SDK_IMPORT_LOCK = threading.Lock()
# Optional SDK type families (module, names, debug message when absent) bound in this order to
# _MCP_SAMPLING_TYPES / _MCP_ELICITATION_TYPES / _MCP_NOTIFICATION_TYPES. Each is gated
# separately so an older SDK only loses that feature, not MCP.
# Optional SDK type families (module, names, debug message when absent), bound in this order to
# _MCP_SAMPLING_TYPES / _MCP_ELICITATION_TYPES / _MCP_NOTIFICATION_TYPES; an older SDK only
# loses that feature, not MCP.
_OPTIONAL_TYPE_FAMILIES = (
("mcp.types", ("CreateMessageResult", "CreateMessageResultWithTools", "ErrorData", "SamplingCapability",
"SamplingToolsCapability", "TextContent", "ToolUseContent"),
@@ -126,9 +123,8 @@ _OPTIONAL_TYPE_FAMILIES = (
"ResourceListChangedNotification"),
"MCP notification types not available -- dynamic tool discovery disabled"),
)
# SDK symbols bound by _ensure_mcp_sdk(). Module __getattr__ (PEP 562) imports the SDK on
# first external access, so mock.patch("tools.mcp_tool.stdio_client") sees a real original
# and the mock is never clobbered (_ensure is idempotent).
# Bound by _ensure_mcp_sdk(); module __getattr__ (PEP 562) imports the SDK on first external
# access so mock.patch("tools.mcp_tool.stdio_client") sees a real original, never clobbered.
_MCP_SDK_LAZY_SYMBOLS = frozenset(
{"StdioServerParameters", "stdio_client", "streamablehttp_client", "streamable_http_client"}
| {n for _mod, names, _msg in _OPTIONAL_TYPE_FAMILIES for n in names})
@@ -159,12 +155,9 @@ def _import_sdk_names(module: str, names: tuple, missing_msg: Optional[str] = No
def _ensure_mcp_sdk() -> bool:
"""Import the optional ``mcp`` SDK on first use; return availability.
Idempotent and thread-safe. Honors a test-patched ``_MCP_AVAILABLE=False`` (no import)
and pre-installed mock symbols (``ClientSession`` already set means no re-import, so
mocks are never clobbered).
"""
"""Import the optional ``mcp`` SDK on first use; return availability. Idempotent and
thread-safe; honors a test-patched ``_MCP_AVAILABLE=False`` (no import) and pre-installed
mocks (``ClientSession`` already set means no re-import)."""
global _MCP_SDK_IMPORT_ATTEMPTED, _MCP_AVAILABLE, _MCP_HTTP_AVAILABLE, _MCP_NEW_HTTP, _MCP_LEGACY_HTTP
global _MCP_SAMPLING_TYPES, _MCP_NOTIFICATION_TYPES, _MCP_ELICITATION_TYPES, sse_client
global _MCP_MESSAGE_HANDLER_SUPPORTED, _MCP_LOGGING_CALLBACK_SUPPORTED, LATEST_HANDSHAKE_VERSION
@@ -213,14 +206,10 @@ _SDK_HTTPX_MOD = None
def sdk_httpx():
"""Return the httpx module the *installed* MCP SDK is built against.
mcp 2.0 moved to ``httpx2`` (same API, separate distribution). Every object crossing the
SDK boundary (the ``AsyncClient`` passed to the transport, OAuth ``Request`` objects, the
exception classes) must come from the module the SDK itself imports, or it fails at the
transport layer. Resolved from the SDK's transport module, not a version number; falls
back to the newest module present. ``None`` only when neither module is importable.
"""
"""The httpx module the *installed* MCP SDK is built against (mcp 2.0 moved to ``httpx2``).
Every object crossing the SDK boundary (AsyncClient, OAuth Request, exception classes) must
come from the module the SDK itself imports or it fails at the transport layer. Resolved
from the SDK's transport module, else the newest present; ``None`` if neither imports."""
global _SDK_HTTPX_MOD
if _SDK_HTTPX_MOD is not None:
return _SDK_HTTPX_MOD
@@ -264,38 +253,29 @@ _DEFAULT_CONNECT_TIMEOUT = 60 # seconds for initial connection per server
_MAX_RECONNECT_RETRIES = 5
_MAX_INITIAL_CONNECT_RETRIES = 3 # retries for the very first connection attempt
_MAX_BACKOFF_SECONDS = 60
# Parked servers (budget exhausted, tools deregistered) self-probe on this cadence: with
# no tools registered nothing else can ever revive them.
# Parked servers (tools deregistered) self-probe on this cadence: nothing else can revive them.
_PARKED_RETRY_INTERVAL = 300
_RECYCLED_RECONNECT_TIMEOUT = 15.0
# Bounded wait for a respawned stdio child when a call finds it dead (gateway restarts kill
# every MCP child); bounded so a broken server still parks via run()'s rapid-drop budget.
_STDIO_RESPAWN_WAIT_SEC = 15.0
# Servers may expire idle sessions on any TTL, so the client MUST ping faster than that TTL;
# short-TTL servers (~15s) need a smaller configured ``keepalive_interval``. The floor stops
# a tiny interval from busy-looping.
# The client MUST ping faster than the server's idle-session TTL (short-TTL servers need a
# smaller configured ``keepalive_interval``); the floor stops a tiny interval busy-looping.
_DEFAULT_KEEPALIVE_INTERVAL = 180
_MIN_KEEPALIVE_INTERVAL = 5
# One bounded cancellation cycle for pending loop tasks at final shutdown, so
# cancellation-resistant tasks cannot hang process exit.
# One bounded cancellation cycle at final shutdown so resistant tasks cannot hang exit.
_MCP_LOOP_DRAIN_TIMEOUT = 3.0
# JSON-RPC 2.0 "method not found" (e.g. a server without the optional ``ping``).
# _ensure_mcp_sdk() overrides it from mcp.types once the SDK is loaded.
# JSON-RPC 2.0 "method not found" (server without optional ``ping``); _ensure_mcp_sdk()
# overrides it from mcp.types once loaded.
_JSONRPC_METHOD_NOT_FOUND = -32601
# Cap on nextCursor pagination so a server returning a cursor forever cannot spin
# discovery; 50 pages at 50-100 items/page covers thousands of entries.
# nextCursor pagination cap so a forever-cursor cannot spin discovery (50 pages = thousands).
_MCP_LIST_MAX_PAGES = 50
async def _paginate_full_list(list_method, items_attr: str, server_name: str,
cache_meta_out: Optional[dict] = None):
"""Drain a paginated ``list_*`` call by following ``nextCursor`` (the SDK fetches one
page per call). ``cache_meta_out`` receives the first page's SEP-2549 hints (``ttl_ms``,
``cache_scope``). Callers must hold the server's ``_rpc_lock`` for a consistent snapshot."""
"""Drain a paginated ``list_*`` call by following ``nextCursor``; ``cache_meta_out`` gets the
first page's SEP-2549 hints. Callers must hold the server's ``_rpc_lock``."""
items: list = []
cursor = None
for _ in range(_MCP_LIST_MAX_PAGES):
@@ -331,10 +311,9 @@ async def _paginate_full_list(list_method, items_attr: str, server_name: str,
# ---------------------------------------------------------------------------
class MCPServerTask(MCPServerRunMixin, MCPServerTransportMixin, MCPServerHealthMixin):
"""One MCP server connection living in one long-lived asyncio Task, so the transport's
anyio cancel scopes are entered and exited in the same Task. Run state machine in
``MCPServerRunMixin``, transport bring-up in ``MCPServerTransportMixin``, keepalive /
refresh / liveness in ``MCPServerHealthMixin``."""
"""One MCP server connection in one long-lived asyncio Task (the transport's anyio cancel
scopes must enter/exit in the same Task). Run state machine, transport bring-up and
keepalive/liveness live in the three mixins."""
__slots__ = (
"name", "session", "tool_timeout", "_task", "_ready", "_shutdown_event", "_reconnect_event",
@@ -353,8 +332,7 @@ class MCPServerTask(MCPServerRunMixin, MCPServerTransportMixin, MCPServerHealthM
self._task: Optional[asyncio.Task] = None
self._ready = asyncio.Event()
self._shutdown_event = asyncio.Event()
# When set, _run_http/_run_stdio exit their async-with cleanly and run() re-enters
# the transport (auth recovery, manual refresh, ...).
# Set -> _run_http/_run_stdio exit cleanly and run() re-enters the transport.
self._reconnect_event = asyncio.Event()
self._tools: list = []
self._error: Optional[Exception] = None
@@ -363,53 +341,47 @@ class MCPServerTask(MCPServerRunMixin, MCPServerTransportMixin, MCPServerHealthM
self._elicitation: Optional[ElicitationHandler] = None
self._registered_tool_names: list[str] = []
self._reconnect_retries: int = 0
# Rapid-drop budget: a (re)established session is UNPROVEN until it survives a full
# keepalive interval or serves a successful call. Only a proven session clears the
# reconnect budget, so a transport that flaps right after the handshake still parks.
# Rapid-drop budget: a session is UNPROVEN until it survives a keepalive interval or a
# successful call; only a proven session clears the budget, so a post-handshake flapper
# still parks.
self._session_proven: bool = False
# Set once tools were ever registered, never cleared (unlike _ready, which clears every
# reconnect cycle): separates first-connect from reconnect failures in run()'s ladders.
# Never cleared (unlike _ready): separates first-connect from reconnect failures.
self._ever_connected: bool = False
# True from park until the session proves healthy again; logs the revival once.
# True from park until proven healthy again; logs the revival once.
self._was_parked: bool = False
# In-flight RPC tasks, so a reconnect/shutdown teardown can fail them fast instead of
# orphaning them on a dying transport; _reconnecting is True while such a deliberate
# teardown runs, letting _track_inflight_rpc turn the cancel into a retryable error.
# In-flight RPC tasks so a deliberate teardown fails them fast; _reconnecting is True
# during that teardown so _track_inflight_rpc turns the cancel into a retryable error.
self._inflight_tasks: set = set()
self._reconnecting: bool = False
# Latched by races (teardown-vs-keepalive, auth-lock corruption); verified lazily by
# ensure_healthy() before the next call.
# Latched by races (teardown-vs-keepalive, auth-lock corruption); ensure_healthy()
# verifies before the next call.
self._suspect_reason: Optional[str] = None
# A teardown that failed >=1 in-flight call makes the next reconnect a RACE RECOVERY:
# it must not charge the rapid-drop budget.
# Teardown that failed in-flight calls => next reconnect is RACE RECOVERY, not a
# budget charge.
self._teardown_race: bool = False
# One-time grace: an auth/permanent-classified failure on a previously PROVEN session
# gets one suspect+reconnect cycle before the park ladder applies.
# One-time grace: auth/permanent failure on a PROVEN session gets one suspect+reconnect
# cycle before parking.
self._permanent_grace_used: bool = False
# Children of the current stdio transport: in-flight calls fail FAST when the child
# dies instead of riding out the tool timeout.
# Children of the current stdio transport: in-flight calls fail FAST when one dies.
self._stdio_child_pids: Set[int] = set()
self._auth_type: str = ""
self._refresh_lock = asyncio.Lock()
# A stdio session is one JSON-RPC stream: a list_tools issued by the notification
# handler while a tool call is in flight can wedge it. Serialize client-initiated
# RPCs per server (HTTP too, for ordering).
# A stdio session is one JSON-RPC stream (a concurrent list_tools can wedge a tool
# call): serialize client-initiated RPCs per server (HTTP too, for ordering).
self._rpc_lock = asyncio.Lock()
self._pending_refresh_tasks: set[asyncio.Task] = set()
# contextvars snapshot of the agent task inside session.call_tool(). The SDK dispatches
# elicitation/create on a separate task that does not inherit HERMES_SESSION_PLATFORM;
# replaying this context in the elicitation callback routes the prompt correctly.
# contextvars snapshot inside session.call_tool(): the SDK runs elicitation/create on a
# task that does not inherit HERMES_SESSION_PLATFORM, so the callback replays this.
self._pending_call_context: Optional[contextvars.Context] = None
self._lifecycle_started_at = self._last_tool_call_at = time.monotonic()
self._idle_timeout_seconds: Optional[float] = None
self._max_lifetime_seconds: Optional[float] = None
self._recycled_reason: Optional[str] = None
# InitializeResult from the handshake: the server's REAL advertised capabilities.
# Handshake InitializeResult: the server's REAL advertised capabilities.
self.initialize_result: Optional[Any] = None
# SEP-2549 cache hints from the last tools/list (ttl_ms, cache_scope).
self._list_cache_meta: dict = {}
# Latched when keepalive ``ping`` returns -32601 (optional utility not implemented);
# later keepalives use list_tools instead. Reset on every fresh transport connection.
# Latched when ``ping`` returns -32601; keepalives then use list_tools. Reset per connect.
self._ping_unsupported: bool = False
# Content types a real Streamable-HTTP endpoint may return on the initial POST/GET;
@@ -422,48 +394,43 @@ class MCPServerTask(MCPServerRunMixin, MCPServerTransportMixin, MCPServerHealthM
# ---------------------------------------------------------------------------
_servers: Dict[str, MCPServerTask] = {}
# Profile registry scope owning each live connection (None outside multiplex): a
# multiplexed /reload-mcp tears down only its own profile's servers.
# Profile registry scope per live connection (None outside multiplex) so a multiplexed
# /reload-mcp tears down only its own profile's servers.
_server_scope_keys: Dict[str, Optional[str]] = {}
_server_connecting: set[str] = set()
_server_connect_errors: Dict[str, str] = {}
# Lazy startup: servers registered from the on-disk schema cache without connecting;
# popped once a real connection is established on first use.
# Lazy startup: servers registered from the schema cache without connecting; popped on
# first real connection.
_lazy_server_configs: Dict[str, dict] = {}
_lazy_server_fingerprints: Dict[str, str] = {}
_lazy_server_tool_names: Dict[str, List[str]] = {}
# Discovery installs a task-local claim around ``_connect_server`` so it can retain a
# recoverable parked task without standalone probe calls publishing failed servers into
# module-global ownership.
# Task-local claim around ``_connect_server``: discovery retains a recoverable parked task
# while standalone probes never publish failed servers into module-global ownership.
_connect_server_claim: contextvars.ContextVar[Optional[Callable[[MCPServerTask], None]]] = (
contextvars.ContextVar("mcp_connect_server_claim", default=None))
# Per-server connect cooldown. A server that fails to spawn never reaches ``_servers``, so
# without this every ``discover_mcp_tools()`` (one per worker session) would respawn it from
# scratch — a restart storm whose unreaped subprocesses destabilise the healthy co-located
# servers. Failed attempts stamp an exponential-backoff deadline that
# ``register_mcp_servers`` honours; a successful connection clears it.
# Per-server connect cooldown: a server that fails to spawn never reaches ``_servers``, so
# without it every ``discover_mcp_tools()`` (one per worker session) would respawn it — a
# restart storm whose unreaped children destabilise healthy servers. Exponential-backoff
# deadline honoured by ``register_mcp_servers``; cleared on success.
_server_connect_retry_after: Dict[str, float] = {} # name -> monotonic deadline
_server_connect_failures: Dict[str, int] = {} # name -> consecutive failures
_CONNECT_RETRY_BASE_BACKOFF_SEC = 30.0
_CONNECT_RETRY_MAX_BACKOFF_SEC = 600.0
# Circuit breaker per server: closed (count < threshold) -> open (calls short-circuit with a
# "stop retrying" message until the cooldown elapses) -> half-open (next call is a probe;
# success closes, failure re-arms). Mutate only via _bump_server_error / _reset_server_error.
# Per-server circuit breaker: closed -> open (calls short-circuit until the cooldown) ->
# half-open (next call probes). Mutate only via _bump_server_error / _reset_server_error.
_server_error_counts: Dict[str, int] = {}
_server_breaker_opened_at: Dict[str, float] = {}
_CIRCUIT_BREAKER_THRESHOLD = 3
_CIRCUIT_BREAKER_COOLDOWN_SEC = 60.0
# Trust-tier gating (``mcp_servers.<name>.trust: full | untrusted``). On an untrusted server
# every write-capable call needs user approval before the RPC fires; a tool is write-capable
# unless its discovery-time ``annotations.readOnlyHint`` is exactly True (malformed fails
# closed). readOnlyHint is a server-supplied HINT and a hostile server can lie, but on an
# untrusted server a lie can only skip approval for calls the operator was already warned
# about — never widen access. Missing ``trust`` defaults to full (backward compatible); any
# unrecognized value normalizes to untrusted (a typo must never disable the gate).
# Classified at CALL time from DISCOVERY data: no schema mutation, prompt cache intact.
# Trust-tier gating (``trust: full | untrusted``): on an untrusted server every write-capable
# call (discovery-time ``readOnlyHint`` not exactly True; malformed fails closed) needs approval
# before the RPC fires. A lying readOnlyHint can only skip approval for calls the operator was
# already warned about, never widen access. Missing trust = full; unrecognized = untrusted (a
# typo must never disable the gate). Classified at CALL time from DISCOVERY data: no schema
# mutation, prompt cache intact.
_server_trust_levels: Dict[str, str] = {}
_tool_read_only_hints: Dict[str, Dict[str, bool]] = {}
@@ -485,11 +452,10 @@ def _reset_server_error(server_name: str) -> None:
_server_breaker_opened_at.pop(server_name, None)
# Raw server names opted into parallel tool calls. Raw identity matters: ``foo-bar`` and
# ``foo_bar`` both sanitize to ``foo_bar`` but must not share policy.
# Raw server names opted into parallel tool calls (``foo-bar``/``foo_bar`` sanitize alike but
# must not share policy).
_parallel_safe_servers: set = set()
# registry tool name -> raw server name, captured at registration. The generated name is
# lossy (punctuation -> ``_``), so never re-parse it.
# registry tool name -> raw server name (the generated name is lossy; never re-parse it).
_mcp_tool_server_names: Dict[str, str] = {}
# Dedicated event loop running in a background daemon thread.
@@ -500,8 +466,7 @@ _lock = threading.Lock()
def _mcp_registry_scope() -> Optional[str]:
"""Registry scope for MCP registrations: under a profile multiplexer each profile's MCP
tools live in its own registry overlay; single-profile processes stay global (None)."""
"""Registry scope for MCP registrations: a profile overlay under a multiplexer, else None."""
from agent.secret_scope import is_multiplex_active
if not is_multiplex_active():
return None
@@ -510,16 +475,14 @@ def _mcp_registry_scope() -> Optional[str]:
def _server_registry_scope(name: str) -> Optional[str]:
"""Scope owning server *name*'s tools: recorded at connect, else current. Teardown runs
on the MCP loop without the discovering profile's context, so the scope captured at
adoption into ``_servers`` is authoritative."""
"""Scope owning *name*'s tools: the one captured at adoption (teardown runs on the MCP
loop without the discovering profile's context), else the current one."""
if name in _server_scope_keys:
return _server_scope_keys[name]
return _mcp_registry_scope()
# Cross-process MCP discovery guard: advisory file lock so gateway + CLI + TUI don't all
# run discovery at once.
# Cross-process discovery guard: advisory file lock so gateway + CLI + TUI don't all discover.
_LOCK_UNAVAILABLE: Any = object() # sentinel: locking broken/unavailable
_MCP_DISCOVERY_LOCK_PATH: Optional[str] = None # resolved lazily
# Bounded wait when another process holds the lock.
+46 -75
View File
@@ -39,25 +39,21 @@ def _enabled(cfg: dict) -> bool:
async def _connect_server(name: str, config: dict) -> _core.MCPServerTask:
"""Create an MCPServerTask, start it and return once ready. Tear it down with
``server.shutdown()`` on the same loop. Raises on bad config, missing HTTP support,
or connect/initialize failure."""
"""Create an MCPServerTask, start it, return once ready (tear down with ``server.shutdown()``
on the same loop). Raises on bad config, missing HTTP support or connect failure."""
server = _core.MCPServerTask(name)
claim = _core._connect_server_claim.get()
claim_token = None
if claim is not None:
claim(server)
# The run task copies this context; the claim is for this attempt only, so don't
# retain the discovery closure for the server's life.
# The run task copies this context: don't retain the discovery closure for its life.
claim_token = _core._connect_server_claim.set(None)
try:
await server.start(config)
except asyncio.CancelledError:
# start() already reaps server._task; a shutdown() here could swallow the cancellation.
raise
raise # start() already reaps server._task; shutdown() here could swallow the cancel
except BaseException:
# Discovery owns claimed tasks (recoverable park vs terminal failure); standalone
# probes have no revival owner and must reap locally.
# Discovery owns claimed tasks (recoverable park); standalone probes must reap locally.
if claim is None:
try:
await server.shutdown()
@@ -128,12 +124,9 @@ def _adopt_server(name: str, server: _core.MCPServerTask) -> None:
def _ensure_lazy_server_connected(server_name: str) -> bool:
"""Connect a lazily-registered server on demand (sync; blocks the caller).
Honours the connect cooldown and the ``_server_connecting`` dedup set and routes through
``_discover_and_register_server`` so park/recycle/cooldown bookkeeping stays in one place.
True when a live session exists after.
"""
"""Connect a lazily-registered server on demand (sync; blocks). Honours the cooldown and the
``_server_connecting`` dedup set; routes through ``_discover_and_register_server`` so
park/recycle/cooldown bookkeeping stays in one place. True when a live session exists."""
with _core._lock:
server = _core._servers.get(server_name)
if server is not None and server.session is not None:
@@ -192,8 +185,7 @@ def _get_connected_server_for_call(server_name: str) -> Optional[_core.MCPServer
async def _discover_and_register_server(name: str, config: dict) -> List[str]:
"""Connect one server, register its tools; return the registered names."""
# The claim callback runs inside _connect_server while this frame is suspended; a list
# append avoids a nonlocal rebind.
# The claim fires inside _connect_server while this frame is suspended (list, not nonlocal).
claimed: List[_core.MCPServerTask] = []
claim_token = _core._connect_server_claim.set(claimed.append)
try:
@@ -205,8 +197,7 @@ async def _discover_and_register_server(name: str, config: dict) -> List[str]:
task_cancelling = task.cancelling() if task is not None and hasattr(task, "cancelling") else 0
if (server is not None and server._error is not None and task is not None
and not task.done() and not task_cancelling):
# Recoverable park: the run task stays alive to self-probe, so adopt it for
# shutdown/revival.
# Recoverable park: the run task self-probes, so adopt it for shutdown/revival.
_adopt_server(name, server)
elif server is not None:
await server.shutdown()
@@ -225,13 +216,9 @@ async def _discover_and_register_server(name: str, config: dict) -> List[str]:
def _select_new_servers(servers: Dict[str, dict]) -> Dict[str, dict]:
"""Pick connect candidates and refresh per-server bookkeeping (under ``_lock``).
Candidates: enabled, not connected, not connecting (dedups concurrent discovery entry
points), not lazily registered, not in backoff. Known servers without a live session are
parked or mid-reconnect; their tools are deregistered so nothing else can nudge them —
signal a reconnect here.
"""
"""Pick connect candidates (enabled, not connected/connecting/lazy, not in backoff) and
refresh per-server bookkeeping. Known servers without a live session are parked or
mid-reconnect with tools deregistered, so nothing else can nudge them: signal a reconnect."""
with _core._lock:
connecting = set(_core._server_connecting)
new_servers = {
@@ -255,9 +242,9 @@ def _select_new_servers(servers: Dict[str, dict]) -> Dict[str, dict]:
def _register_lazy_from_cache(new_servers: Dict[str, dict]) -> Tuple[Dict[str, dict], int, int]:
"""Register ``lazy: true`` servers with a valid schema-cache entry without connecting; a
missing/stale entry (or a failed registration) falls back to eager. Returns (servers still
needing an eager connect, lazy tool count, lazy server count)."""
"""Register ``lazy: true`` servers from a valid schema-cache entry without connecting
(missing/stale entry or failed registration -> eager). Returns (eager servers, lazy tool
count, lazy server count)."""
eager_servers: Dict[str, dict] = dict(new_servers)
lazy_registered = 0
lazy_server_count = 0
@@ -302,10 +289,9 @@ async def _discover_all(new_servers: Dict[str, dict]) -> None:
def _run_discovery_pass(new_servers: Dict[str, dict]) -> None:
"""Run ``_discover_all`` on the MCP loop with the interrupt flag parked and the
``_server_connecting`` set cleaned up when the pass dies early."""
# Clear a stale interrupt flag (executor threads are reused) so a prior session's
# interrupt cannot cancel this discovery pass.
"""Run ``_discover_all`` on the MCP loop with the interrupt flag parked; clean up
``_server_connecting`` when the pass dies early."""
# Executor threads are reused: a prior session's stale interrupt must not cancel this pass.
from tools.interrupt import is_interrupted as _is_interrupted, set_interrupt as _set_interrupt
_was_interrupted = _is_interrupted()
if _was_interrupted:
@@ -313,7 +299,7 @@ def _run_discovery_pass(new_servers: Dict[str, dict]) -> None:
try:
_core._run_on_mcp_loop(lambda: _discover_all(new_servers), timeout=120)
except (TimeoutError, InterruptedError) as _e:
# Entries stranded in _server_connecting would block future reconnect attempts.
# Stranded _server_connecting entries would block future reconnects.
how = "timed out" if isinstance(_e, TimeoutError) else "interrupted"
with _core._lock:
stale = [n for n in new_servers if n in _core._server_connecting]
@@ -330,8 +316,7 @@ def _run_discovery_pass(new_servers: Dict[str, dict]) -> None:
def _connected_summary(names, *, lazy_tools: int = 0, lazy_servers: int = 0) -> Tuple[int, int, int]:
"""(tool count, connected server count, failed count) for a set of candidate names,
folding in lazily registered servers."""
"""(tool count, connected count, failed count) for candidate names, plus lazy servers."""
with _core._lock:
connected = [n for n in names if n in _core._servers and n not in _core._server_connect_errors]
tool_count = sum(len(getattr(_core._servers[n], "_registered_tool_names", [])) for n in connected)
@@ -350,9 +335,8 @@ def _log_summary(prefix: str, names, **lazy) -> None:
def register_mcp_servers(servers: Dict[str, dict]) -> List[str]:
"""Connect the given ``{name: config}`` servers and register their tools. Idempotent for
connected names; ``enabled: false`` servers are skipped without disconnecting existing
sessions. Returns every registered MCP tool name."""
"""Connect ``{name: config}`` servers and register their tools; idempotent for connected
names, ``enabled: false`` skipped without disconnecting. Returns every MCP tool name."""
if not _core._ensure_mcp_sdk():
logger.debug("MCP SDK not available -- skipping explicit MCP registration")
return []
@@ -376,9 +360,8 @@ def register_mcp_servers(servers: Dict[str, dict]) -> List[str]:
def _acquire_discovery_lock_with_retry():
"""Cross-process guard: a lock loser waits for the holder, then runs its own discovery;
if locking is unavailable or the wait expires, run unguarded (fail-soft). Returns the
cookie (None / _LOCK_UNAVAILABLE when unguarded)."""
"""Cross-process guard: a lock loser waits for the holder then discovers itself; unavailable
locking or an expired wait runs unguarded (fail-soft). None / _LOCK_UNAVAILABLE = unguarded."""
cookie = _core._try_acquire_mcp_discovery_lock()
if cookie is not None:
return cookie
@@ -397,9 +380,8 @@ def _acquire_discovery_lock_with_retry():
def discover_mcp_tools() -> List[str]:
"""Entry point: load config, connect servers, register tools. Safe without the ``mcp``
package (returns []). Idempotent: only servers missing from a previous call are retried.
Returns all MCP tool names."""
"""Entry point: load config, connect servers, register tools. [] without the ``mcp``
package; idempotent (only servers missing from a previous call are retried)."""
servers = _core._load_mcp_config()
if not servers:
logger.debug("No MCP servers configured")
@@ -424,8 +406,8 @@ def discover_mcp_tools() -> List[str]:
def is_mcp_tool_parallel_safe(tool_name: str) -> bool:
"""True when the tool's server opted into ``supports_parallel_tool_calls``. Uses the
provenance captured at registration, never the (ambiguous) ``mcp__{server}__{tool}`` shape."""
"""True when the tool's server opted into ``supports_parallel_tool_calls`` (provenance
captured at registration, never the ambiguous ``mcp__{server}__{tool}`` shape)."""
if not tool_name.startswith(_core.MCP_TOOL_NAME_PREFIX):
return False
with _core._lock:
@@ -434,12 +416,8 @@ def is_mcp_tool_parallel_safe(tool_name: str) -> bool:
def get_mcp_status() -> List[dict]:
"""Status of every configured server for banner/TUI display.
Each dict has name, transport, tools, connected, disabled, status (one of connected /
disabled / connecting / failed / configured) and, for failed, error. ``enabled: false``
is reported as disabled, not failed.
"""
"""Per-server status dicts for banner/TUI: name, transport, tools, connected, disabled,
status (connected / disabled / connecting / failed / configured) and error for failed."""
configured = _core._load_mcp_config()
if not configured:
return []
@@ -448,35 +426,30 @@ def get_mcp_status() -> List[dict]:
connecting = set(_core._server_connecting)
connect_errors = dict(_core._server_connect_errors)
def _entry(name: str, transport: str, status: str, **extra) -> dict:
return {"name": name, "transport": transport, "tools": 0, "connected": False,
"disabled": status == "disabled", "status": status, **extra}
result: List[dict] = []
for name, cfg in configured.items():
transport = cfg.get("transport", "http") if "url" in cfg else "stdio"
enabled = _enabled(cfg) # evaluated unconditionally: malformed values warn even when connected
server = active_servers.get(name)
if server and server.session is not None:
entry = _entry(name, transport, "connected", connected=True)
live = server is not None and server.session is not None
status = ("connected" if live else "disabled" if not enabled else "connecting" if name in connecting
else "failed" if name in connect_errors else "configured")
entry = {"name": name, "transport": cfg.get("transport", "http") if "url" in cfg else "stdio",
"tools": 0, "connected": False, "disabled": status == "disabled", "status": status}
if live:
entry["connected"] = True
entry["tools"] = (len(server._registered_tool_names) if hasattr(server, "_registered_tool_names")
else len(server._tools))
if server._sampling:
entry["sampling"] = dict(server._sampling.metrics)
elif not _enabled(cfg):
entry = _entry(name, transport, "disabled")
elif name in connecting:
entry = _entry(name, transport, "connecting")
elif name in connect_errors:
entry = _entry(name, transport, "failed", error=connect_errors[name])
else:
entry = _entry(name, transport, "configured")
elif status == "failed":
entry["error"] = connect_errors[name]
result.append(entry)
return result
def probe_mcp_server_tools() -> Dict[str, List[tuple]]:
"""Connect to each enabled server, list ``(tool_name, description)`` and disconnect,
without registering anything. Failed servers are omitted."""
"""Connect each enabled server, list ``(tool_name, description)``, disconnect; nothing is
registered and failed servers are omitted."""
if not _core._ensure_mcp_sdk():
return {}
enabled = {k: v for k, v in (_core._load_mcp_config() or {}).items() if _enabled(v)}
@@ -509,15 +482,13 @@ def probe_mcp_server_tools() -> Dict[str, List[tuple]]:
def has_registered_mcp_tools() -> bool:
"""True if any MCP server has registered tools (cheap; no registry walk). Checks
registered TOOLS, not connected servers, so the per-turn refresh hook stays idle for
zero-tool servers."""
"""True if any MCP server has registered TOOLS (not merely connected), so the per-turn
refresh hook stays idle for zero-tool servers."""
with _core._lock:
return bool(_core._mcp_tool_server_names)
def get_registered_mcp_server_names() -> set:
"""Server names that registered at least one tool (the live, filtered signal — not
merely what config.yaml lists)."""
"""Server names that registered at least one tool (live, filtered — not config.yaml)."""
with _core._lock:
return set(_core._mcp_tool_server_names.values())
+42 -67
View File
@@ -67,11 +67,9 @@ def _check_circuit_breaker(server_name: str) -> Optional[str]:
def _acquire_call_server(server_name: str, tool_timeout: float):
"""``(server, None)`` when a call may be dispatched, else ``(None, error)``.
No session: a reconnect may be completing (fresh session swaps in asynchronously), so wait
briefly before charging a breaker strike. 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 and return a clean "reconnecting" error."""
"""``(server, None)`` when a call may be dispatched, else ``(None, error)``. No session: a
reconnect may be completing, so wait briefly before a breaker strike; still down -> ask the
server task to rebuild (probing a dead transport would re-arm the breaker forever)."""
not_connected = tool_error(f"MCP server '{server_name}' is not connected")
server = _core._get_connected_server_for_call(server_name)
if not server:
@@ -142,11 +140,9 @@ def _retry_once(server_name: str, retry_call, op_description: str, what: str):
# --------------------------------------------------------------- recovery ladder
def _handle_auth_error_and_retry(server_name: str, exc: BaseException, retry_call, op_description: str):
"""OAuth recovery + one retry; None when *exc* is not an auth error.
``MCPOAuthManager.handle_401`` decides 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."""
"""OAuth recovery + one retry; None when *exc* is not an auth error. ``handle_401`` decides
viability; if viable, signal a reconnect (fresh credentials), wait ready, retry once. Any
failure returns the structured ``needs_reauth`` error so the model stops refreshing."""
if not _core._is_auth_error(exc):
return None
from tools.mcp_oauth_manager import get_manager
@@ -158,9 +154,8 @@ def _handle_auth_error_and_retry(server_name: str, exc: BaseException, retry_cal
recovered = False
if recovered:
srv = _lookup_reconnectable_server(server_name)
# 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.
# Recovery + reconnect is independent evidence of viability: close the breaker here, not
# only on retry success (else a failing retry pins it open forever).
if srv is not None and _core._signal_reconnect_and_wait(
server_name, srv, op_description=f"{op_description} after OAuth recovery", timeout=15):
_core._reset_server_error(server_name)
@@ -176,9 +171,8 @@ def _handle_auth_error_and_retry(server_name: str, exc: BaseException, retry_cal
def _handle_session_expired_and_retry(server_name: str, exc: BaseException, retry_call, op_description: str):
"""Transport reconnect + one retry on session expiry; None to fall through (not
session-expired, no server record / loop, reconnect did not ready in time, retry failed).
Skips ``handle_401`` — the token is still valid, only the server-side session is stale."""
"""Transport reconnect + one retry on session expiry; None to fall through. Skips
``handle_401``: the token is valid, only the server-side session is stale."""
if not _is_session_expired_error(exc):
return None
srv = _lookup_reconnectable_server(server_name, require_loop=True)
@@ -194,15 +188,13 @@ def _handle_session_expired_and_retry(server_name: str, exc: BaseException, retr
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)."""
"""Stdio subprocess gone when (or while) a call ran. Deliberately NOT a TimeoutError."""
def _handle_stdio_child_exited_and_retry(server_name: str, exc: Exception, retry_call, op_description: str):
"""Respawn a dead stdio child and retry once; None if not our error.
Never spawns anything itself: 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."""
"""Respawn a dead stdio child and retry once; None if not our error. Never spawns itself: it
sets ``_reconnect_event`` and waits, so spawn frequency stays governed by ``run()``'s
rapid-drop budget. Single-shot: a child that dies again reports and stops."""
if not isinstance(exc, _StdioChildExited):
return None
reconnected = False
@@ -214,8 +206,7 @@ def _handle_stdio_child_exited_and_retry(server_name: str, exc: Exception, retry
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.
# No MCP loop to wait on (non-async adapters, tests): still request the respawn.
_core._signal_reconnect(srv)
if not reconnected:
return _strike(
@@ -227,8 +218,7 @@ def _handle_stdio_child_exited_and_retry(server_name: str, exc: Exception, retry
try:
return _record_call_outcome(server_name, 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.
# Died again right after respawn: broken server; 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)
return _strike(
@@ -246,12 +236,10 @@ def _handle_stdio_child_exited_and_retry(server_name: str, exc: Exception, retry
def _invoke_with_recovery(server_name: str, call_once: Callable[[], str], op: str,
recoverers, on_final_failure: Callable[[BaseException], None],
record_outcome: bool = False) -> str:
"""Run ``call_once``, walking the recovery ladder on failure. Each recoverer
``(server_name, exc, retry_call, op) -> Optional[str]`` returns None when the exception is
not its kind; order matters: dead stdio child -> auth -> session expiry. Unrecovered
exceptions go through ``on_final_failure`` (breaker strike / logging) and become the generic
call-failed error. ``record_outcome`` applies breaker bookkeeping to the FIRST attempt only;
retries own their bookkeeping inside the recoverers."""
"""Run ``call_once``, walking ``recoverers`` (``(server_name, exc, retry_call, op) ->
Optional[str]``, None = not its kind; order matters) on failure. Unrecovered exceptions go
through ``on_final_failure`` and become the generic call-failed error. ``record_outcome``
applies breaker bookkeeping to the FIRST attempt only; retries own theirs."""
try:
result = call_once()
return _record_call_outcome(server_name, result) if record_outcome else result
@@ -276,10 +264,9 @@ def _mark_server_call_started(server: Any) -> None:
@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.
A deliberate reconnect/shutdown teardown (``_fail_inflight_calls`` sets ``_reconnecting``
first) turns the cancel into a clean retryable RuntimeError; external cancels (caller
timeout, user interrupt) propagate unchanged. Doubles without ``_inflight_tasks`` skip tracking."""
"""Register the running RPC so teardown can fail it fast. A deliberate teardown
(``_reconnecting`` set first) turns the cancel into a retryable RuntimeError; external
cancels propagate unchanged. Doubles without ``_inflight_tasks`` skip tracking."""
inflight = getattr(server, "_inflight_tasks", None)
task = asyncio.current_task()
tracked = task is not None and inflight is not None
@@ -298,12 +285,11 @@ async def _track_inflight_rpc(server: Any, server_name: str, op: str):
async def _call_tool_racing_stdio_death(server, server_name: str, tool_name: str, args: dict):
"""``session.call_tool`` that fails fast when the stdio child is/gets dead.
Pre-call: an already-dead child must not hold the slot for the full tool timeout
(``server.session`` is stale so the transport-down path never fired). Mid-call: race the
RPC against ``_watch_stdio_children``. Both raise :class:`_StdioChildExited` for the
respawn-and-retry path, which owns the reconnect signal (nothing clears ``server.session``).
callable()/``is True`` checks because MagicMock attributes return truthy Mocks."""
"""``session.call_tool`` that fails fast when the stdio child is/gets dead: pre-call (a dead
child must not hold the slot for the full timeout; ``server.session`` is stale) and mid-call
(race against ``_watch_stdio_children``). Both raise :class:`_StdioChildExited` for the
respawn path, which owns the reconnect signal. callable()/``is True`` because MagicMock
attributes are truthy."""
_stdio_dead = getattr(server, "_stdio_children_dead", None)
if callable(_stdio_dead) and _stdio_dead() is True:
raise _StdioChildExited(f"MCP stdio subprocess for '{server_name}' had already exited when the call was dispatched")
@@ -337,9 +323,8 @@ def _error_result_text(result) -> str:
def _render_content_blocks(result, server_name: str) -> str:
"""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."""
"""Text passes through; image/audio blocks are cached (MEDIA: tags); resource blocks are
materialized rather than silently dropped."""
parts: List[str] = []
for block in (result.content or []):
if getattr(block, "text", None):
@@ -373,10 +358,8 @@ def _capped_structured_content(result):
def _render_call_tool_result(result, server_name: str) -> str:
"""Pure: ``CallToolResult`` -> the handler's JSON string. ``content`` is the primary
(model-oriented) payload; ``structuredContent`` supplements it (or becomes ``result`` when
there is no text). Server-level ``_meta`` is surfaced minus protocol-reserved keys.
``.is_error`` is ``.isError`` before mcp 2.0."""
"""Pure: ``CallToolResult`` -> handler JSON. ``content`` is primary; ``structuredContent``
supplements it (or becomes ``result`` without text); ``_meta`` minus reserved keys."""
if mcp_field(result, "is_error", "isError", False):
return tool_error(_sanitize_error(_truncate_mcp_text_result(_error_result_text(result) or "MCP tool returned an error")))
text_result = _render_content_blocks(result, server_name)
@@ -444,10 +427,9 @@ def _make_tool_handler(server_name: str, tool_name: str, tool_timeout: float):
def _make_utility_handler(server_name: str, tool_timeout: float, op: str, log_label: str,
rpc, render, required: Optional[str] = None):
"""Shared shape of the four utility handlers (resources/prompts): ``rpc(session, args,
server_name)`` is awaited under ``_rpc_lock``, ``render(result, server_name)`` builds the
JSON-able payload, ``required`` names a parameter validated before any transport work. The
wrapper owns the connected check and the auth / session-expired recovery ladder."""
"""Shared shape of the four utility handlers: ``rpc(session, args, server_name)`` awaited
under ``_rpc_lock``, ``render(result, server_name)`` -> JSON-able payload, ``required``
validated before any transport work; owns the connected check and recovery ladder."""
def _handler(args: dict, **kwargs) -> str:
server = _core._get_connected_server_for_call(server_name)
@@ -471,9 +453,8 @@ def _make_utility_handler(server_name: str, tool_timeout: float, op: str, log_la
def _pick(obj, *specs) -> dict:
"""``{out_key: value}`` for each ``(out_key, attr[, truthy])`` present on *obj*.
``hasattr`` (not a default) so SDK models and test stubs behave alike; with ``truthy``
the field is also skipped when falsy. Output key order = spec order."""
"""``{out_key: value}`` for each ``(out_key, attr[, truthy])`` present on *obj* (``hasattr``
so SDK models and stubs behave alike; ``truthy`` also skips falsy). Key order = spec order."""
entry = {}
for out_key, attr, *truthy in specs:
if not hasattr(obj, attr):
@@ -490,10 +471,9 @@ def _render_resource_list(all_resources, server_name: str) -> dict:
entry = _pick(r, ("uri", "uri"), ("name", "name"), ("description", "description", True))
if "uri" in entry:
entry["uri"] = str(entry["uri"])
# Key stays camelCase — this is the tool's own JSON output shape.
mime = mcp_field(r, "mime_type", "mimeType")
if mime:
entry["mimeType"] = mime
entry["mimeType"] = mime # camelCase: this is the tool's own JSON output shape
resources.append(entry)
return {"resources": resources}
@@ -504,8 +484,7 @@ def _render_read_resource(result, server_name: str) -> dict:
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).
# Binary contents go to the document cache (same contract as EmbeddedResource blocks).
rendered = _render_mcp_resource_block(SimpleNamespace(type="resource", resource=block), server_name)
parts.append(rendered or f"[binary data, {len(block.blob)} bytes]")
return {"result": "\n".join(parts)}
@@ -516,9 +495,8 @@ def _render_prompt_list(all_prompts, server_name: str) -> dict:
for p in all_prompts:
entry = _pick(p, ("name", "name"), ("description", "description", True))
if getattr(p, "arguments", None):
entry["arguments"] = [
{"name": a.name, **_pick(a, ("description", "description", True), ("required", "required"))}
for a in p.arguments]
entry["arguments"] = [{"name": a.name, **_pick(a, ("description", "description", True), ("required", "required"))}
for a in p.arguments]
prompts.append(entry)
return {"prompts": prompts}
@@ -530,10 +508,7 @@ def _render_get_prompt(result, server_name: str) -> dict:
if hasattr(msg, "content"):
entry["content"] = strip_unicode_tags(msg.content.text if hasattr(msg.content, "text") else str(msg.content))
messages.append(entry)
resp = {"messages": messages}
if getattr(result, "description", None):
resp["description"] = result.description
return resp
return {"messages": messages, **_pick(result, ("description", "description", True))}
def _utility_factory(op: str, log_label: str, rpc, render, required: Optional[str] = None):
+15 -23
View File
@@ -74,12 +74,10 @@ def _forget_mcp_tool_server(tool_name: str) -> None:
def _select_utility_schemas(server_name: str, server: "MCPServerTask", config: dict) -> List[dict]:
"""Utility schemas allowed by config (``tools.resources``/``tools.prompts``) and by the
server's advertised capabilities. ``initialize_result.capabilities`` is the source of
truth: its sub-objects are non-None iff the server advertises that request family (a
``hasattr(server.session, ...)`` gate never filters anything — ClientSession defines all
four methods). Without an initialize_result (test fixtures, older paths) fall back to
that legacy session-method check."""
"""Utility schemas allowed by config (``tools.resources``/``tools.prompts``) and advertised
capabilities. ``initialize_result.capabilities`` is the truth (sub-object non-None iff the
family is served); without it fall back to the legacy session-method check, which never
filters anything since ClientSession defines all four methods."""
tools_filter = config.get("tools") or {}
enabled = {f: _parse_boolish(tools_filter.get(f), default=True) for f in ("resources", "prompts")}
init_result = getattr(server, "initialize_result", None)
@@ -204,11 +202,9 @@ def _utility_candidates(name: str, entries: Iterable[Any], tool_timeout) -> List
def _resolve_name_collisions(name: str, candidates: List[_Candidate]) -> List[_Candidate]:
"""Preflight registry-name collisions among one server's candidates. Exact duplicates
(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). Returns the
survivors in order."""
"""Preflight name collisions: exact duplicates dropped silently; a utility normalizing onto
a native tool's name is shadowed (native wins); any other multi-origin collision skips every
colliding entry (fail closed). Returns survivors in order."""
unique: List[_Candidate] = []
seen: set[tuple[str, str]] = set()
origins_by_name: Dict[str, set[str]] = {}
@@ -264,9 +260,8 @@ def _log_foreign_owner(name: str, c: _Candidate, existing_toolset: str, lazy: bo
def _register_candidates(name: str, candidates: List[_Candidate], *, check_fn: Callable,
scope: Callable[[], Optional[str]], lazy: bool) -> List[str]:
"""Register candidates under toolset ``mcp-{name}``; returns the names that landed. The
ownership pre-check is advisory only — servers connect in parallel, so
``ToolRegistry.register()`` is the atomic ownership gate and its verdict is re-read after
every call."""
ownership pre-check is advisory (servers connect in parallel): ``registry.register()`` is
the atomic gate and its verdict is re-read after every call."""
from tools.registry import registry
toolset_name = f"mcp-{name}"
registered: List[str] = []
@@ -315,11 +310,9 @@ def _write_schema_cache(name: str, server: "MCPServerTask", config: dict, should
def _register_server_tools(name: str, server: "MCPServerTask", config: dict) -> List[str]:
"""Register an already-connected server's tools (plus utility tools); used by initial
discovery and list_changed refresh. Returns the registered names. 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."""
"""Register a connected server's tools plus utilities (initial discovery and list_changed
refresh); returns the names. Toolset aliases derive from the live registry, not
``toolsets.TOOLSETS``; lossy normalization collisions (``read-file``/``read_file``) fail closed."""
should_register = _make_tool_filter(name, config)
_record_tool_trust_metadata(name, config, server._tools)
candidates = _tool_candidates(name, server._tools, should_register, server.tool_timeout)
@@ -333,10 +326,9 @@ def _register_server_tools(name: str, server: "MCPServerTask", config: dict) ->
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 goes through ``_get_connected_server_for_call`` ->
``_ensure_lazy_server_connected``. Trust metadata is recorded first so the call-time gate
is identical whether the server was spawned live or registered from cache."""
"""Lazy startup: register from a cached manifest with no child process (first real call goes
through ``_ensure_lazy_server_connected``). Trust metadata is recorded first so the
call-time gate is identical for live and cached registrations."""
from tools.mcp_schema_cache import config_fingerprint, tools_from_cache_entry, utility_tools_from_cache_entry
tool_timeout = _resolve_tool_timeout(config)
cached_tools = _CachedMCPTool.from_cache_dicts(tools_from_cache_entry(entry))
+48 -80
View File
@@ -15,8 +15,7 @@ logger = logging.getLogger("tools.mcp_tool")
@dataclass
class _RetryBudget:
"""Per-run() retry counters shared by the branch helpers (``_reconnect_retries`` lives on
the task because handlers and tests read it)."""
"""Per-run() retry counters (``_reconnect_retries`` stays on the task: handlers/tests read it)."""
initial_retries: int = 0
backoff: float = 1.0
@@ -49,15 +48,11 @@ class MCPServerRunMixin:
return True
async def _wait_for_lifecycle_event(self) -> str:
"""Serve the connection until a lifecycle event; return its kind.
``"shutdown"`` exits the run loop; ``"reconnect"`` tears the session down and
re-enters the transport (event cleared before return); ``"recycle"`` means a stdio
idle/lifetime limit elapsed and the transport restarts lazily on the next call.
Shutdown wins a tie. Between events a keepalive (``ping``, list_tools fallback) runs
every ``keepalive_interval`` — which must stay below the server's session TTL — and a
failure triggers a reconnect.
"""
"""Serve until a lifecycle event: ``"shutdown"`` (exits run), ``"reconnect"`` (session torn
down, transport re-entered; event cleared first) or ``"recycle"`` (stdio idle/lifetime
limit; restarts lazily on next call). Shutdown wins a tie. A keepalive (``ping``,
list_tools fallback) runs every ``keepalive_interval`` (must stay below the server's
session TTL); a failure triggers a reconnect."""
keepalive_interval = max(
_core._MIN_KEEPALIVE_INTERVAL,
float(self._config.get("keepalive_interval", _core._DEFAULT_KEEPALIVE_INTERVAL)))
@@ -76,9 +71,8 @@ class MCPServerRunMixin:
break
if self._recycle_if_due():
return "recycle"
# Timeout: probe for a stale session — but NEVER while an RPC is in flight (a
# concurrent ping can wedge the single stdio stream, and a busy server is
# provably alive anyway).
# Timeout: probe for a stale session — NEVER while an RPC is in flight (a
# concurrent ping can wedge the stdio stream; a busy server is alive anyway).
if self.session:
if self._rpc_lock.locked() or any(not t.done() for t in self._inflight_tasks):
continue
@@ -101,16 +95,14 @@ class MCPServerRunMixin:
if self._shutdown_event.is_set():
self._fail_inflight_calls("shutdown")
return "shutdown"
# Deliberate teardown: fail in-flight RPCs NOW rather than letting them ride the dying
# transport to the full tool timeout.
# Deliberate teardown: fail in-flight RPCs NOW instead of riding out the tool timeout.
self._fail_inflight_calls("reconnect")
self._reconnect_event.clear()
return "reconnect"
async def _wait_for_reconnect_or_shutdown(self, timeout: Optional[float] = None) -> str:
"""Wait, while parked, for a reconnect request or shutdown. Returns ``"shutdown"`` or
``"reconnect"`` (explicit request or, with ``timeout``, the periodic self-probe); the
reconnect event is cleared first. Shutdown wins a tie."""
"""Parked wait: ``"shutdown"`` or ``"reconnect"`` (explicit, or the ``timeout`` self-probe;
event cleared first). Shutdown wins a tie."""
shutdown_task, reconnect_task = self._event_waiters()
try:
await asyncio.wait({shutdown_task, reconnect_task}, return_when=asyncio.FIRST_COMPLETED, timeout=timeout)
@@ -122,15 +114,11 @@ class MCPServerRunMixin:
return "reconnect"
async def _park(self, revival_reason: str) -> bool:
"""Drop this server's tools and wait for a reconnect request.
The run task must NOT exit: it is the only listener on ``_reconnect_event``, so
returning leaves the server unrevivable for the life of the process. Parking
deregisters the tools, so no call can reach the breaker probe or ``_signal_reconnect``;
the wait is therefore TIMED (one self-probe per ``_PARKED_RETRY_INTERVAL``), and an
explicit ``_reconnect_event.set()`` wakes it immediately. True when shutdown was
requested instead.
"""
"""Drop this server's tools and wait for a reconnect request; True when shutdown came instead.
The run task must NOT exit (it is the only ``_reconnect_event`` listener, so returning
leaves the server unrevivable). With tools deregistered no call can reach the breaker
probe, so the wait is TIMED (one self-probe per ``_PARKED_RETRY_INTERVAL``); an explicit
``_reconnect_event.set()`` wakes it immediately."""
self._was_parked = True
self._deregister_tools()
self._reconnect_event.clear()
@@ -141,12 +129,9 @@ class MCPServerRunMixin:
return False
async def _prepare_run(self, config: dict) -> bool:
"""Bind config, build sampling/elicitation handlers, validate HTTP.
Returns False when the server must not start: a bad remote URL or a non-MCP endpoint
(both fail fast, non-retryably, with ``_error`` set and ``_ready`` fired) instead of
burning the reconnect ladder inside the SDK's httpx layer on every retry.
"""
"""Bind config, build sampling/elicitation handlers, validate HTTP. False when the server
must not start (bad remote URL / non-MCP endpoint: fail fast with ``_error`` set and
``_ready`` fired instead of burning the reconnect ladder inside the SDK's httpx layer)."""
self._config = config
self.tool_timeout = _core._resolve_tool_timeout(config)
self._auth_type = (config.get("auth") or "").lower().strip()
@@ -170,11 +155,9 @@ class MCPServerRunMixin:
return True
try:
_core._validate_remote_mcp_url(self.name, config.get("url"))
# Content-type preflight (Streamable HTTP only; SSE legitimately serves
# text/event-stream): a URL at a web-app root returns HTML and would make the SDK
# hang for the full connect_timeout. Skipped once _ready was ever set (endpoint
# already validated) and for OAuth servers, where a token-less probe sees
# HTML/401 and would block the flow.
# Content-type preflight (Streamable HTTP only; SSE serves text/event-stream): a
# web-app root returns HTML and would hang the SDK for connect_timeout. Skipped once
# _ready was ever set and for OAuth servers (a token-less probe sees HTML/401).
if (config.get("transport") != "sse" and not config.get("skip_preflight")
and not self._ready.is_set() and self._auth_type != "oauth"):
await self._preflight_content_type(
@@ -190,14 +173,10 @@ class MCPServerRunMixin:
return True
async def run(self, config: dict):
"""Long-lived coroutine: connect, discover, serve, reconnect.
State machine: connecting -> connected -> (degraded -> parked -> revived)*. Unproven
drops and transport errors charge a rapid-drop budget with jittered exponential
backoff; exhausting it (or a permanent error) parks the server via :meth:`_park`
rather than exiting, so it stays revivable. The branch helpers return True to keep
looping and False to exit the loop.
"""
"""Long-lived: connecting -> connected -> (degraded -> parked -> revived)*. Unproven drops
and transport errors charge a rapid-drop budget with jittered backoff; exhausting it (or
a permanent error) parks via :meth:`_park` rather than exiting, so the server stays
revivable. Branch helpers return True to keep looping, False to exit."""
if not await self._prepare_run(config):
return
self._reconnect_retries = 0
@@ -208,8 +187,7 @@ class MCPServerRunMixin:
if not await self._on_clean_return(await run_transport(config), budget):
break
except asyncio.CancelledError:
# Not a connection failure: re-raise so cancellation reaches asyncio and
# shutdown()'s ``await self._task`` completes.
# Not a connection failure: re-raise so shutdown()'s ``await self._task`` completes.
self.session = None
raise
except Exception as exc:
@@ -222,9 +200,8 @@ class MCPServerRunMixin:
self._stdio_child_pids = set()
async def _on_clean_return(self, lifecycle_reason: str, budget: "_RetryBudget") -> bool:
"""Transport returned cleanly: shutdown, stdio recycle, or a requested rebuild (auth
recovery / manual refresh / keepalive failure). A rebuild is not a failure for the
retry counters."""
"""Clean transport return: shutdown, stdio recycle, or a requested rebuild (not a failure
for the retry counters)."""
if self._shutdown_event.is_set():
return False
if lifecycle_reason == "recycle":
@@ -235,9 +212,8 @@ class MCPServerRunMixin:
return await self._wait_for_reconnect_or_shutdown() != "shutdown"
# Per-cycle chatter stays DEBUG; WARNINGs mark state transitions.
logger.debug("MCP server '%s': reconnecting (OAuth recovery or manual refresh)", self.name)
# A clean return is NOT proof of health (a flapping transport handshakes fine and drops
# moments later). Only a PROVEN session clears the budget; a teardown race is
# recovery, not a failure, and must never reach the park on its own.
# A clean return is NOT proof of health (a flapper handshakes fine, then drops). Only a
# PROVEN session clears the budget; a teardown race is recovery, never a park charge.
if self._teardown_race and not self._session_proven:
logger.info("MCP server '%s': reconnect after teardown race "
"(in-flight calls were failed); not charging the "
@@ -256,23 +232,21 @@ class MCPServerRunMixin:
self.name, _core._MAX_RECONNECT_RETRIES, _core._PARKED_RETRY_INTERVAL)
if not await self._park_and_rearm("from parked state", budget):
return False
# Clear readiness too: a stale _ready lets handler-side recovery mistake the old
# session for a fresh one.
# Clear readiness too: a stale _ready lets handler recovery mistake old for fresh.
self._ready.clear()
self.session = None
return True
async def _park_and_rearm(self, revival_reason: str, budget: "_RetryBudget") -> bool:
"""Park; on revival leave a budget of ONE probe per wake so a still-dead server parks
again instead of burning 5 rapid retries. False on shutdown."""
"""Park; on revival leave ONE probe per wake so a still-dead server re-parks instead of
burning 5 rapid retries. False on shutdown."""
if await self._park(revival_reason):
return False
self._reconnect_retries, budget.backoff = _core._MAX_RECONNECT_RETRIES, 1.0
return True
async def _park_initial_failure(self, exc: Exception, revival_reason: str, budget: "_RetryBudget") -> bool:
"""Publish ``exc`` to the waiting ``start()``, park, and on revival reset every counter
so the ladder starts fresh. False on shutdown."""
"""Publish ``exc`` to ``start()``, park, and on revival reset every counter. False on shutdown."""
self._error = exc
self._ready.set()
if await self._park(revival_reason):
@@ -288,10 +262,8 @@ class MCPServerRunMixin:
budget.backoff = min(budget.backoff * 2, _core._MAX_BACKOFF_SECONDS)
async def _on_transport_error(self, exc: Exception, budget: "_RetryBudget") -> bool:
"""Transport raised: classify, then run the initial-connect or the reconnect ladder.
Returns False when the run loop must exit."""
# Unwrap anyio TaskGroup wrappers: the group's str() is useless and hides the root
# cause from the classification below.
"""Transport raised: classify, then run the initial-connect or reconnect ladder. False = exit."""
# Unwrap anyio TaskGroup wrappers: the group's str() hides the root cause.
root = _core._unwrap_exception_group(exc)
failure_class = _core._classify_mcp_failure(root)
if self._is_recycled_stdio():
@@ -299,8 +271,8 @@ class MCPServerRunMixin:
"failed, marking unavailable while retrying: %s: %s",
self.name, type(root).__name__, root)
self._recycled_reason = None
# Initial-connect ladder: a transient blip at startup must not kill the server. Gated
# on _ever_connected (never cleared), not _ready (cleared every reconnect cycle).
# Initial-connect ladder (a startup blip must not kill the server); gated on
# _ever_connected, not _ready (which clears every reconnect cycle).
if not self._ever_connected:
return await self._on_initial_connect_error(exc, root, failure_class, budget)
if self._shutdown_event.is_set():
@@ -328,9 +300,8 @@ class MCPServerRunMixin:
async def _on_initial_connect_error(self, exc: Exception, root: BaseException,
failure_class: str, budget: "_RetryBudget") -> bool:
if failure_class == "permanent":
# Deterministic failure (bad command, non-MCP URL, 401/403): park at once instead
# of burning the ladder. Auth failures park rather than return so the task stays
# alive to pick up fresh tokens later.
# Deterministic failure (bad command, non-MCP URL, 401/403): park at once; auth
# failures park (not return) so the task can pick up fresh tokens later.
detail = (f"authentication, parking until credentials change; re-authenticate with "
f"`hermes mcp login {self.name}`" if _core._is_auth_error(root)
else "connection with a permanent error, parking without retries")
@@ -356,8 +327,8 @@ class MCPServerRunMixin:
return not self._shutdown_event.is_set()
async def _on_permanent_error(self, root: BaseException, budget: "_RetryBudget") -> bool:
# An auth failure on a PROVEN session is often a corrupt OAuth lock from a raced
# teardown, not revoked credentials: grant ONE suspect+reconnect cycle first.
# Auth failure on a PROVEN session is often a raced-teardown OAuth lock, not revoked
# credentials: grant ONE suspect+reconnect cycle first.
if _core._is_auth_error(root) and self._session_proven and not self._permanent_grace_used:
self._permanent_grace_used = True
self.mark_suspect(f"auth error on proven session: {root}")
@@ -382,11 +353,9 @@ class MCPServerRunMixin:
try:
await self._ready.wait()
except asyncio.CancelledError:
# The caller's connect timeout (discover_mcp_tools wraps start() in
# asyncio.wait_for) cancels *this* coroutine, but the ensure_future'd run() task
# is independent and would otherwise keep running detached — parked on a hung
# transport with no owner to reap it. Propagate the cancellation so the transport
# context managers unwind and release the child process / FDs.
# The caller's connect timeout cancels *this* coroutine; the ensure_future'd run()
# task would otherwise keep running detached on a hung transport with no owner.
# Propagate so the transport context managers unwind and release child / FDs.
if self._task and not self._task.done():
self._task.cancel()
raise
@@ -418,9 +387,8 @@ class MCPServerRunMixin:
self.session = None
def _deregister_tools(self) -> None:
"""Drop this server's tools from the global registry (idempotent). Called on shutdown
AND when the reconnect budget is exhausted, so a dead server never leaves phantom tool
definitions bloating the prompt cache and producing "not connected" errors."""
"""Drop this server's tools from the registry (idempotent); on shutdown AND budget
exhaustion, so a dead server never leaves phantom tools in the prompt."""
from tools.registry import registry
for tool_name in list(getattr(self, "_registered_tool_names", [])):
registry.deregister(tool_name, scope=_core._server_registry_scope(self.name))