Merge branch 'simp/r3-33-F' into simp/r3-33
This commit is contained in:
@@ -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
|
||||
@@ -24,13 +25,9 @@ def _jsonrpc_code(exc: BaseException):
|
||||
|
||||
|
||||
def _jsonrpc_matches(exc: BaseException, code, codes: tuple, markers: tuple) -> bool:
|
||||
"""Structural *code* in *codes*, else any lowercased *marker* in ``str(exc)``. Never
|
||||
``isinstance`` on SDK exception types: they arrive wrapped in ExceptionGroups and drift
|
||||
across generations."""
|
||||
if code in codes:
|
||||
return True
|
||||
msg = str(exc).lower()
|
||||
return bool(msg) and any(marker in msg for marker in markers)
|
||||
"""Structural *code* in *codes*, else any *marker* in ``str(exc).lower()``. Never ``isinstance``
|
||||
on SDK exception types: they arrive wrapped in ExceptionGroups and drift across generations."""
|
||||
return code in codes or any(marker in str(exc).lower() for marker in markers)
|
||||
|
||||
|
||||
def _handshake_rejected_as_modern(exc: BaseException) -> bool:
|
||||
@@ -43,37 +40,32 @@ def _handshake_rejected_as_modern(exc: BaseException) -> bool:
|
||||
|
||||
|
||||
def _is_method_not_found_error(exc: BaseException) -> bool:
|
||||
"""True if *exc* is a JSON-RPC ``method not found`` (-32601). ``ping`` is optional in MCP;
|
||||
servers lacking it answer -32601. The substring fallback includes "Unknown method: <name>",
|
||||
which some servers use; without it the ping→list_tools keepalive fallback never latches
|
||||
and reconnect-loops."""
|
||||
"""True if *exc* is a JSON-RPC ``method not found`` (-32601; ``ping`` is optional in MCP). The
|
||||
substring fallback includes "Unknown method: <name>" — without it the ping→list_tools keepalive
|
||||
fallback never latches and reconnect-loops."""
|
||||
return _jsonrpc_matches(
|
||||
exc, _jsonrpc_code(exc), (_core._JSONRPC_METHOD_NOT_FOUND,),
|
||||
(str(_core._JSONRPC_METHOD_NOT_FOUND), "method not found", "unknown method", "not found: ping"),
|
||||
)
|
||||
(str(_core._JSONRPC_METHOD_NOT_FOUND), "method not found", "unknown method", "not found: ping"))
|
||||
|
||||
|
||||
class InvalidMcpUrlError(ValueError):
|
||||
"""A remote MCP server's ``url`` is not parseable http(s)://. Validated once at startup so
|
||||
we fail fast instead of burning the reconnect-backoff loop on every attempt."""
|
||||
"""A remote MCP server's ``url`` is not parseable http(s):// — validated once at startup so we
|
||||
fail fast instead of burning the reconnect-backoff loop."""
|
||||
|
||||
|
||||
class NonMcpEndpointError(ConnectionError):
|
||||
"""An HTTP MCP URL served a non-MCP 2xx response (e.g. ``text/html``); real Streamable-HTTP
|
||||
endpoints answer ``application/json`` or ``text/event-stream``. Non-retryable: every attempt
|
||||
gets the same page, so the backoff loop is skipped and the server is failed immediately.
|
||||
Subclasses ConnectionError so broad catches still see a connection problem."""
|
||||
"""An HTTP MCP URL served a non-MCP 2xx (e.g. ``text/html``). Non-retryable: every attempt gets
|
||||
the same page, so backoff is skipped and the server fails immediately. Subclasses ConnectionError
|
||||
so broad catches still see a connection problem."""
|
||||
|
||||
|
||||
def _unwrap_exception_group(exc: BaseException) -> BaseException:
|
||||
"""Extract the root-cause leaf from anyio ``(Base)ExceptionGroup`` wrappers (group ``str()``
|
||||
is opaque, so log sites must unwrap). Two rules: a ``KeyboardInterrupt``/``SystemExit`` leaf
|
||||
anywhere is re-raised (never flattened into a loggable error); a non-cancellation leaf is
|
||||
preferred over the ``CancelledError`` noise anyio sprays on siblings."""
|
||||
"""Root-cause leaf of anyio ``(Base)ExceptionGroup`` wrappers (group ``str()`` is opaque). A
|
||||
``KeyboardInterrupt``/``SystemExit`` leaf anywhere is re-raised, never flattened into a loggable
|
||||
error; a non-cancellation leaf is preferred over the ``CancelledError`` noise anyio sprays on siblings."""
|
||||
while isinstance(exc, BaseExceptionGroup) and exc.exceptions:
|
||||
fatal, _rest = exc.split((KeyboardInterrupt, SystemExit))
|
||||
if fatal is not None:
|
||||
leaf: BaseException = fatal
|
||||
leaf: BaseException = exc.split((KeyboardInterrupt, SystemExit))[0]
|
||||
if leaf is not None:
|
||||
while isinstance(leaf, BaseExceptionGroup) and leaf.exceptions:
|
||||
leaf = leaf.exceptions[0]
|
||||
raise leaf
|
||||
@@ -89,29 +81,20 @@ def _contains_only_cancellation(exc: BaseException) -> bool:
|
||||
|
||||
|
||||
def _classify_mcp_failure(exc: BaseException) -> str:
|
||||
"""``'permanent'`` (deterministic — ``run()`` parks immediately instead of burning the retry
|
||||
ladder: auth 401/403, NonMcpEndpointError, InvalidMcpUrlError, missing stdio command
|
||||
FileNotFoundError / ENOENT) or ``'transient'`` (keeps backoff retry)."""
|
||||
"""``'permanent'`` (``run()`` parks instead of burning the retry ladder: auth 401/403,
|
||||
NonMcpEndpointError, InvalidMcpUrlError, missing stdio command) or ``'transient'`` (backoff retry)."""
|
||||
root = _unwrap_exception_group(exc)
|
||||
permanent = (
|
||||
_core._is_auth_error(root)
|
||||
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)
|
||||
)
|
||||
permanent = (_core._is_auth_error(root)
|
||||
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 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``),
|
||||
and empty hosts."""
|
||||
"""The stripped URL if it is a valid http(s) URL; else InvalidMcpUrlError naming the server
|
||||
(non-string, other scheme — stdio servers use ``command`` — or empty host)."""
|
||||
def _bad(detail: str) -> InvalidMcpUrlError:
|
||||
return InvalidMcpUrlError(f"Invalid MCP URL for '{server_name}': {detail}")
|
||||
|
||||
@@ -134,10 +117,9 @@ def _validate_remote_mcp_url(server_name: str, url: Any) -> str:
|
||||
|
||||
|
||||
def _resolve_client_cert(server_name: str, config: dict):
|
||||
"""``client_cert`` / ``client_key`` in httpx's ``cert=`` shape: None when neither is set; a
|
||||
single path for a combined PEM; ``(cert, key)`` or ``(cert, key, password)`` for the
|
||||
pair/list forms. ``~`` is expanded and missing files raise a server-scoped
|
||||
FileNotFoundError instead of an opaque TLS handshake error."""
|
||||
"""``client_cert`` / ``client_key`` in httpx's ``cert=`` shape: None, a combined-PEM path,
|
||||
``(cert, key)`` or ``(cert, key, password)``. ``~`` is expanded; missing files raise a
|
||||
server-scoped FileNotFoundError instead of an opaque TLS handshake error."""
|
||||
raw_cert = config.get("client_cert")
|
||||
raw_key = config.get("client_key")
|
||||
if raw_cert is None and raw_key is None:
|
||||
@@ -152,29 +134,26 @@ def _resolve_client_cert(server_name: str, config: dict):
|
||||
raise FileNotFoundError(f"{prefix}{label} not found at {expanded!r}")
|
||||
return expanded
|
||||
|
||||
if isinstance(raw_cert, (list, tuple)):
|
||||
if raw_key is not None:
|
||||
raise ValueError(f"{prefix}specify either client_cert as a list [cert, key] OR "
|
||||
f"client_cert + client_key, not both")
|
||||
if len(raw_cert) not in (2, 3):
|
||||
raise ValueError(f"{prefix}client_cert list form must have 2 or 3 elements (got {len(raw_cert)})")
|
||||
pair = (_expand(raw_cert[0], "client_cert[0]"), _expand(raw_cert[1], "client_cert[1]"))
|
||||
if len(raw_cert) == 2:
|
||||
return pair
|
||||
if not isinstance(raw_cert[2], str):
|
||||
raise ValueError(f"{prefix}client_cert[2] (key passphrase) must be a string")
|
||||
return (*pair, raw_cert[2])
|
||||
cert_path = _expand(raw_cert, "client_cert")
|
||||
if not isinstance(raw_cert, (list, tuple)):
|
||||
cert_path = _expand(raw_cert, "client_cert")
|
||||
return (cert_path, _expand(raw_key, "client_key")) if raw_key is not None else cert_path # combined PEM
|
||||
if raw_key is not None:
|
||||
return (cert_path, _expand(raw_key, "client_key"))
|
||||
return cert_path # single combined PEM (cert + key)
|
||||
raise ValueError(f"{prefix}specify either client_cert as a list [cert, key] OR "
|
||||
f"client_cert + client_key, not both")
|
||||
if len(raw_cert) not in (2, 3):
|
||||
raise ValueError(f"{prefix}client_cert list form must have 2 or 3 elements (got {len(raw_cert)})")
|
||||
pair = (_expand(raw_cert[0], "client_cert[0]"), _expand(raw_cert[1], "client_cert[1]"))
|
||||
if len(raw_cert) == 2:
|
||||
return pair
|
||||
if not isinstance(raw_cert[2], str):
|
||||
raise ValueError(f"{prefix}client_cert[2] (key passphrase) must be a string")
|
||||
return (*pair, raw_cert[2])
|
||||
|
||||
|
||||
def _resolve_identity_header(server_name: str, config: dict):
|
||||
"""Optional per-server ``identity_header`` ``{name, value_from: "static"|"profile", value}``
|
||||
(``value`` required for static) → ``(name, value)`` or None. Invalid configs warn and are
|
||||
ignored — an identity header must never break the connection. ``profile`` resolves once at
|
||||
connect time; no per-call mutation."""
|
||||
"""``identity_header`` ``{name, value_from: "static"|"profile", value}`` → ``(name, value)`` or
|
||||
None. Invalid configs warn and are ignored — an identity header must never break the connection.
|
||||
``profile`` resolves once at connect time."""
|
||||
raw = config.get("identity_header")
|
||||
if raw is None:
|
||||
return None
|
||||
@@ -189,55 +168,48 @@ def _resolve_identity_header(server_name: str, config: dict):
|
||||
if not isinstance(name, str) or not name.strip():
|
||||
return _ignore("requires a non-empty 'name'")
|
||||
value_from = (raw.get("value_from") or "static").strip().lower()
|
||||
if value_from == "static":
|
||||
value = raw.get("value")
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
return _ignore("with value_from: static requires a non-empty string 'value'")
|
||||
return (name.strip(), value)
|
||||
if value_from == "profile":
|
||||
from hermes_cli.profiles import get_active_profile_name
|
||||
return (name.strip(), get_active_profile_name())
|
||||
return _ignore("value_from must be 'static' or 'profile' (got %r)", value_from)
|
||||
if value_from != "static":
|
||||
return _ignore("value_from must be 'static' or 'profile' (got %r)", value_from)
|
||||
value = raw.get("value")
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
return _ignore("with value_from: static requires a non-empty string 'value'")
|
||||
return (name.strip(), value)
|
||||
|
||||
|
||||
def _apply_identity_header(server_name: str, config: dict, headers: dict) -> dict:
|
||||
"""Merge the resolved identity header into ``headers`` in place. An explicit per-server
|
||||
``headers`` entry with the same name (any casing) wins — the identity header never silently
|
||||
overrides user config."""
|
||||
resolved = _resolve_identity_header(server_name, config)
|
||||
if resolved is None:
|
||||
"""Merge the identity header into ``headers`` in place; an explicit entry of the same name (any
|
||||
casing) wins — never silently override user config."""
|
||||
name, value = _resolve_identity_header(server_name, config) or (None, None)
|
||||
if name is None:
|
||||
return headers
|
||||
name, value = resolved
|
||||
if any(key.lower() == name.lower() for key in headers):
|
||||
logger.debug("MCP server '%s': identity_header '%s' already set via explicit "
|
||||
"headers config — keeping the explicit value", server_name, name)
|
||||
return headers
|
||||
headers[name] = value
|
||||
else:
|
||||
headers[name] = value
|
||||
return headers
|
||||
|
||||
|
||||
def _make_redirect_header_stripper(original_url, *, strict: bool = False,
|
||||
configured_header_names: "set[str] | frozenset[str]" = frozenset()):
|
||||
"""httpx response hook guarding cross-origin redirects: always strips ``Authorization`` when
|
||||
a redirect leaves the original origin; with *strict* (Agent Plugins v1
|
||||
``strict_redirect_headers``) every configured header (lowercase names in
|
||||
*configured_header_names*) is stripped too — the v1 spec forbids forwarding
|
||||
package-configured headers cross-origin."""
|
||||
"""httpx response hook: strips ``Authorization`` when a redirect leaves the original origin;
|
||||
with *strict* (Agent Plugins v1 ``strict_redirect_headers``) every configured header (lowercase
|
||||
names in *configured_header_names*) is stripped too — v1 forbids forwarding them cross-origin."""
|
||||
origin = (original_url.scheme, original_url.host, original_url.port)
|
||||
|
||||
async def _strip_on_cross_origin_redirect(response):
|
||||
if not (response.is_redirect and response.next_request):
|
||||
return
|
||||
target = response.next_request.url
|
||||
if (target.scheme, target.host, target.port) == origin:
|
||||
target = response.next_request.url if response.is_redirect and response.next_request else None
|
||||
if target is None or (target.scheme, target.host, target.port) == origin:
|
||||
return
|
||||
headers = response.next_request.headers
|
||||
headers.pop("authorization", None)
|
||||
headers.pop("Authorization", None)
|
||||
if strict:
|
||||
for _name in configured_header_names:
|
||||
while _name in headers:
|
||||
del headers[_name]
|
||||
for _name in configured_header_names if strict else ():
|
||||
while _name in headers:
|
||||
del headers[_name]
|
||||
|
||||
return _strip_on_cross_origin_redirect
|
||||
|
||||
@@ -245,14 +217,11 @@ def _make_redirect_header_stripper(original_url, *, strict: bool = False,
|
||||
def _exc_children(exc: BaseException) -> List[BaseException]:
|
||||
"""Sub-exceptions of a group, else ``__cause__``/``__context__`` when they are exceptions."""
|
||||
nested = getattr(exc, "exceptions", None)
|
||||
if nested:
|
||||
return list(nested)
|
||||
return [c for c in (exc.__cause__, exc.__context__) if isinstance(c, BaseException)]
|
||||
return list(nested) if nested else [c for c in (exc.__cause__, exc.__context__) if isinstance(c, 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):
|
||||
@@ -265,24 +234,21 @@ def _format_connect_error(exc: BaseException) -> str:
|
||||
def _flatten_messages(current: BaseException) -> List[str]:
|
||||
# A group's own str() is opaque — only its children speak.
|
||||
text = "" if getattr(current, "exceptions", None) else str(current).strip()
|
||||
messages = [text] if text else []
|
||||
for child in _exc_children(current):
|
||||
messages.extend(_flatten_messages(child))
|
||||
messages = ([text] if text else []) + [m for child in _exc_children(current) for m in _flatten_messages(child)]
|
||||
return messages or [current.__class__.__name__]
|
||||
|
||||
missing = _find_missing(exc)
|
||||
if missing:
|
||||
message = f"missing executable '{missing}'"
|
||||
if os.path.basename(missing) in {"npx", "npm", "node"}:
|
||||
message += (" (ensure Node.js is installed and PATH includes its bin directory, "
|
||||
"or set mcp_servers.<name>.command to an absolute path and include "
|
||||
"that directory in mcp_servers.<name>.env.PATH)")
|
||||
return _sanitize_error(message)
|
||||
deduped = list(dict.fromkeys(_flatten_messages(exc)))
|
||||
return _sanitize_error("; ".join(deduped[:3]))
|
||||
if not missing:
|
||||
return _sanitize_error("; ".join(list(dict.fromkeys(_flatten_messages(exc)))[:3]))
|
||||
message = f"missing executable '{missing}'"
|
||||
if os.path.basename(missing) in {"npx", "npm", "node"}:
|
||||
message += (" (ensure Node.js is installed and PATH includes its bin directory, "
|
||||
"or set mcp_servers.<name>.command to an absolute path and include "
|
||||
"that directory in mcp_servers.<name>.env.PATH)")
|
||||
return _sanitize_error(message)
|
||||
|
||||
|
||||
# Lazily-built caches so this module imports even without the SDK OAuth module.
|
||||
# Lazily-built caches so this module imports without the SDK OAuth module.
|
||||
_AUTH_ERROR_TYPES: tuple = ()
|
||||
_HTTP_STATUS_ERROR_TYPES: Optional[tuple] = None
|
||||
|
||||
@@ -297,74 +263,59 @@ def _optional_types(module: str, *names: str) -> list:
|
||||
|
||||
|
||||
def _http_status_error_types() -> tuple:
|
||||
"""``HTTPStatusError`` classes from both httpx flavours: a 401 may come from the SDK's own
|
||||
stack (``httpx2`` on mcp >= 2.0) or from Hermes' pinned ``httpx``; the classes are unrelated."""
|
||||
"""``HTTPStatusError`` from both httpx flavours: a 401 may come from the SDK's own stack
|
||||
(``httpx2`` on mcp >= 2.0) or Hermes' pinned ``httpx``; the classes are unrelated."""
|
||||
global _HTTP_STATUS_ERROR_TYPES
|
||||
if _HTTP_STATUS_ERROR_TYPES is None:
|
||||
sdk_mod = _core.sdk_httpx()
|
||||
found: list = [sdk_mod.HTTPStatusError] if sdk_mod is not None else []
|
||||
found += [cls for cls in _optional_types("httpx", "HTTPStatusError") if cls not in found]
|
||||
_HTTP_STATUS_ERROR_TYPES = tuple(found)
|
||||
_HTTP_STATUS_ERROR_TYPES = tuple(dict.fromkeys(
|
||||
([sdk_mod.HTTPStatusError] if sdk_mod is not None else []) + _optional_types("httpx", "HTTPStatusError")))
|
||||
return _HTTP_STATUS_ERROR_TYPES
|
||||
|
||||
|
||||
def _get_auth_error_types() -> tuple:
|
||||
"""Cached exception types indicating MCP OAuth failure: SDK ``OAuthFlowError``/``OAuthTokenError``
|
||||
(+ legacy ``UnauthorizedError``), our ``OAuthNonInteractiveError``, and ``HTTPStatusError`` from
|
||||
both httpx flavours — the latter needs the 401 status check in :func:`_is_auth_error`."""
|
||||
"""Cached MCP OAuth failure types: SDK ``OAuthFlowError``/``OAuthTokenError`` (+ legacy
|
||||
``UnauthorizedError``), our ``OAuthNonInteractiveError``, and both ``HTTPStatusError`` flavours
|
||||
(which still need the 401 check in :func:`_is_auth_error`)."""
|
||||
global _AUTH_ERROR_TYPES
|
||||
if not _AUTH_ERROR_TYPES:
|
||||
_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())
|
||||
)
|
||||
_AUTH_ERROR_TYPES = (*_optional_types("mcp.client.auth", "OAuthFlowError", "OAuthTokenError"),
|
||||
*_optional_types("mcp.client.auth", "UnauthorizedError"), # older SDKs
|
||||
*_optional_types("tools.mcp_oauth", "OAuthNonInteractiveError"),
|
||||
*_http_status_error_types())
|
||||
return _AUTH_ERROR_TYPES
|
||||
|
||||
|
||||
def _is_auth_error(exc: BaseException) -> bool:
|
||||
"""True if ``exc`` indicates an MCP OAuth failure. ``HTTPStatusError`` counts only with
|
||||
status 401; other HTTP errors fall through to the generic error path."""
|
||||
types = _get_auth_error_types()
|
||||
if not types or not isinstance(exc, types):
|
||||
"""True if ``exc`` indicates an MCP OAuth failure; ``HTTPStatusError`` counts only with status 401."""
|
||||
if not isinstance(exc, _get_auth_error_types()):
|
||||
return False
|
||||
status_error_types = _http_status_error_types()
|
||||
if status_error_types and isinstance(exc, status_error_types):
|
||||
return getattr(exc.response, "status_code", None) == 401
|
||||
return True
|
||||
return getattr(exc.response, "status_code", None) == 401 if isinstance(exc, _http_status_error_types()) else True
|
||||
|
||||
|
||||
# Lower-cased substrings meaning the server-side transport session expired / was GC'd.
|
||||
# The OAuth token is still valid — only the transport needs rebuilding.
|
||||
# Lower-cased substrings meaning the transport session expired / was GC'd (OAuth token still valid).
|
||||
_SESSION_EXPIRED_MARKERS: tuple = (
|
||||
"invalid or expired session", "expired session", "session expired", "session not found",
|
||||
"unknown session", "session terminated", "closedresourceerror", "closed resource",
|
||||
"transport is closed", "connection closed", "broken pipe", "end of file",
|
||||
)
|
||||
"transport is closed", "connection closed", "broken pipe", "end of file")
|
||||
|
||||
# Node budget for ``_is_session_expired_error``. The visited set breaks cycles; the budget
|
||||
# bounds pathological acyclic graphs. Kept well above ``sys.getrecursionlimit()`` so deep
|
||||
# task-group nesting is still fully scanned.
|
||||
# Node budget for ``_is_session_expired_error`` (the visited set breaks cycles; this bounds acyclic
|
||||
# blow-ups). Well above ``sys.getrecursionlimit()`` so deep task-group nesting is fully scanned.
|
||||
_EXC_TRAVERSAL_MAX_NODES = 10_000
|
||||
|
||||
|
||||
def _is_session_expired_error(exc: BaseException) -> bool:
|
||||
"""True if ``exc`` looks like an MCP transport session expiry. Streamable-HTTP servers GC
|
||||
session state (idle TTL, restart, pod rotation) while the OAuth token stays valid, so unlike
|
||||
:func:`_is_auth_error` the fix is a transport reconnect (``_reconnect_event``), not an OAuth refresh."""
|
||||
# AnyIO stream exceptions are often message-less (``str(ClosedResourceError()) == ""``),
|
||||
# so type checks are needed in addition to marker matching.
|
||||
"""True if ``exc`` looks like a transport session expiry (Streamable-HTTP servers GC session
|
||||
state on idle TTL / restart / pod rotation while the OAuth token stays valid) — the fix is a
|
||||
transport reconnect, not an OAuth refresh. Iterative walk over ``exceptions`` / ``__cause__`` /
|
||||
``__context__`` with a visited set AND a node budget; every reachable node is inspected so an
|
||||
InterruptedError anywhere overrides transport markers, and the chain walk matters because SDK
|
||||
wrappers raise a generic RuntimeError *from* a message-less ClosedResourceError."""
|
||||
# AnyIO stream exceptions are often message-less, so type checks complement 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
|
||||
found = False
|
||||
budget = _EXC_TRAVERSAL_MAX_NODES
|
||||
while stack and budget > 0:
|
||||
current = stack.pop()
|
||||
@@ -374,11 +325,10 @@ def _is_session_expired_error(exc: BaseException) -> bool:
|
||||
budget -= 1
|
||||
if isinstance(current, InterruptedError):
|
||||
return False
|
||||
# Messages vary across SDK versions and servers: match a narrow allow-list of stable
|
||||
# substrings, not exception type, to avoid false positives.
|
||||
# Messages vary across SDK versions/servers: a narrow allow-list of stable substrings avoids
|
||||
# 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)):
|
||||
transport_error_found = True
|
||||
found = found or isinstance(current, transport_error_types) or any(m in msg for m in _SESSION_EXPIRED_MARKERS)
|
||||
stack.extend(getattr(current, "exceptions", ()))
|
||||
stack.extend((getattr(current, "__cause__", None), getattr(current, "__context__", None)))
|
||||
return transport_error_found
|
||||
return found
|
||||
|
||||
@@ -1,13 +1,13 @@
|
||||
"""Background-loop plumbing for tools.mcp_tool: the cross-process discovery file lock,
|
||||
scheduling coroutines onto the MCP loop from caller threads (with profile HOME override
|
||||
and dashboard OAuth flow propagation) and the loop's exception handler. Split from
|
||||
tools/mcp_tool.py; origin state (``_lock``, ``_mcp_loop``) is read through ``_core`` so
|
||||
``mock.patch("tools.mcp_tool.X")`` keeps working."""
|
||||
"""Background-loop plumbing for tools.mcp_tool: cross-process discovery file lock, scheduling
|
||||
coroutines onto the MCP loop from caller threads (with profile HOME override and dashboard OAuth
|
||||
flow propagation) and the loop's exception handler. Origin state (``_lock``, ``_mcp_loop``) is
|
||||
read through ``_core`` so ``mock.patch("tools.mcp_tool.X")`` keeps working."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import concurrent.futures
|
||||
import contextlib
|
||||
import errno
|
||||
import logging
|
||||
import os
|
||||
@@ -20,11 +20,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
|
||||
held: both fcntl and portalocker locks are tied to the descriptor."""
|
||||
|
||||
def __init__(self, fh: Any) -> None:
|
||||
self._fh = fh
|
||||
@@ -32,76 +29,56 @@ 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: 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)
|
||||
return True
|
||||
fcntl.flock(fh.fileno(), fcntl.LOCK_EX | fcntl.LOCK_NB)
|
||||
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
|
||||
return True
|
||||
import portalocker
|
||||
try:
|
||||
portalocker.lock(fh, portalocker.LOCK_EX | portalocker.LOCK_NB)
|
||||
except portalocker.LockException:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _try_acquire_mcp_discovery_lock() -> Any:
|
||||
"""Return a ``_LockCookie`` (acquired), ``None`` (held by another process)
|
||||
or ``_LOCK_UNAVAILABLE`` (locking broken: run discovery unguarded)."""
|
||||
# The cached path lives on the ORIGIN module (tests reset
|
||||
# ``tools.mcp_tool._MCP_DISCOVERY_LOCK_PATH = None``), so write it there.
|
||||
"""``_LockCookie`` (acquired), ``None`` (held by another process) or ``_LOCK_UNAVAILABLE``
|
||||
(locking broken: run discovery unguarded)."""
|
||||
# The cached path lives on the ORIGIN module (tests reset ``tools.mcp_tool._MCP_DISCOVERY_LOCK_PATH``).
|
||||
from tools import mcp_tool as _origin
|
||||
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()
|
||||
@@ -109,27 +86,21 @@ def _try_acquire_mcp_discovery_lock() -> Any:
|
||||
|
||||
|
||||
def _mcp_loop_exception_handler(loop, context):
|
||||
"""Suppress the benign 'Event loop is closed' RuntimeError that httpx
|
||||
finalizers raise against the dead loop during shutdown; forward the rest."""
|
||||
"""Suppress the benign 'Event loop is closed' RuntimeError httpx finalizers raise against the
|
||||
dead loop during shutdown; forward the rest."""
|
||||
exc = context.get("exception")
|
||||
if isinstance(exc, RuntimeError) and "Event loop is closed" in str(exc):
|
||||
return
|
||||
loop.default_exception_handler(context)
|
||||
if not (isinstance(exc, RuntimeError) and "Event loop is closed" in str(exc)):
|
||||
loop.default_exception_handler(context)
|
||||
|
||||
|
||||
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)."""
|
||||
"""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
|
||||
home_override = None
|
||||
if not home_override:
|
||||
return coro
|
||||
|
||||
@@ -146,14 +117,10 @@ 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
|
||||
flow = None
|
||||
if flow is None:
|
||||
return coro
|
||||
|
||||
@@ -172,12 +139,8 @@ 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 (or zero-arg factory — avoids leaking a never-awaited coroutine when the
|
||||
loop is down) on the MCP loop and block until done, polling so user interrupts are honored."""
|
||||
from tools.interrupt import is_interrupted
|
||||
from agent.async_utils import safe_schedule_threadsafe
|
||||
|
||||
@@ -186,58 +149,37 @@ 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.
|
||||
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",
|
||||
)
|
||||
# run_coroutine_threadsafe copies 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_dashboard_oauth_flow(_core._wrap_with_home_override(
|
||||
coro_or_factory() if callable(coro_or_factory) else coro_or_factory))
|
||||
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)"
|
||||
)
|
||||
wait_timeout = min(wait_timeout, remaining)
|
||||
|
||||
remaining = 0.1 if deadline is None else deadline - time.monotonic()
|
||||
if remaining <= 0:
|
||||
future.cancel()
|
||||
raise TimeoutError(f"MCP call timed out after {time.monotonic() - start_time:.1f}s "
|
||||
f"(configured timeout: {float(timeout):.1f}s)")
|
||||
try:
|
||||
return future.result(timeout=wait_timeout)
|
||||
return future.result(timeout=min(0.1, remaining))
|
||||
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 it 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: the event lives on the MCP loop,
|
||||
so set via ``call_soon_threadsafe`` when it runs (direct ``.set()`` otherwise). False when the
|
||||
server has no reconnect machinery."""
|
||||
event = getattr(server, "_reconnect_event", None)
|
||||
if event is None:
|
||||
return False
|
||||
@@ -253,97 +195,58 @@ def reconnect_mcp_server(server_name: str) -> bool:
|
||||
"""Ask a currently-live MCP server to rebuild after external re-auth."""
|
||||
with _core._lock:
|
||||
server = _core._servers.get(server_name)
|
||||
if server is None:
|
||||
return False
|
||||
return _core._signal_reconnect(server)
|
||||
return server is not None and _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``.
|
||||
"""
|
||||
poll_interval = 0.25
|
||||
iterations = max(1, int(max(float(timeout), 0.0) / poll_interval))
|
||||
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 stale; retrying blindly burns breaker strikes). With ``old_session`` the observed
|
||||
session must differ. Iteration-bounded, not deadline-bounded: tests freeze ``time.monotonic``."""
|
||||
iterations = max(1, int(max(float(timeout), 0.0) / 0.25))
|
||||
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:
|
||||
time.sleep(poll_interval)
|
||||
time.sleep(0.25)
|
||||
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, else the readiness poll returns at once on the dead session."""
|
||||
loop = _core._mcp_loop
|
||||
if loop is None or not loop.is_running():
|
||||
return False
|
||||
old_session = getattr(srv, "session", None)
|
||||
|
||||
def _request_reconnect() -> None:
|
||||
ready = getattr(srv, "_ready", None)
|
||||
if ready is not None and hasattr(ready, "clear"):
|
||||
ready, reconnect_event = getattr(srv, "_ready", None), getattr(srv, "_reconnect_event", None)
|
||||
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,
|
||||
)
|
||||
old_session = getattr(srv, "session", None)
|
||||
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 loop thread if not running. The loop/thread handles live on the ORIGIN
|
||||
module (tests read and reset ``tools.mcp_tool._mcp_loop``), so they are written there."""
|
||||
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,
|
||||
)
|
||||
loop = _origin._mcp_loop = asyncio.new_event_loop()
|
||||
loop.set_exception_handler(_core._mcp_loop_exception_handler)
|
||||
_origin._mcp_thread = threading.Thread(target=loop.run_forever, name="mcp-event-loop", daemon=True)
|
||||
_origin._mcp_thread.start()
|
||||
|
||||
|
||||
@@ -354,53 +257,41 @@ def _stop_mcp_loop(*, only_if_idle: bool = False) -> bool:
|
||||
if only_if_idle and (_core._servers or _core._server_connecting):
|
||||
logger.debug("Leaving MCP event loop running; active servers are registered or connecting")
|
||||
return False
|
||||
loop = _origin._mcp_loop
|
||||
thread = _origin._mcp_thread
|
||||
_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.
|
||||
stop_owned_by_loop = False
|
||||
if loop.is_running():
|
||||
from agent.async_utils import safe_schedule_threadsafe
|
||||
loop, thread = _origin._mcp_loop, _origin._mcp_thread
|
||||
_origin._mcp_loop = _origin._mcp_thread = None
|
||||
if loop is None:
|
||||
return True
|
||||
# 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; everything else ends here.
|
||||
future = None
|
||||
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,
|
||||
)
|
||||
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,
|
||||
)
|
||||
except BaseException as exc:
|
||||
logger.warning("Error draining MCP loop tasks: %s", exc)
|
||||
elif not loop.is_closed():
|
||||
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)
|
||||
if future is not None:
|
||||
try:
|
||||
loop.run_until_complete(
|
||||
_core._drain_mcp_loop_tasks(timeout=_core._MCP_LOOP_DRAIN_TIMEOUT)
|
||||
)
|
||||
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)
|
||||
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:
|
||||
thread.join(timeout=5)
|
||||
if thread.is_alive():
|
||||
logger.warning("MCP event loop thread did not stop within 5.0s")
|
||||
logger.warning("Error draining MCP loop tasks: %s", exc)
|
||||
elif not loop.is_closed():
|
||||
try:
|
||||
loop.close()
|
||||
except Exception as exc:
|
||||
logger.warning("Unable to close MCP event loop cleanly: %s", exc)
|
||||
# The loop is gone, so no session can be in flight: reap active too.
|
||||
_core._kill_orphaned_mcp_children(include_active=True)
|
||||
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 future is None and loop.is_running(): # drain-and-stop wasn't scheduled: stop it ourselves
|
||||
loop.call_soon_threadsafe(loop.stop)
|
||||
if thread is not None:
|
||||
thread.join(timeout=5)
|
||||
if thread.is_alive():
|
||||
logger.warning("MCP event loop thread did not stop within 5.0s")
|
||||
try:
|
||||
loop.close()
|
||||
except Exception as exc:
|
||||
logger.warning("Unable to close MCP event loop cleanly: %s", exc)
|
||||
# The loop is gone, so no session can be in flight: reap active too.
|
||||
_core._kill_orphaned_mcp_children(include_active=True)
|
||||
return True
|
||||
|
||||
@@ -13,16 +13,11 @@ 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 (marks a tool *result* block) 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)
|
||||
|
||||
|
||||
def _is_tool_use(block) -> bool:
|
||||
return hasattr(block, "name") and hasattr(block, "input")
|
||||
|
||||
|
||||
def _tool_result_text(block) -> str:
|
||||
"""Text of a ToolResultContent block ("" when it carries no content)."""
|
||||
content = getattr(block, "content", None)
|
||||
@@ -45,26 +40,19 @@ 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)."""
|
||||
"""One MCP SamplingMessage -> OpenAI 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)]
|
||||
|
||||
others = [b for b in blocks if _tool_use_id(b) is _MISSING]
|
||||
tool_uses = [b for b in others if hasattr(b, "name") and hasattr(b, "input")]
|
||||
content_blocks = [b for b in others if not (hasattr(b, "name") and hasattr(b, "input"))]
|
||||
out = [{"role": "tool", "tool_call_id": _tool_use_id(tr), "content": _tool_result_text(tr)} for tr in tool_results]
|
||||
if tool_uses:
|
||||
msg_dict: dict = {"role": msg.role, "tool_calls": [_tool_call_dict(tu, i) for i, tu in enumerate(tool_uses)]}
|
||||
@@ -72,45 +60,32 @@ def _convert_sampling_message(msg) -> List[dict]:
|
||||
if text_parts:
|
||||
msg_dict["content"] = "\n".join(text_parts)
|
||||
out.append(msg_dict)
|
||||
elif len(content_blocks) == 1 and hasattr(content_blocks[0], "text"):
|
||||
out.append({"role": msg.role, "content": content_blocks[0].text})
|
||||
elif content_blocks:
|
||||
if len(content_blocks) == 1 and hasattr(content_blocks[0], "text"):
|
||||
out.append({"role": msg.role, "content": content_blocks[0].text})
|
||||
else:
|
||||
parts = [p for p in map(_content_part, content_blocks) if p is not None]
|
||||
if parts:
|
||||
out.append({"role": msg.role, "content": parts})
|
||||
parts = [p for p in map(_content_part, content_blocks) if p is not None]
|
||||
if parts:
|
||||
out.append({"role": msg.role, "content": parts})
|
||||
return out
|
||||
|
||||
|
||||
def _parse_tool_call_arguments(server_name: str, args) -> dict:
|
||||
"""LLM tool_calls arguments -> dict; malformed JSON / non-dict values are
|
||||
wrapped as ``{"_raw": ...}`` rather than dropped."""
|
||||
"""LLM tool_calls arguments -> dict; malformed JSON / non-dicts become ``{"_raw": ...}``, not dropped."""
|
||||
if isinstance(args, str):
|
||||
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)}
|
||||
|
||||
|
||||
def _response_total_tokens(response, default):
|
||||
return getattr(getattr(response, "usage", None), "total_tokens", 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``.
|
||||
|
||||
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.
|
||||
"""
|
||||
"""``sampling_callback`` for one MCP server (per-instance rate-limit, metrics, tool-loop state).
|
||||
Runs on the MCP loop; the sync LLM call is offloaded via ``asyncio.to_thread``. Deprecated
|
||||
upstream (MCP 2026-07-28, SEP-2577): stays functional for handshake-era servers, 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}
|
||||
@@ -141,22 +116,19 @@ class SamplingHandler:
|
||||
"""Config override > server hint > None (use default)."""
|
||||
if self.model_override:
|
||||
return self.model_override
|
||||
for hint in (getattr(preferences, "hints", None) or []):
|
||||
if getattr(hint, "name", None):
|
||||
return hint.name
|
||||
return None
|
||||
hints = getattr(preferences, "hints", None) or []
|
||||
return next((hint.name for hint in hints if getattr(hint, "name", None)), 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
|
||||
def _error(message: str, code: int = -1):
|
||||
"""Return ErrorData (MCP spec) or raise as fallback."""
|
||||
if _core._MCP_SAMPLING_TYPES:
|
||||
return _core.ErrorData(code=code, message=message)
|
||||
raise Exception(message)
|
||||
if not _core._MCP_SAMPLING_TYPES:
|
||||
raise Exception(message)
|
||||
return _core.ErrorData(code=code, message=message)
|
||||
|
||||
def _fail(self, message: str):
|
||||
"""Count an error and return the ErrorData for it."""
|
||||
@@ -164,77 +136,55 @@ 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, getattr(getattr(response, "usage", None), "total_tokens", "?"), *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 a tool_calls response, under ``max_tool_rounds`` (0 disables)."""
|
||||
self.metrics["tool_use_count"] += 1
|
||||
if self.max_tool_rounds == 0:
|
||||
self._tool_loop_count = 0
|
||||
return self._error(f"Tool loops disabled for server '{self.server_name}' (max_tool_rounds=0)")
|
||||
self._tool_loop_count += 1
|
||||
if self._tool_loop_count > self.max_tool_rounds:
|
||||
if self.max_tool_rounds == 0 or 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)"
|
||||
)
|
||||
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
|
||||
]
|
||||
f"Tool loops disabled for server '{self.server_name}' (max_tool_rounds=0)" if self.max_tool_rounds == 0
|
||||
else 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]
|
||||
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)"
|
||||
)
|
||||
model = self._resolve_model(mcp_field(params, "model_preferences", "modelPreferences"))
|
||||
resolved_model = model or self.model_override or ""
|
||||
f"Sampling rate limit exceeded for server '{self.server_name}' ({self.max_rpm} requests/minute)")
|
||||
resolved_model = self._resolve_model(mcp_field(params, "model_preferences", "modelPreferences")) 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)."""
|
||||
"""Sampling params -> zero-arg sync ``call_llm`` thunk (run off-loop); server tools are forwarded."""
|
||||
from agent.auxiliary_client import call_llm
|
||||
|
||||
messages = self._convert_messages(params)
|
||||
@@ -242,96 +192,69 @@ class SamplingHandler:
|
||||
if system_prompt:
|
||||
messages.insert(0, {"role": "system", "content": system_prompt})
|
||||
max_tokens = min(mcp_field(params, "max_tokens", "maxTokens", self.max_tokens_cap), self.max_tokens_cap)
|
||||
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, max_tokens=max_tokens,
|
||||
temperature=getattr(params, "temperature", None), tools=tools, timeout=self.timeout)
|
||||
|
||||
async def __call__(self, context, params):
|
||||
"""SDK sampling callback (``SamplingFnT``). Returns CreateMessageResult,
|
||||
CreateMessageResultWithTools, or ErrorData."""
|
||||
"""SDK ``SamplingFnT``: CreateMessageResult, CreateMessageResultWithTools, or ErrorData."""
|
||||
resolved_model, err = self._admit(params)
|
||||
if err is not None:
|
||||
return err
|
||||
sync_call = self._build_llm_call(params, resolved_model)
|
||||
sync_call = self._build_llm_call(params, resolved_model) # outside the try: its errors propagate, not _fail
|
||||
try:
|
||||
response = await asyncio.wait_for(asyncio.to_thread(sync_call), timeout=self.timeout)
|
||||
except asyncio.TimeoutError:
|
||||
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)
|
||||
if isinstance(total_tokens, int):
|
||||
self.metrics["tokens_used"] += total_tokens
|
||||
total_tokens = getattr(getattr(response, "usage", None), "total_tokens", 0)
|
||||
self.metrics["tokens_used"] += total_tokens if isinstance(total_tokens, int) else 0
|
||||
if choice.finish_reason == "tool_calls" and getattr(choice.message, "tool_calls", None):
|
||||
return self._build_tool_use_result(choice, response)
|
||||
return self._build_text_result(choice, response)
|
||||
|
||||
|
||||
def _format_elicitation_schema_summary(schema: dict, server_name: str) -> str:
|
||||
"""Render a flat-object requested_schema as a human-readable field list
|
||||
(names, types, descriptions) so the user knows what they're approving."""
|
||||
"""Flat-object requested_schema -> readable field list 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 {}
|
||||
field_type = str(spec.get("type", "") or "")
|
||||
field_desc = str(spec.get("description", "") or "")
|
||||
suffix = f" ({field_type})" if field_type else ""
|
||||
lines.append(f" - {field_name}{suffix}: {field_desc}" if field_desc else f" - {field_name}{suffix}")
|
||||
field_type, field_desc = str(spec.get("type", "") or ""), str(spec.get("description", "") or "")
|
||||
lines.append(f" - {field_name}" + (f" ({field_type})" if field_type else "") + (f": {field_desc}" if field_desc else ""))
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
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
|
||||
"""``elicitation_callback`` for one MCP server. Form-mode routes through Hermes' approval system
|
||||
(CLI, TUI, Telegram, ...); URL-mode is declined. 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.
|
||||
# 5 min mirrors the gateway approval default so async surfaces (Telegram, Slack) can 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 for isolated unit tests.
|
||||
self.owner = owner
|
||||
self.metrics = {"requests": 0, "accepted": 0, "declined": 0, "errors": 0}
|
||||
|
||||
@@ -342,15 +265,12 @@ class ElicitationHandler:
|
||||
def _result(self, action: str, metric: str):
|
||||
"""Count *metric* and return ``ElicitResult(action)`` (accept carries empty content)."""
|
||||
self.metrics[metric] += 1
|
||||
if action == "accept":
|
||||
return _core.ElicitResult(action="accept", content={})
|
||||
return _core.ElicitResult(action=action)
|
||||
return _core.ElicitResult(action=action, **({"content": {}} if action == "accept" else {}))
|
||||
|
||||
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; gateway-platform detection needs them).
|
||||
``Context.run`` runs 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 +282,28 @@ 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,
|
||||
)
|
||||
if getattr(params, "mode", "form") == "url": # OAuth/payment: needs a browser + elicitation/complete; unsupported
|
||||
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 (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)
|
||||
try: # lazy import inside avoids import-order coupling with early-bootstrap tools.approval
|
||||
invoke_consent = self._consent_thunk(message, _format_elicitation_schema_summary(schema, self.server_name))
|
||||
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.
|
||||
try:
|
||||
try: # off-thread: inline, the sync consent flow would freeze the MCP loop and every RPC on it
|
||||
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")))
|
||||
|
||||
@@ -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,13 +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"}}}')
|
||||
|
||||
|
||||
def _content_type_base(resp) -> str:
|
||||
@@ -32,18 +29,6 @@ def _is_2xx(resp) -> bool:
|
||||
return 200 <= resp.status_code < 300
|
||||
|
||||
|
||||
def _capture_pgids(pids: Set[int]) -> Dict[int, int]:
|
||||
"""pgid per live pid. Captured while the child is alive — getpgid fails once it exits,
|
||||
and the sweep needs it to reach reparented descendants."""
|
||||
pgids: Dict[int, int] = {}
|
||||
for pid in pids:
|
||||
try:
|
||||
pgids[pid] = os.getpgid(pid)
|
||||
except (AttributeError, ProcessLookupError, OSError): # Windows / already exited
|
||||
pass
|
||||
return pgids
|
||||
|
||||
|
||||
def _pgroup_alive(pgid: Optional[int]) -> bool:
|
||||
"""Signal 0 to the group succeeds iff any member is alive (POSIX only)."""
|
||||
_killpg = getattr(os, "killpg", None)
|
||||
@@ -57,9 +42,8 @@ def _pgroup_alive(pgid: Optional[int]) -> bool:
|
||||
|
||||
|
||||
async def _osv_malware_preflight(server_name: str, command: str, args: list) -> None:
|
||||
"""OSV malware preflight: off-loop (blocking HTTPS) with a wall-clock bound so a stalled
|
||||
handshake can't freeze discovery; fail-open on timeout. Must run against the REAL
|
||||
command/args — the watchdog wrap rewrites argv to the supervisor, turning the check into a no-op."""
|
||||
"""OSV malware preflight, off-loop with a wall-clock bound (fail-open on timeout). Must run on
|
||||
the REAL command/args — the watchdog wrap rewrites argv to the supervisor (check becomes a no-op)."""
|
||||
from tools.osv_check import check_package_for_malware
|
||||
try:
|
||||
malware_error = await asyncio.wait_for(
|
||||
@@ -79,9 +63,8 @@ class MCPServerTransportMixin:
|
||||
__slots__ = ()
|
||||
|
||||
def _advertises_tools(self) -> bool:
|
||||
"""Whether the server advertises the ``tools`` capability. Prompt-/resource-only servers
|
||||
omit it, and ``tools/list`` against them raises ``MCPError(-32601)``. True when no
|
||||
capability info was captured (legacy fallback: always call list_tools)."""
|
||||
"""Whether the server advertises ``tools`` (prompt-/resource-only servers omit it and
|
||||
``tools/list`` raises -32601). True when no capability info was captured (legacy fallback)."""
|
||||
caps = getattr(self.initialize_result, "capabilities", None)
|
||||
return caps is None or getattr(caps, "tools", None) is not None
|
||||
|
||||
@@ -97,53 +80,45 @@ class MCPServerTransportMixin:
|
||||
return kwargs
|
||||
|
||||
async def _negotiate_session(self, session, connect_timeout: float):
|
||||
"""Negotiate the protocol era (``initialize`` vs ``server/discover``) and return its result.
|
||||
Per-server ``protocol`` key: ``auto`` (default) tries the legacy handshake FIRST and falls
|
||||
back to ``server/discover`` only when the server signals modern-only (-32022 / initialize
|
||||
-32601) — the reverse of the SDK's discover-first mode, on purpose: zero extra round-trips
|
||||
for the handshake-era servers that dominate today. ``stateless`` probes discover first (one
|
||||
legacy retry on any error); ``legacy`` is handshake only, no fallback. Both result types
|
||||
expose ``.capabilities``. A handshake TIMEOUT never triggers a fallback — it propagates."""
|
||||
def initialize():
|
||||
return asyncio.wait_for(session.initialize(), timeout=connect_timeout)
|
||||
|
||||
def discover():
|
||||
return asyncio.wait_for(session.discover(), timeout=connect_timeout)
|
||||
"""Negotiate the protocol era (``initialize`` vs ``server/discover``); both results expose
|
||||
``.capabilities``. ``protocol: auto`` (default) tries the legacy handshake FIRST and falls back
|
||||
to discover only on a modern-only signal (-32022 / initialize -32601) — deliberately the
|
||||
reverse of the SDK's discover-first mode: zero extra round-trips for the handshake-era
|
||||
servers that dominate today. ``stateless`` probes discover first (one legacy retry on any
|
||||
error); ``legacy`` is handshake only. A handshake TIMEOUT never falls back — it propagates."""
|
||||
def call(method: str):
|
||||
return asyncio.wait_for(getattr(session, method)(), timeout=connect_timeout)
|
||||
|
||||
async def attempt(primary, fallback, should_fallback, log_fmt, *log_extra):
|
||||
try:
|
||||
return await primary()
|
||||
except asyncio.TimeoutError:
|
||||
raise
|
||||
return await call(primary)
|
||||
except Exception as exc:
|
||||
if not should_fallback(exc):
|
||||
if isinstance(exc, asyncio.TimeoutError) or not should_fallback(exc):
|
||||
raise
|
||||
logger.info(log_fmt, self.name, exc, *log_extra)
|
||||
return await fallback()
|
||||
return await call(fallback)
|
||||
|
||||
mode = str((self._config or {}).get("protocol", "auto")).lower().strip()
|
||||
if mode in ("stateless", "modern", "2026-07-28"):
|
||||
return await attempt(discover, initialize, lambda exc: True,
|
||||
return await attempt("discover", "initialize", lambda exc: True,
|
||||
"MCP server '%s': server/discover rejected (%s) despite "
|
||||
"protocol=%s — falling back to the legacy handshake", mode)
|
||||
if mode in ("legacy", "handshake"):
|
||||
return await initialize()
|
||||
return await call("initialize")
|
||||
if mode != "auto":
|
||||
logger.warning("MCP server '%s': unknown protocol=%r — treating as 'auto' "
|
||||
"(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"),
|
||||
"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:
|
||||
"""Handshake, discover, publish readiness, then serve until a lifecycle event.
|
||||
Clears stale breaker state from a prior outage but leaves the session UNPROVEN: a
|
||||
completed handshake is not proof of health (flapping transports handshake fine and drop
|
||||
moments later); only keepalive or tool-call success clears the reconnect budget."""
|
||||
"""Handshake, discover, publish readiness, then serve until a lifecycle event. Clears stale
|
||||
breaker state but leaves the session UNPROVEN: flapping transports handshake fine and drop
|
||||
moments later, so only keepalive or tool-call success clears the reconnect budget."""
|
||||
self.initialize_result = await self._negotiate_session(session, connect_timeout)
|
||||
self.session = session
|
||||
if mark_lifecycle:
|
||||
@@ -159,8 +134,8 @@ class MCPServerTransportMixin:
|
||||
return reason
|
||||
|
||||
async def _serve_transport(self, transport_cm, label: str, connect_timeout: float) -> str:
|
||||
"""Open *transport_cm*, wrap its streams in a ClientSession and serve it. Streams are
|
||||
unpacked positionally: mcp 1.x yields ``(read, write, get_session_id)``, 2.x ``(read, write)``.
|
||||
"""Open *transport_cm*, wrap its streams in a ClientSession and serve it. Streams are indexed,
|
||||
not unpacked: mcp 1.x yields ``(read, write, get_session_id)``, 2.x ``(read, write)``.
|
||||
A transport TaskGroup drop maps to ``"reconnect"`` instead of backoff/park."""
|
||||
try:
|
||||
async with transport_cm as _streams:
|
||||
@@ -171,47 +146,39 @@ class MCPServerTransportMixin:
|
||||
|
||||
# ------------------------------------------------------------------ stdio
|
||||
|
||||
def _resolve_stdio_config(self, config: dict):
|
||||
"""``(command, args, safe_env)`` from config, with the command resolved against the safe env."""
|
||||
command = config.get("command")
|
||||
if not command:
|
||||
raise ValueError(f"MCP server '{self.name}' has no 'command' in config")
|
||||
safe_env = _core._build_safe_env(config.get("env"))
|
||||
command, safe_env = _core._resolve_stdio_command(command, safe_env)
|
||||
return command, config.get("args", []), safe_env
|
||||
|
||||
def _track_spawned_children(self, new_pids: Set[int]) -> None:
|
||||
"""Ledger the freshly spawned stdio children (pids, pgids, machine spawn ledger)."""
|
||||
new_pgids = _capture_pgids(new_pids)
|
||||
"""Ledger the freshly spawned stdio children (pids, pgids, machine spawn ledger). pgids are
|
||||
captured while alive (getpgid fails once it exits; the sweep needs it for reparented descendants)."""
|
||||
new_pgids: Dict[int, int] = {}
|
||||
for pid in new_pids:
|
||||
try:
|
||||
new_pgids[pid] = os.getpgid(pid)
|
||||
except (AttributeError, ProcessLookupError, OSError): # Windows / already exited
|
||||
pass
|
||||
with _core._lock:
|
||||
for _pid in new_pids:
|
||||
_stdio_pids[_pid] = self.name
|
||||
_stdio_pids.update(dict.fromkeys(new_pids, 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 (startup sweeps reap orphans after an unclean exit); best-effort.
|
||||
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)
|
||||
|
||||
def _release_spawned_children(self, new_pids: Set[int]) -> None:
|
||||
"""Drop the ledger entries; any child (or its pgroup) still alive means SDK teardown
|
||||
failed (common on cancel mid-way on Linux, where setsid() children escape the cgroup)
|
||||
— mark it orphaned for the next cleanup sweep."""
|
||||
"""Drop the ledger entries; a child (or its pgroup) still alive means SDK teardown failed
|
||||
(common on mid-way cancel on Linux: setsid() children escape) — mark it orphaned for the sweep."""
|
||||
from gateway.status import _pid_exists
|
||||
with _core._lock:
|
||||
for pid in new_pids:
|
||||
_stdio_pids.pop(pid, None)
|
||||
# ``os.kill(pid, 0)`` is NOT a no-op on Windows; use the cross-platform check.
|
||||
# The child may have exited while descendants remain in its pgroup.
|
||||
# ``os.kill(pid, 0)`` is NOT a no-op on Windows; the child may be gone while
|
||||
# descendants remain in its pgroup.
|
||||
if _pid_exists(pid) or _pgroup_alive(_stdio_pgids.get(pid)):
|
||||
_orphan_stdio_pids.add(pid)
|
||||
_orphan_stdio_pid_servers[pid] = self.name
|
||||
else:
|
||||
# Nothing to reap — drop the pgid so PID reuse can't surface stale pgroup state.
|
||||
else: # nothing to reap — drop the pgid so PID reuse can't surface stale pgroup state
|
||||
_stdio_pgids.pop(pid, None)
|
||||
|
||||
async def _run_stdio(self, config: dict):
|
||||
@@ -223,46 +190,40 @@ class MCPServerTransportMixin:
|
||||
if not _core._ensure_mcp_sdk():
|
||||
raise ImportError(f"MCP server '{self.name}' requires the 'mcp' Python SDK, but "
|
||||
"it is not installed. Run `hermes setup` to install MCP support, then retry.")
|
||||
command, 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.
|
||||
command = config.get("command")
|
||||
if not command:
|
||||
raise ValueError(f"MCP server '{self.name}' has no 'command' in config")
|
||||
command, safe_env = _core._resolve_stdio_command(command, _core._build_safe_env(config.get("env")))
|
||||
await _osv_malware_preflight(self.name, command, config.get("args", []))
|
||||
# Parent-death watchdog so kill -9 / crash can't leave the child tree running (POSIX-only).
|
||||
# AFTER the OSV preflight so the check inspects the real package.
|
||||
command, args = _wrap_command_with_watchdog(command, args)
|
||||
command, args = _wrap_command_with_watchdog(command, config.get("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",
|
||||
)
|
||||
command=command, args=args, env=safe_env or None, cwd=config.get("cwd"),
|
||||
# Windows pipes can split non-UTF-8 bytes at chunk boundaries; substitute, don't raise.
|
||||
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
|
||||
# reconnect). Worker thread: the reaper blocks up to 2s (SIGTERM → wait → SIGKILL).
|
||||
# Reap orphans of prior attempts first, else each retry piles up zombie pairs. Unscoped on
|
||||
# purpose (also reaps servers that never reconnect). Off-loop: the reaper blocks up to 2s.
|
||||
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()
|
||||
pids_before = _core._snapshot_child_pids() # so the new child can be identified after spawn
|
||||
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.
|
||||
# Subprocess stderr goes to ~/.hermes/logs/mcp-stderr.log so banners can't 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.
|
||||
errlog = _core._get_mcp_stderr_log()
|
||||
async with _core.stdio_client(server_params, errlog=errlog) as (read_stream, write_stream):
|
||||
# New PIDs for force-kill cleanup, minus non-MCP children (slash_worker, LSP) that
|
||||
# race into the window: they share the TUI's pgid, so leaking them into _stdio_pgids
|
||||
# would make 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)
|
||||
self._stdio_child_pids = set(new_pids) # so in-flight calls fail fast when the child dies
|
||||
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 here: ``connect_timeout`` only bounds the caller's ``.result()``.
|
||||
# A server that never answers ``initialize`` would otherwise hang forever, skip the
|
||||
# ``finally`` and leak child + pipes 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:
|
||||
@@ -274,67 +235,56 @@ class MCPServerTransportMixin:
|
||||
|
||||
async def _preflight_content_type(self, url: str, *, headers: Optional[dict] = None,
|
||||
ssl_verify: bool = True, client_cert=None, timeout: float = 5.0) -> None:
|
||||
"""Probe *url* for an MCP-shaped response before the SDK connects. A URL pointing at a
|
||||
plain web page makes the SDK sit out the full ``connect_timeout`` before an opaque
|
||||
``CancelledError``; this raises :class:`NonMcpEndpointError` within ``timeout`` instead.
|
||||
Allow-list based: only a 2xx with a definite non-MCP content type is rejected, and only
|
||||
after a JSON-RPC ``initialize`` POST also fails to look like MCP (some servers serve a UI
|
||||
on GET but speak Streamable HTTP via POST). Missing content type, non-2xx, or transport
|
||||
errors pass silently — the real handshake stays the source of truth. Uses its own httpx
|
||||
client OUTSIDE the SDK's anyio task group so the error isn't wrapped in an ExceptionGroup."""
|
||||
"""Probe *url* before the SDK connects: a plain web page makes the SDK sit out the full
|
||||
``connect_timeout`` before an opaque ``CancelledError``; this raises NonMcpEndpointError within
|
||||
``timeout`` instead. Allow-list based: only a 2xx with a definite non-MCP content type is
|
||||
rejected, and only after a JSON-RPC ``initialize`` POST also fails to look like MCP (some
|
||||
servers serve a UI on GET but speak MCP via POST). Missing content type, non-2xx or transport
|
||||
errors pass silently — the real handshake stays the source of truth. Own httpx client, OUTSIDE
|
||||
the SDK's anyio task group, so the error isn't wrapped in an ExceptionGroup."""
|
||||
try:
|
||||
import httpx as _httpx
|
||||
except ImportError:
|
||||
return # No httpx → skip probe; SDK import would have failed first.
|
||||
|
||||
client_kwargs: dict = {"verify": ssl_verify, "follow_redirects": True, "timeout": _httpx.Timeout(timeout)}
|
||||
if client_cert is not None:
|
||||
client_kwargs["cert"] = client_cert
|
||||
probe_headers = dict(headers) if headers else {}
|
||||
try:
|
||||
async with _httpx.AsyncClient(**client_kwargs) as client:
|
||||
async with _httpx.AsyncClient(verify=ssl_verify, follow_redirects=True, timeout=_httpx.Timeout(timeout),
|
||||
**({"cert": client_cert} if client_cert is not None else {})) as client:
|
||||
# HEAD is cheapest; fall back to GET on 405/501.
|
||||
resp = await client.head(url, headers=probe_headers)
|
||||
if resp.status_code in (405, 501):
|
||||
resp = await client.get(url, headers=probe_headers)
|
||||
# Non-MCP content type on HEAD/GET: try a JSON-RPC POST before rejecting, so
|
||||
# POST-only servers aren't false positives.
|
||||
# Non-MCP content type on HEAD/GET: try a JSON-RPC POST so POST-only servers pass.
|
||||
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:
|
||||
return # DNS/connect/timeout/transport error — let the SDK try.
|
||||
|
||||
# Only judge 2xx: a 4xx/5xx may be an auth challenge or transient error the real
|
||||
# handshake handles correctly. No content type advertised → don't second-guess the SDK.
|
||||
if not _is_2xx(resp):
|
||||
return
|
||||
# Only judge 2xx (4xx/5xx may be an auth challenge the handshake handles); no content type
|
||||
# advertised → don't second-guess the SDK.
|
||||
ct_base = _content_type_base(resp)
|
||||
if not ct_base or ct_base in self._MCP_CONTENT_TYPES:
|
||||
if not _is_2xx(resp) or not ct_base or ct_base in self._MCP_CONTENT_TYPES:
|
||||
return
|
||||
raise NonMcpEndpointError(
|
||||
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
|
||||
pumps run in an anyio TaskGroup, so a transient stream drop escapes as a
|
||||
``BaseExceptionGroup``; unmapped, ``run()`` would back off and eventually park the server
|
||||
for 300s (deregistering its tools) over a sub-second glitch. Re-raise when it is not a
|
||||
transient drop: shutdown in progress (``_shutdown_event`` is set before the task is
|
||||
cancelled), the group carries KeyboardInterrupt/SystemExit or a real CancelledError (must
|
||||
propagate), or no live session was reached this attempt (``_ready`` unset — connect
|
||||
failures must go through backoff, not hot-loop)."""
|
||||
"""Map an SDK transport TaskGroup failure to a clean ``"reconnect"``: HTTP/SSE stream pumps
|
||||
run in an anyio TaskGroup, so a transient drop escapes as a ``BaseExceptionGroup`` that would
|
||||
otherwise back off and park the server for 300s over a sub-second glitch. Re-raise when it is
|
||||
not a transient drop: shutdown in progress (``_shutdown_event`` is set before cancel), the
|
||||
group carries KeyboardInterrupt/SystemExit or a real CancelledError, or no live session was
|
||||
reached this attempt (``_ready`` unset — connect failures must back off, not hot-loop)."""
|
||||
if (self._shutdown_event.is_set()
|
||||
or eg.split((KeyboardInterrupt, SystemExit))[0] is not None
|
||||
or eg.split(asyncio.CancelledError)[0] is not None
|
||||
@@ -345,9 +295,9 @@ class MCPServerTransportMixin:
|
||||
return "reconnect"
|
||||
|
||||
def _build_oauth_auth(self, url: str, config: dict):
|
||||
"""OAuth 2.1 PKCE via the central MCPOAuthManager so one provider is reused across
|
||||
reconnects and shared with config-time CLI paths. On setup failure (e.g.
|
||||
non-interactive without cached tokens) re-raise so only this server is reported failed."""
|
||||
"""OAuth 2.1 PKCE via the central MCPOAuthManager (one provider reused across reconnects and
|
||||
CLI paths). Setup failures (e.g. non-interactive without cached tokens) re-raise so only this
|
||||
server is reported failed."""
|
||||
if self._auth_type != "oauth":
|
||||
return None
|
||||
try:
|
||||
@@ -368,46 +318,44 @@ class MCPServerTransportMixin:
|
||||
raise ImportError(f"MCP server '{self.name}' requires SSE transport but "
|
||||
"mcp.client.sse.sse_client is not available. "
|
||||
"Upgrade the mcp package to get SSE support.")
|
||||
# sse_read_timeout bounds the gap between SSE events. SSE servers commonly idle for
|
||||
# minutes, so tool_timeout (60s) would drop the stream; 300s matches the Streamable
|
||||
# HTTP read timeout.
|
||||
sse_kwargs: dict = {"url": url, "headers": headers or None, "timeout": float(connect_timeout), "sse_read_timeout": 300.0}
|
||||
if oauth_auth is not None:
|
||||
# Forward OAuth to sse_client, else OAuth SSE servers 401 silently.
|
||||
sse_kwargs["auth"] = oauth_auth
|
||||
# sse_read_timeout bounds the gap between events: SSE servers idle for minutes, so 300s
|
||||
# (matching the Streamable HTTP read timeout), not tool_timeout. ``auth`` must be forwarded
|
||||
# or OAuth SSE servers 401 silently.
|
||||
sse_kwargs: dict = {"url": url, "headers": headers or None, "timeout": float(connect_timeout),
|
||||
"sse_read_timeout": 300.0, **({"auth": oauth_auth} if oauth_auth is not None else {})}
|
||||
if client_cert is not None or ssl_verify is not True:
|
||||
# sse_client has no verify/cert kwargs: wrap the SDK defaults (follow_redirects=True)
|
||||
# in an httpx_client_factory, forwarding the SDK's (headers, auth, timeout) and
|
||||
# layering TLS on top. The client MUST come from the SDK's own httpx module
|
||||
# (httpx2 on mcp >= 2.0) — see sdk_httpx().
|
||||
# sse_client has no verify/cert kwargs: an httpx_client_factory forwards the SDK's
|
||||
# (headers, auth, timeout) and layers TLS on top. The client MUST come from the SDK's
|
||||
# own httpx module (httpx2 on mcp >= 2.0) — see sdk_httpx().
|
||||
_httpx_mod = _core.sdk_httpx()
|
||||
|
||||
def _mcp_http_client_factory(headers=None, timeout=None, auth=None):
|
||||
kwargs: dict = {"follow_redirects": True, "verify": ssl_verify,
|
||||
"timeout": timeout if timeout is not None else _httpx_mod.Timeout(30.0, read=300.0)}
|
||||
kwargs.update({k: v for k, v in (("headers", headers), ("auth", auth), ("cert", client_cert)) if v is not None})
|
||||
return _httpx_mod.AsyncClient(**kwargs)
|
||||
|
||||
sse_kwargs["httpx_client_factory"] = _mcp_http_client_factory
|
||||
sse_kwargs["httpx_client_factory"] = lambda headers=None, timeout=None, auth=None: _httpx_mod.AsyncClient(
|
||||
follow_redirects=True, verify=ssl_verify,
|
||||
timeout=timeout if timeout is not None else _httpx_mod.Timeout(30.0, read=300.0),
|
||||
**{k: v for k, v in (("headers", headers), ("auth", auth), ("cert", client_cert)) if v is not None})
|
||||
return _core.sse_client(**sse_kwargs)
|
||||
|
||||
def _streamable_http_transport(self, url: str, headers: dict, connect_timeout: float,
|
||||
ssl_verify, client_cert, oauth_auth,
|
||||
strict_cfg_headers: bool, configured_header_names: set):
|
||||
"""Streamable HTTP context manager (mcp >= 1.24.0: caller-owned httpx client)."""
|
||||
"""Streamable HTTP context manager: mcp >= 1.24.0 gets a caller-owned httpx client; on the
|
||||
deprecated API (mcp < 1.24.0) the SDK owns the client."""
|
||||
if not _core._MCP_NEW_HTTP:
|
||||
return self._legacy_http_transport(url, headers, connect_timeout, ssl_verify, oauth_auth, strict_cfg_headers)
|
||||
# Build an explicit AsyncClient matching the SDK's create_mcp_http_client defaults. It
|
||||
# MUST come from the SDK's httpx module (httpx2 on mcp >= 2.0) since the SDK sends its
|
||||
# own Request objects through it — see sdk_httpx().
|
||||
if strict_cfg_headers:
|
||||
# Fail closed: without an owned client redirects can't be hooked.
|
||||
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.")
|
||||
return _core.streamablehttp_client(url, headers=headers, timeout=float(connect_timeout), verify=ssl_verify,
|
||||
**({"auth": oauth_auth} if oauth_auth is not None else {}))
|
||||
# Explicit AsyncClient matching the SDK's create_mcp_http_client defaults; MUST come from
|
||||
# the SDK's httpx module (httpx2 on mcp >= 2.0) since the SDK sends its own Requests through it.
|
||||
httpx = _core.sdk_httpx()
|
||||
_strip_auth_on_cross_origin_redirect = _make_redirect_header_stripper(
|
||||
httpx.URL(url), strict=strict_cfg_headers, configured_header_names=configured_header_names)
|
||||
client_kwargs: dict = {"follow_redirects": True, "timeout": httpx.Timeout(float(connect_timeout), read=300.0),
|
||||
"verify": ssl_verify, "event_hooks": {"response": [_strip_auth_on_cross_origin_redirect]}}
|
||||
if headers:
|
||||
client_kwargs["headers"] = headers
|
||||
client_kwargs.update({k: v for k, v in (("auth", oauth_auth), ("cert", client_cert)) if v is not None})
|
||||
"verify": ssl_verify, **({"headers": headers} if headers else {}),
|
||||
"event_hooks": {"response": [_strip_auth_on_cross_origin_redirect]},
|
||||
**{k: v for k, v in (("auth", oauth_auth), ("cert", client_cert)) if v is not None}}
|
||||
|
||||
@asynccontextmanager
|
||||
async def _owned_client_streams():
|
||||
@@ -418,20 +366,6 @@ class MCPServerTransportMixin:
|
||||
|
||||
return _owned_client_streams()
|
||||
|
||||
def _legacy_http_transport(self, url: str, headers: dict, connect_timeout: float,
|
||||
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.
|
||||
raise ImportError(f"MCP server '{self.name}' requires mcp >= 1.24.0 to "
|
||||
"enforce the portable redirect-header boundary "
|
||||
"(strict_redirect_headers). Upgrade the mcp package.")
|
||||
http_kwargs: dict = {"headers": headers, "timeout": float(connect_timeout), "verify": ssl_verify}
|
||||
if oauth_auth is not None:
|
||||
http_kwargs["auth"] = oauth_auth
|
||||
return _core.streamablehttp_client(url, **http_kwargs)
|
||||
|
||||
async def _run_http(self, config: dict):
|
||||
"""Run the server using HTTP/StreamableHTTP (or SSE) transport."""
|
||||
_core._ensure_mcp_sdk()
|
||||
@@ -441,30 +375,23 @@ 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.
|
||||
strict_cfg_headers = bool(config.get("strict_redirect_headers"))
|
||||
# Agent Plugins v1 strict_redirect_headers: configured headers MUST NOT follow a cross-origin
|
||||
# redirect. Capture their names BEFORE client-generated headers are merged in.
|
||||
configured_header_names = {key.lower() for key in headers}
|
||||
# Optional per-user identity header; explicit headers of the same name win.
|
||||
headers = _apply_identity_header(self.name, config, headers)
|
||||
# Some servers require MCP-Protocol-Version on the initial request; seed it (case-insensitive
|
||||
# user override wins) from the HANDSHAKE version, not the latest: ``initialize()``'s body speaks
|
||||
# the handshake era, and a 2026-07-28 header would route it onto the server's per-request-envelope
|
||||
# ladder, which rejects that body. The header must agree with what the body actually speaks.
|
||||
# Some servers require MCP-Protocol-Version on the first request; seed it (user override
|
||||
# wins) from the HANDSHAKE version, not the latest: a 2026-07-28 header would route the
|
||||
# handshake-era ``initialize()`` body onto the per-request-envelope ladder, which rejects it.
|
||||
if not any(key.lower() == "mcp-protocol-version" for key in headers):
|
||||
headers["mcp-protocol-version"] = _core.LATEST_HANDSHAKE_VERSION
|
||||
connect_timeout = config.get("connect_timeout", _core._DEFAULT_CONNECT_TIMEOUT)
|
||||
ssl_verify = config.get("ssl_verify", True)
|
||||
client_cert = _resolve_client_cert(self.name, config)
|
||||
oauth_auth = self._build_oauth_auth(url, config)
|
||||
|
||||
common = (url, headers, connect_timeout, config.get("ssl_verify", True), _resolve_client_cert(self.name, config),
|
||||
self._build_oauth_auth(url, config), bool(config.get("strict_redirect_headers")))
|
||||
if config.get("transport") == "sse":
|
||||
transport = self._sse_transport(url, headers, connect_timeout, ssl_verify, client_cert, oauth_auth, strict_cfg_headers)
|
||||
label = "SSE"
|
||||
transport, label = self._sse_transport(*common), "SSE"
|
||||
else:
|
||||
transport = self._streamable_http_transport(url, headers, connect_timeout, ssl_verify, client_cert,
|
||||
oauth_auth, strict_cfg_headers, configured_header_names)
|
||||
transport = self._streamable_http_transport(*common, configured_header_names)
|
||||
label = "HTTP" if _core._MCP_NEW_HTTP else "legacy HTTP"
|
||||
return await self._serve_transport(transport, label, float(connect_timeout))
|
||||
|
||||
@@ -473,8 +400,7 @@ class MCPServerTransportMixin:
|
||||
async def _discover_tools(self):
|
||||
"""Discover tools from the connected session. Capability-gated: prompt-/resource-only
|
||||
servers raise ``MCPError(-32601)`` on ``tools/list``, which would abort the connection."""
|
||||
# Fresh transport: re-probe with cheap ``ping`` in case the server gained support
|
||||
# across the reconnect.
|
||||
# Fresh transport: re-probe ``ping`` in case the server gained support across the reconnect.
|
||||
self._ping_unsupported = False
|
||||
if self.session is None:
|
||||
return
|
||||
@@ -482,21 +408,19 @@ class MCPServerTransportMixin:
|
||||
logger.info("MCP server '%s': does not advertise 'tools' capability — "
|
||||
"skipping tools/list (prompts/resources remain available)", self.name)
|
||||
self._tools = []
|
||||
self._register_discovered_tools_if_needed()
|
||||
return
|
||||
async with self._rpc_lock:
|
||||
self._list_cache_meta = {}
|
||||
self._tools = await _core._paginate_full_list(
|
||||
self.session.list_tools, "tools", self.name, cache_meta_out=self._list_cache_meta)
|
||||
else:
|
||||
async with self._rpc_lock:
|
||||
self._list_cache_meta = {}
|
||||
self._tools = await _core._paginate_full_list(
|
||||
self.session.list_tools, "tools", self.name, cache_meta_out=self._list_cache_meta)
|
||||
self._register_discovered_tools_if_needed()
|
||||
|
||||
def _register_discovered_tools_if_needed(self) -> None:
|
||||
"""Publish freshly discovered tools for a registry-owned server if none are registered.
|
||||
Initial registration normally happens in ``_discover_and_register_server`` after
|
||||
``start()``. On reconnect, outage handling may clear ``_ready`` and deregister stale
|
||||
tools; ownership via ``_servers`` authorizes publishing before readiness is restored so a
|
||||
revival never comes back with zero tools. A server retained after a recoverable initial
|
||||
failure is likewise owned before its first session, which authorizes its first publication."""
|
||||
"""Publish freshly discovered tools for a registry-owned server if none are registered
|
||||
(initial registration normally happens in ``_discover_and_register_server``). On reconnect,
|
||||
outage handling may clear ``_ready`` and deregister stale tools; ownership via ``_servers``
|
||||
authorizes publishing before readiness is restored so a revival never comes back with zero
|
||||
tools — likewise a server retained after a recoverable initial failure."""
|
||||
if self._registered_tool_names:
|
||||
return
|
||||
if not self._ready.is_set():
|
||||
@@ -504,8 +428,7 @@ 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.
|
||||
with _core._lock:
|
||||
if _core._servers.get(self.name) is self:
|
||||
_core._server_connect_errors.pop(self.name, None)
|
||||
|
||||
Reference in New Issue
Block a user