refactor(tools): compact MCP transport/loop/sampling/errors helpers, defensive collapse

This commit is contained in:
Teknium
2026-09-02 22:16:12 -07:00
parent 113f04616b
commit 1ad360b3aa
4 changed files with 174 additions and 339 deletions

View File

@@ -1,4 +1,6 @@
"""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."""
"""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 asyncio
import errno
@@ -12,7 +14,6 @@ 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
@@ -30,7 +31,7 @@ def _jsonrpc_matches(exc: BaseException, code, codes: tuple, markers: tuple) ->
if code in codes:
return True
msg = str(exc).lower()
return bool(msg) and any(marker in msg for marker in markers)
return any(marker in msg for marker in markers)
def _handshake_rejected_as_modern(exc: BaseException) -> bool:
@@ -98,16 +99,11 @@ def _classify_mcp_failure(exc: BaseException) -> str:
or isinstance(root, (NonMcpEndpointError, InvalidMcpUrlError, FileNotFoundError))
or (isinstance(root, OSError) and getattr(root, "errno", None) == errno.ENOENT)
# 401/403 HTTPStatusError that _is_auth_error's type-gate missed (auth types not importable here).
or _response_status(root) in (401, 403)
or getattr(getattr(root, "response", None), "status_code", None) in (401, 403)
)
return "permanent" if permanent else "transient"
def _response_status(exc: BaseException):
"""``exc.response.status_code`` for httpx-shaped errors, else None."""
return getattr(getattr(exc, "response", None), "status_code", None)
def _validate_remote_mcp_url(server_name: str, url: Any) -> str:
"""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``),
@@ -252,7 +248,6 @@ def _exc_children(exc: BaseException) -> List[BaseException]:
def _format_connect_error(exc: BaseException) -> str:
"""Render nested MCP connection errors into an actionable short message."""
def _find_missing(current: BaseException) -> Optional[str]:
if isinstance(current, FileNotFoundError):
if getattr(current, "filename", None):
@@ -318,8 +313,7 @@ def _get_auth_error_types() -> tuple:
_optional_types("mcp.client.auth", "OAuthFlowError", "OAuthTokenError")
+ _optional_types("mcp.client.auth", "UnauthorizedError") # older SDKs
+ _optional_types("tools.mcp_oauth", "OAuthNonInteractiveError")
+ list(_http_status_error_types())
)
+ list(_http_status_error_types()))
return _AUTH_ERROR_TYPES
@@ -352,16 +346,15 @@ _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."""
:func:`_is_auth_error` the fix is a transport reconnect (``_reconnect_event``), not an OAuth refresh.
Iterative walk 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* a message-less ClosedResourceError."""
# AnyIO stream exceptions are often message-less (``str(ClosedResourceError()) == ""``),
# so type checks are needed in addition to marker matching.
transport_error_types = tuple(_optional_types("anyio", "BrokenResourceError", "ClosedResourceError", "EndOfStream"))
# 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
@@ -377,7 +370,7 @@ def _is_session_expired_error(exc: BaseException) -> bool:
# 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 isinstance(current, transport_error_types) or (msg and any(marker in msg for marker in _SESSION_EXPIRED_MARKERS)):
if isinstance(current, transport_error_types) or any(marker in msg for marker in _SESSION_EXPIRED_MARKERS):
transport_error_found = True
stack.extend(getattr(current, "exceptions", ()))
stack.extend((getattr(current, "__cause__", None), getattr(current, "__context__", None)))

View File

@@ -8,6 +8,7 @@ from __future__ import annotations
import asyncio
import concurrent.futures
import contextlib
import errno
import logging
import os
@@ -20,11 +21,8 @@ logger = logging.getLogger("tools.mcp_tool")
class _LockCookie:
"""Holds a cross-process file lock; ``release()`` drops it.
The file object MUST stay open while the lock is held: both the fcntl and
the portalocker lock are tied to the descriptor's lifetime.
"""
"""Holds a cross-process file lock; ``release()`` drops it. The file object MUST stay open
while the lock is held: both the fcntl and the portalocker lock are tied to the descriptor."""
def __init__(self, fh: Any) -> None:
self._fh = fh
@@ -32,47 +30,37 @@ class _LockCookie:
def release(self) -> None:
if self._fh is None:
return
# Best effort on every step: an unlock/close failure must never
# propagate out of discovery.
try:
# Best effort on every step: an unlock/close failure must never propagate out of discovery.
with contextlib.suppress(Exception):
if os.name == "posix":
import fcntl
fcntl.flock(self._fh.fileno(), fcntl.LOCK_UN)
else:
import portalocker
portalocker.unlock(self._fh)
except Exception:
pass
try:
with contextlib.suppress(Exception):
self._fh.close()
except Exception:
pass
self._fh = None
def _acquire_lock_on_fh(fh: Any) -> bool:
"""Non-blocking exclusive lock (fcntl on POSIX, portalocker elsewhere).
False when another process holds it; unexpected errors propagate so the
caller can treat locking as unavailable.
"""
fd = fh.fileno()
"""Non-blocking exclusive lock (fcntl on POSIX, portalocker elsewhere). False when another
process holds it; unexpected errors propagate so the caller can treat locking as unavailable."""
if os.name == "posix":
import fcntl
try:
fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
fcntl.flock(fh.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)
return True
except OSError as e:
if e.errno in (errno.EACCES, errno.EAGAIN, errno.EWOULDBLOCK):
return False
raise
else:
import portalocker
try:
portalocker.lock(fh, portalocker.LOCK_EX | portalocker.LOCK_NB)
return True
except portalocker.LockException:
return False
import portalocker
try:
portalocker.lock(fh, portalocker.LOCK_EX | portalocker.LOCK_NB)
return True
except portalocker.LockException:
return False
def _try_acquire_mcp_discovery_lock() -> Any:
@@ -84,24 +72,15 @@ def _try_acquire_mcp_discovery_lock() -> Any:
try:
from hermes_constants import get_hermes_home
if _origin._MCP_DISCOVERY_LOCK_PATH is None:
_origin._MCP_DISCOVERY_LOCK_PATH = str(
get_hermes_home() / ".mcp-discovery.lock"
)
lock_path = _origin._MCP_DISCOVERY_LOCK_PATH
_origin._MCP_DISCOVERY_LOCK_PATH = str(get_hermes_home() / ".mcp-discovery.lock")
fh = open(_origin._MCP_DISCOVERY_LOCK_PATH, "w", encoding="utf-8")
except Exception:
return _core._LOCK_UNAVAILABLE
try:
fh = open(lock_path, "w", encoding="utf-8")
except Exception:
return _core._LOCK_UNAVAILABLE
try:
acquired = _core._acquire_lock_on_fh(fh)
except Exception:
fh.close()
return _core._LOCK_UNAVAILABLE
if acquired:
return _core._LockCookie(fh)
fh.close()
@@ -121,12 +100,7 @@ def _wrap_with_home_override(coro: "Coroutine") -> "Coroutine":
"""Carry the caller's context-local HERMES_HOME override into ``coro``
(task-local on the MCP loop, so concurrent scopes don't interfere)."""
try:
from hermes_constants import (
get_hermes_home_override,
reset_hermes_home_override,
set_hermes_home_override,
)
from hermes_constants import get_hermes_home_override, reset_hermes_home_override, set_hermes_home_override
home_override = get_hermes_home_override()
except Exception:
return coro
@@ -146,11 +120,7 @@ def _wrap_with_home_override(coro: "Coroutine") -> "Coroutine":
def _wrap_with_dashboard_oauth_flow(coro):
"""Propagate a dashboard OAuth flow onto the dedicated MCP loop task."""
try:
from tools.mcp_dashboard_oauth import (
dashboard_oauth_flow,
get_dashboard_oauth_flow,
)
from tools.mcp_dashboard_oauth import dashboard_oauth_flow, get_dashboard_oauth_flow
flow = get_dashboard_oauth_flow()
except Exception:
return coro
@@ -172,12 +142,9 @@ def _running_loop() -> Optional[asyncio.AbstractEventLoop]:
def _run_on_mcp_loop(coro_or_factory, timeout: float = 30):
"""Schedule a coroutine on the MCP loop and block until done.
Accepts a coroutine or a zero-arg factory (a factory avoids leaking a
never-awaited coroutine when the loop is down). Polls in short intervals
so the calling thread can honor user interrupts.
"""
"""Schedule a coroutine on the MCP loop and block until done. Accepts a coroutine or a
zero-arg factory (a factory avoids leaking a never-awaited coroutine when the loop is down).
Polls in short intervals so the calling thread can honor user interrupts."""
from tools.interrupt import is_interrupted
from agent.async_utils import safe_schedule_threadsafe
@@ -186,58 +153,42 @@ def _run_on_mcp_loop(coro_or_factory, timeout: float = 30):
if asyncio.iscoroutine(coro_or_factory):
coro_or_factory.close()
raise RuntimeError("MCP event loop is not running")
coro = coro_or_factory() if callable(coro_or_factory) else coro_or_factory
# Tasks created via run_coroutine_threadsafe copy the LOOP thread's
# context, so a per-request profile scope would vanish here; re-establish
# it inside the task's own context.
# Tasks created via run_coroutine_threadsafe copy the LOOP thread's context, so a per-request
# profile scope would vanish here; re-establish it inside the task's own context.
coro = _core._wrap_with_home_override(coro)
coro = _core._wrap_with_dashboard_oauth_flow(coro)
future = safe_schedule_threadsafe(
coro, loop,
logger=logger,
log_message="MCP scheduling failed",
)
future = safe_schedule_threadsafe(coro, loop, logger=logger, log_message="MCP scheduling failed")
if future is None:
raise RuntimeError("MCP event loop unavailable (failed to schedule)")
start_time = time.monotonic()
deadline = None if timeout is None else start_time + timeout
while True:
if is_interrupted():
future.cancel()
raise InterruptedError("User sent a new message")
wait_timeout = 0.1
if deadline is not None:
remaining = deadline - time.monotonic()
if remaining <= 0:
future.cancel()
elapsed = time.monotonic() - start_time
raise TimeoutError(
f"MCP call timed out after {elapsed:.1f}s "
f"(configured timeout: {float(timeout):.1f}s)"
)
raise TimeoutError(f"MCP call timed out after {elapsed:.1f}s "
f"(configured timeout: {float(timeout):.1f}s)")
wait_timeout = min(wait_timeout, remaining)
try:
return future.result(timeout=wait_timeout)
except concurrent.futures.TimeoutError:
# Aliases builtin TimeoutError, so this also fires for the
# coroutine's own timeout: a done future must yield its outcome.
# Aliases builtin TimeoutError, so this also fires for the coroutine's own timeout: a
# done future must yield its outcome.
if future.done():
return future.result()
continue
def _signal_reconnect(server: Any) -> bool:
"""Ask a server task to rebuild its transport, thread-safely.
Handlers run on caller threads while the event lives on the MCP loop, so
it is set via ``call_soon_threadsafe`` when the loop runs (direct
``.set()`` otherwise). False when the server has no reconnect machinery.
"""
"""Ask a server task to rebuild its transport, thread-safely. Handlers run on caller threads
while the event lives on the MCP loop, so it is set via ``call_soon_threadsafe`` when the loop
runs (direct ``.set()`` otherwise). False when the server has no reconnect machinery."""
event = getattr(server, "_reconnect_event", None)
if event is None:
return False
@@ -258,30 +209,20 @@ def reconnect_mcp_server(server_name: str) -> bool:
return _core._signal_reconnect(server)
def _wait_for_server_session_ready(
srv: Any,
*,
old_session: Any = None,
timeout: float = 15.0,
) -> bool:
"""Poll until the server exposes a usable, ready session.
During a reconnect ``srv.session`` is briefly None or still the stale
object; retrying blindly there burns breaker strikes. With
``old_session`` the observed session must differ from it. Iteration-
bounded, not deadline-bounded: tests freeze ``time.monotonic``.
"""
def _wait_for_server_session_ready(srv: Any, *, old_session: Any = None, timeout: float = 15.0) -> bool:
"""Poll until the server exposes a usable, ready session. During a reconnect ``srv.session``
is briefly None or still the stale object; retrying blindly there burns breaker strikes. With
``old_session`` the observed session must differ from it. Iteration-bounded, not
deadline-bounded: tests freeze ``time.monotonic``."""
poll_interval = 0.25
iterations = max(1, int(max(float(timeout), 0.0) / poll_interval))
for i in range(iterations):
session = getattr(srv, "session", None)
ready = getattr(srv, "_ready", None)
is_ready = True
if ready is not None and hasattr(ready, "is_set"):
try:
is_ready = bool(ready.is_set())
except Exception:
is_ready = True
try:
is_ready = bool(ready.is_set()) if hasattr(ready, "is_set") else True
except Exception:
is_ready = True
if session is not None and session is not old_session and is_ready:
return True
if i < iterations - 1:
@@ -289,19 +230,10 @@ def _wait_for_server_session_ready(
return False
def _signal_reconnect_and_wait(
server_name: str,
srv: Any,
*,
op_description: str,
timeout: float = 15.0,
) -> bool:
"""Request a transport rebuild and wait for the fresh session.
``_ready`` is cleared on the loop BEFORE ``_reconnect_event`` is set;
otherwise the readiness poll returns immediately and retries against the
same dead session.
"""
def _signal_reconnect_and_wait(server_name: str, srv: Any, *, op_description: str, timeout: float = 15.0) -> bool:
"""Request a transport rebuild and wait for the fresh session. ``_ready`` is cleared on the
loop BEFORE ``_reconnect_event`` is set; otherwise the readiness poll returns immediately and
retries against the same dead session."""
loop = _core._mcp_loop
if loop is None or not loop.is_running():
return False
@@ -309,41 +241,28 @@ def _signal_reconnect_and_wait(
def _request_reconnect() -> None:
ready = getattr(srv, "_ready", None)
if ready is not None and hasattr(ready, "clear"):
if hasattr(ready, "clear"):
ready.clear()
reconnect_event = getattr(srv, "_reconnect_event", None)
if reconnect_event is not None and hasattr(reconnect_event, "set"):
if hasattr(reconnect_event, "set"):
reconnect_event.set()
logger.info(
"MCP server '%s': %s requesting transport reconnect",
server_name, op_description,
)
logger.info("MCP server '%s': %s requesting transport reconnect", server_name, op_description)
loop.call_soon_threadsafe(_request_reconnect)
return _core._wait_for_server_session_ready(
srv,
old_session=old_session,
timeout=timeout,
)
return _core._wait_for_server_session_ready(srv, old_session=old_session, timeout=timeout)
def _ensure_mcp_loop():
"""Start the background event loop thread if not already running.
The loop/thread handles live on the ORIGIN module (tests read and reset
``tools.mcp_tool._mcp_loop``), so they are written there, never here.
"""
"""Start the background event loop thread if not already running. The loop/thread handles live
on the ORIGIN module (tests read and reset ``tools.mcp_tool._mcp_loop``), so they are written
there, never here."""
from tools import mcp_tool as _origin
with _core._lock:
if _origin._mcp_loop is not None and _origin._mcp_loop.is_running():
return
_origin._mcp_loop = asyncio.new_event_loop()
_origin._mcp_loop.set_exception_handler(_core._mcp_loop_exception_handler)
_origin._mcp_thread = threading.Thread(
target=_origin._mcp_loop.run_forever,
name="mcp-event-loop",
daemon=True,
)
_origin._mcp_thread = threading.Thread(target=_origin._mcp_loop.run_forever, name="mcp-event-loop", daemon=True)
_origin._mcp_thread.start()
@@ -359,38 +278,29 @@ def _stop_mcp_loop(*, only_if_idle: bool = False) -> bool:
_origin._mcp_loop = None
_origin._mcp_thread = None
if loop is not None:
# Drain before stopping: tasks still suspended when the loop closes
# get resumed by the GC against a closed loop. shutdown_mcp_servers
# only reaps servers held in _servers; everything else ends up here.
# Drain before stopping: tasks still suspended when the loop closes get resumed by the GC
# against a closed loop. shutdown_mcp_servers only reaps servers held in _servers;
# everything else ends up here.
stop_owned_by_loop = False
if loop.is_running():
from agent.async_utils import safe_schedule_threadsafe
future = safe_schedule_threadsafe(
_core._drain_and_stop_mcp_loop(), loop,
logger=logger,
log_message="MCP loop drain: failed to schedule",
log_level=logging.WARNING,
)
_core._drain_and_stop_mcp_loop(), loop, logger=logger,
log_message="MCP loop drain: failed to schedule", log_level=logging.WARNING)
if future is not None:
stop_owned_by_loop = True
try:
future.result(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT + 1)
except TimeoutError:
logger.warning(
"Timed out waiting for MCP loop drain after %.1fs",
_core._MCP_LOOP_DRAIN_TIMEOUT + 1,
)
logger.warning("Timed out waiting for MCP loop drain after %.1fs", _core._MCP_LOOP_DRAIN_TIMEOUT + 1)
except BaseException as exc:
logger.warning("Error draining MCP loop tasks: %s", exc)
elif not loop.is_closed():
try:
loop.run_until_complete(
_core._drain_mcp_loop_tasks(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT)
)
loop.run_until_complete(_core._drain_mcp_loop_tasks(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT))
except BaseException as exc:
logger.warning("Error draining stopped MCP loop tasks: %s", exc)
if not stop_owned_by_loop and loop.is_running():
loop.call_soon_threadsafe(loop.stop)
if thread is not None:

View File

@@ -13,9 +13,8 @@ logger = logging.getLogger("tools.mcp_tool")
def _tool_use_id(block):
"""Tool-use id (the discriminator for a tool *result* block), read under both
SDK spellings — on mcp 2.x a bare ``hasattr(b, "toolUseId")`` is False and
would silently drop tool results."""
"""Tool-use id (the discriminator for a tool *result* block), read under both SDK spellings —
on mcp 2.x a bare ``hasattr(b, "toolUseId")`` is False and would silently drop tool results."""
return mcp_field(block, "tool_use_id", "toolUseId", _MISSING)
@@ -45,22 +44,15 @@ def _content_part(block) -> Optional[dict]:
def _tool_call_dict(tu, index: int) -> dict:
args = tu.input
return {
"id": getattr(tu, "id", f"call_{index}"),
"type": "function",
"function": {
"name": tu.name,
"arguments": json.dumps(args, ensure_ascii=False) if isinstance(args, dict) else str(args),
},
}
return {"id": getattr(tu, "id", f"call_{index}"), "type": "function", "function": {
"name": tu.name, "arguments": json.dumps(args, ensure_ascii=False) if isinstance(args, dict) else str(args)}}
def _convert_sampling_message(msg) -> List[dict]:
"""One MCP SamplingMessage -> OpenAI-format messages (tool results first,
then either an assistant tool_calls message or plain content)."""
blocks = msg.content_as_list if hasattr(msg, "content_as_list") else (
msg.content if isinstance(msg.content, list) else [msg.content]
)
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)]
@@ -89,10 +81,8 @@ def _parse_tool_call_arguments(server_name: str, args) -> dict:
try:
return json.loads(args)
except (json.JSONDecodeError, ValueError):
logger.warning(
"MCP server '%s': malformed tool_calls arguments from LLM (wrapping as raw): %.100s",
server_name, args,
)
logger.warning("MCP server '%s': malformed tool_calls arguments from LLM (wrapping as raw): %.100s",
server_name, args)
return {"_raw": args}
return args if isinstance(args, dict) else {"_raw": str(args)}
@@ -102,15 +92,13 @@ def _response_total_tokens(response, default):
class SamplingHandler:
"""Handles sampling/createMessage requests for one MCP server; 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``.
"""Handles sampling/createMessage requests for one MCP server; 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``.
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.
"""
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."""
_STOP_REASON_MAP = {"stop": "endTurn", "length": "maxTokens", "tool_calls": "toolUse"}
_LOG_LEVELS = {"debug": logging.DEBUG, "info": logging.INFO, "warning": logging.WARNING}
@@ -147,8 +135,7 @@ class SamplingHandler:
return None
def _convert_messages(self, params) -> List[dict]:
"""Convert MCP SamplingMessages to OpenAI format (``content_as_list``
when the SDK provides it; per-block duck-typed dispatch)."""
"""MCP SamplingMessages -> OpenAI format (per-block duck-typed dispatch)."""
return [m for msg in params.messages for m in _convert_sampling_message(msg)]
@staticmethod
@@ -164,14 +151,12 @@ class SamplingHandler:
return self._error(message)
def _log_response(self, response, suffix: str = "", *args) -> None:
logger.log(
self.audit_level, "MCP server '%s' sampling response: model=%s, tokens=%s" + suffix,
self.server_name, response.model, _response_total_tokens(response, "?"), *args,
)
logger.log(self.audit_level, "MCP server '%s' sampling response: model=%s, tokens=%s" + suffix,
self.server_name, response.model, _response_total_tokens(response, "?"), *args)
def _build_tool_use_result(self, choice, response):
"""Build a CreateMessageResultWithTools from an LLM tool_calls response,
subject to tool-loop governance (``max_tool_rounds``; 0 disables)."""
"""CreateMessageResultWithTools from an LLM tool_calls response, subject to tool-loop
governance (``max_tool_rounds``; 0 disables)."""
self.metrics["tool_use_count"] += 1
if self.max_tool_rounds == 0:
self._tool_loop_count = 0
@@ -180,61 +165,46 @@ class SamplingHandler:
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}' (max {self.max_tool_rounds} rounds)"
)
f"Tool loop limit exceeded for server '{self.server_name}' (max {self.max_tool_rounds} rounds)")
content_blocks = [
_core.ToolUseContent(
type="tool_use", id=tc.id, name=tc.function.name,
input=_parse_tool_call_arguments(self.server_name, tc.function.arguments),
)
for tc in choice.message.tool_calls
]
_core.ToolUseContent(type="tool_use", id=tc.id, name=tc.function.name,
input=_parse_tool_call_arguments(self.server_name, tc.function.arguments))
for tc in choice.message.tool_calls]
self._log_response(response, ", tool_calls=%d", len(content_blocks))
return _core.CreateMessageResultWithTools(
role="assistant", content=content_blocks, model=response.model, stopReason="toolUse",
)
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)."""
"""CreateMessageResult from a normal text response (resets the tool loop)."""
self._tool_loop_count = 0
self._log_response(response)
return _core.CreateMessageResult(
role="assistant",
role="assistant", model=response.model,
content=_core.TextContent(type="text", text=_sanitize_error(choice.message.content or "")),
model=response.model,
stopReason=self._STOP_REASON_MAP.get(choice.finish_reason, "endTurn"),
)
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()),
}
return {"sampling_callback": self,
"sampling_capabilities": _core.SamplingCapability(tools=_core.SamplingToolsCapability())}
def _admit(self, params):
"""Rate-limit + allowed_models gate. Returns ``(resolved_model, None)``
or ``(None, ErrorData)``."""
"""Rate-limit + allowed_models gate. Returns ``(resolved_model, None)`` or ``(None, ErrorData)``."""
if not self._check_rate_limit():
logger.warning("MCP server '%s' sampling rate limit exceeded (%d/min)", self.server_name, self.max_rpm)
return None, self._fail(
f"Sampling rate limit exceeded for server '{self.server_name}' ({self.max_rpm} requests/minute)"
)
f"Sampling rate limit exceeded for server '{self.server_name}' ({self.max_rpm} requests/minute)")
model = self._resolve_model(mcp_field(params, "model_preferences", "modelPreferences"))
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,
)
return None, self._fail(
f"Model '{resolved_model}' not allowed for server "
f"'{self.server_name}'. Allowed: {', '.join(self.allowed_models)}"
)
logger.warning("MCP server '%s' requested model '%s' not in allowed_models", self.server_name, resolved_model)
return None, self._fail(f"Model '{resolved_model}' not allowed for server "
f"'{self.server_name}'. Allowed: {', '.join(self.allowed_models)}")
return resolved_model, None
def _build_llm_call(self, params, resolved_model: str) -> Callable[[], object]:
"""Translate the sampling params into a zero-arg sync ``call_llm`` thunk
(run off-loop so the MCP loop is not blocked)."""
"""Translate the sampling params into a zero-arg sync ``call_llm`` thunk (run off-loop so
the MCP loop is not blocked). Server-provided tools are forwarded."""
from agent.auxiliary_client import call_llm
messages = self._convert_messages(params)
@@ -243,29 +213,15 @@ class SamplingHandler:
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)
temperature = getattr(params, "temperature", None)
# Forward server-provided tools.
server_tools = getattr(params, "tools", None)
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
] if server_tools else None
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),
)
return lambda: call_llm(
task="mcp", model=resolved_model or None, messages=messages, temperature=temperature,
max_tokens=max_tokens, tools=tools, timeout=self.timeout,
)
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] if server_tools else None
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))
return lambda: call_llm(task="mcp", model=resolved_model or None, messages=messages, temperature=temperature,
max_tokens=max_tokens, tools=tools, timeout=self.timeout)
async def __call__(self, context, params):
"""SDK sampling callback (``SamplingFnT``). Returns CreateMessageResult,
@@ -280,11 +236,9 @@ class SamplingHandler:
return self._fail(f"Sampling LLM call timed out after {self.timeout}s for server '{self.server_name}'")
except Exception as exc:
return self._fail(f"Sampling LLM call failed: {_sanitize_error(_exc_str(exc))}")
# Empty choices happen on content filtering / provider errors.
if not getattr(response, "choices", None):
return self._fail(f"LLM returned empty response (no choices) for server '{self.server_name}'")
choice = response.choices[0]
self.metrics["requests"] += 1
total_tokens = _response_total_tokens(response, 0)
@@ -296,12 +250,11 @@ class SamplingHandler:
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."""
"""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():
spec = field_spec if isinstance(field_spec, dict) else {}
@@ -313,25 +266,24 @@ def _format_elicitation_schema_summary(schema: dict, server_name: str) -> str:
class ElicitationHandler:
"""Handles ``elicitation/create`` requests for one MCP server; 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."""
"""Handles ``elicitation/create`` requests for one MCP server; 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.
# 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
# consent answer -> (ElicitResult action, metric); anything else declines.
_ANSWER_RESULTS = {"accept": ("accept", "accepted"), "cancel": ("cancel", "errors")}
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.
# 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.
# 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}
@@ -347,10 +299,9 @@ class ElicitationHandler:
return _core.ElicitResult(action=action)
def _consent_thunk(self, message: str, description: str) -> Callable[[], str]:
"""Sync consent call, replaying the agent's contextvars snapshot when the
owner captured one: the recv-loop task does NOT inherit them, and
gateway-platform detection needs them. ``Context.run`` executes a
context once, so it is copied per elicitation."""
"""Sync consent call, replaying the agent's contextvars snapshot when the owner captured
one: the recv-loop task does NOT inherit them, and gateway-platform detection needs them.
``Context.run`` executes a context once, so it is copied per elicitation."""
from tools.approval import request_elicitation_consent
kwargs = {"timeout_seconds": int(self.timeout), "surface": f"mcp-elicitation/{self.server_name}"}
@@ -362,42 +313,34 @@ class ElicitationHandler:
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.
if getattr(params, "mode", "form") == "url":
logger.info(
"MCP server '%s' requested URL-mode elicitation; declining (URL-mode elicitation not implemented)",
self.server_name,
)
logger.info("MCP server '%s' requested URL-mode elicitation; declining (URL-mode elicitation not implemented)",
self.server_name)
return self._result("decline", "declined")
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.
# ``requestedSchema`` on mcp 1.x, ``requested_schema`` on 2.0 (pydantic aliases don't apply
# to attribute access) — read both or the user approves without seeing the fields.
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:
invoke_consent = self._consent_thunk(message, description)
except Exception as exc: # pragma: no cover -- defensive
logger.error("MCP server '%s' elicitation: approval system unavailable: %s", self.server_name, exc)
return self._result("decline", "errors")
# Offload the sync consent flow to a thread — inline it would freeze the
# MCP loop and every other RPC on this session.
# Offload the sync consent flow to a thread — inline it would freeze the MCP loop and
# every other RPC on this session.
try:
answer = await asyncio.wait_for(
asyncio.to_thread(invoke_consent), timeout=self.timeout + self._OUTER_TIMEOUT_GRACE_SECONDS,
)
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))
return self._result("cancel", "errors")
except Exception as exc:
logger.error("MCP server '%s' elicitation failed: %s", self.server_name, exc, exc_info=True)
return self._result("decline", "errors")
return self._result(*self._ANSWER_RESULTS.get(answer, ("decline", "declined")))

View File

@@ -1,4 +1,6 @@
"""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."""
"""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
@@ -14,12 +16,8 @@ logger = logging.getLogger("tools.mcp_tool")
# JSON-RPC ``initialize`` body used by the content-type preflight POST.
_PROBE_INITIALIZE_BODY = (
'{"jsonrpc":"2.0","id":"_probe",'
'"method":"initialize",'
'"params":{"protocolVersion":"2025-03-26",'
'"capabilities":{},'
'"clientInfo":{"name":"hermes-probe",'
'"version":"0.1"}}}'
'{"jsonrpc":"2.0","id":"_probe","method":"initialize","params":{"protocolVersion":"2025-03-26",'
'"capabilities":{},"clientInfo":{"name":"hermes-probe","version":"0.1"}}}'
)
@@ -133,13 +131,10 @@ class MCPServerTransportMixin:
"(valid: auto, stateless, legacy)", self.name, mode)
# mcp 1.x has no server/discover client — nothing to fall back to.
return await attempt(
initialize, discover,
lambda exc: _handshake_rejected_as_modern(exc) and hasattr(session, "discover"),
"MCP server '%s': legacy handshake rejected (%s) — "
"retrying via server/discover (2026-07-28 stateless server)")
initialize, discover, lambda exc: _handshake_rejected_as_modern(exc) and hasattr(session, "discover"),
"MCP server '%s': legacy handshake rejected (%s) — retrying via server/discover (2026-07-28 stateless server)")
async def _serve_session(self, session, connect_timeout: float,
label: str = "", mark_lifecycle: bool = False) -> str:
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
@@ -187,12 +182,11 @@ class MCPServerTransportMixin:
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.
# 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)
@@ -225,44 +219,43 @@ class MCPServerTransportMixin:
"it is not installed. Run `hermes setup` to install MCP support, then retry.")
command, args, safe_env = self._resolve_stdio_config(config)
await _osv_malware_preflight(self.name, command, args)
# Parent-death watchdog: an ungraceful Hermes exit (kill -9, crash) can't leave the
# child and its descendants running. POSIX-only (process groups); no-op elsewhere.
# AFTER the OSV preflight so the check inspects the real package.
# Parent-death watchdog: an ungraceful Hermes exit (kill -9, crash) can't leave the child
# and its descendants running. 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",
)
encoding_error_handler="replace")
session_kwargs = self._session_kwargs()
# 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
# 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).
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.
# 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.
# (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:
self._track_spawned_children(new_pids)
# 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, **session_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.
# 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:
@@ -302,11 +295,9 @@ class MCPServerTransportMixin:
ct = _content_type_base(resp)
if ct and ct not in self._MCP_CONTENT_TYPES and _is_2xx(resp):
post_resp = await client.post(
url,
url, content=_PROBE_INITIALIZE_BODY,
headers={**probe_headers, "Content-Type": "application/json",
"Accept": "application/json, text/event-stream"},
content=_PROBE_INITIALIZE_BODY,
)
"Accept": "application/json, text/event-stream"})
if _is_2xx(post_resp) and _content_type_base(post_resp) in self._MCP_CONTENT_TYPES:
resp = post_resp
except _httpx.HTTPError:
@@ -323,8 +314,7 @@ class MCPServerTransportMixin:
f"MCP server '{self.name}' at {url} returned Content-Type '{ct_base}', not an MCP "
f"response (expected one of: {', '.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/)."
)
"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
@@ -422,8 +412,8 @@ class MCPServerTransportMixin:
ssl_verify, oauth_auth, strict_cfg_headers: bool):
"""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.
# 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.")
@@ -441,9 +431,9 @@ class MCPServerTransportMixin:
"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.
# 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.
@@ -458,7 +448,6 @@ class MCPServerTransportMixin:
ssl_verify = config.get("ssl_verify", True)
client_cert = _resolve_client_cert(self.name, config)
oauth_auth = self._build_oauth_auth(url, config)
if config.get("transport") == "sse":
transport = self._sse_transport(url, headers, connect_timeout, ssl_verify, client_cert, oauth_auth, strict_cfg_headers)
label = "SSE"
@@ -504,8 +493,8 @@ class MCPServerTransportMixin:
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.
# 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)