From 1f7beadae7a461c85bb4ab8f02f200efffc65f04 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 15:41:56 -0700 Subject: [PATCH] refactor(mcp): dedupe error classification ladders, optional-type import helper, cause-chain walker --- tools/mcp_tool_errors.py | 319 ++++++++++++++++----------------------- 1 file changed, 126 insertions(+), 193 deletions(-) diff --git a/tools/mcp_tool_errors.py b/tools/mcp_tool_errors.py index fe4d727a92..b55567f2b0 100644 --- a/tools/mcp_tool_errors.py +++ b/tools/mcp_tool_errors.py @@ -1,8 +1,9 @@ """MCP connection/transport error classification: URL validation, TLS client certs, identity headers, redirect header stripping, exception-group unwrapping, auth/session-expired/method-not-found detection and connect-error formatting. Split from tools/mcp_tool.py.""" -import logging import asyncio import errno +import importlib +import logging import os import re from typing import Any, List, Optional @@ -17,20 +18,22 @@ logger = logging.getLogger("tools.mcp_tool") _JSONRPC_UNSUPPORTED_PROTOCOL_VERSION = -32022 +def _jsonrpc_code(exc: BaseException): + """Structural ``MCPError.error.code`` (None when absent).""" + return getattr(getattr(exc, "error", None), "code", None) + + def _handshake_rejected_as_modern(exc: BaseException) -> bool: """True when a failed ``initialize`` signals a stateless-only (2026-07-28) server. Structural code check first, then substring fallback — never ``isinstance`` on SDK exception types (they arrive wrapped in ExceptionGroups and drift across generations). """ - err = getattr(exc, "error", None) - code = getattr(err, "code", None) or getattr(exc, "code", None) + code = _jsonrpc_code(exc) or getattr(exc, "code", None) if code in (_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, _core._JSONRPC_METHOD_NOT_FOUND): return True msg = str(exc).lower() - if not msg: - return False - return ( + return bool(msg) and ( "unsupported protocol version" in msg or str(_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION) in msg or _is_method_not_found_error(exc) @@ -41,18 +44,14 @@ def _is_method_not_found_error(exc: BaseException) -> bool: """True if *exc* is a JSON-RPC ``method not found`` (-32601). ``ping`` is optional in MCP; servers lacking it answer -32601. Structural - ``MCPError.error.code`` check first, then substring fallback — including - "Unknown method: ", which some servers use; without it the - ping→list_tools keepalive fallback never latches and reconnect-loops. + code check first, then substring fallback — including "Unknown method: ", + which some servers use; without it the ping→list_tools keepalive fallback + never latches and reconnect-loops. """ - err = getattr(exc, "error", None) - code = getattr(err, "code", None) - if code == _core._JSONRPC_METHOD_NOT_FOUND: + if _jsonrpc_code(exc) == _core._JSONRPC_METHOD_NOT_FOUND: return True msg = str(exc).lower() - if not msg: - return False - return ( + return bool(msg) and ( str(_core._JSONRPC_METHOD_NOT_FOUND) in msg or "method not found" in msg or "unknown method" in msg @@ -93,12 +92,10 @@ def _unwrap_exception_group(exc: BaseException) -> BaseException: while isinstance(leaf, BaseExceptionGroup) and leaf.exceptions: leaf = leaf.exceptions[0] raise leaf - chosen = exc.exceptions[0] - for sub in exc.exceptions: - if not _contains_only_cancellation(sub): - chosen = sub - break - exc = chosen + exc = next( + (sub for sub in exc.exceptions if not _contains_only_cancellation(sub)), + exc.exceptions[0], + ) return exc @@ -117,20 +114,20 @@ def _classify_mcp_failure(exc: BaseException) -> str: stdio command (FileNotFoundError / ENOENT). Everything else keeps backoff retry. """ root = _unwrap_exception_group(exc) - if _core._is_auth_error(root): - return "permanent" - if isinstance(root, (NonMcpEndpointError, InvalidMcpUrlError)): - return "permanent" - if isinstance(root, FileNotFoundError): - return "permanent" - if isinstance(root, OSError) and getattr(root, "errno", None) == errno.ENOENT: - return "permanent" - # 401/403 HTTPStatusError that _is_auth_error's type-gate missed - # (auth types not importable in this environment). - status = getattr(getattr(root, "response", None), "status_code", None) - if status in (401, 403): - return "permanent" - return "transient" + 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 in this environment). + or _response_status(root) 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: @@ -139,37 +136,25 @@ def _validate_remote_mcp_url(server_name: str, url: Any) -> str: Raises InvalidMcpUrlError naming the server for non-strings, missing/other schemes (stdio servers use ``command``, not ``url``), and empty hosts. """ + def _bad(detail: str) -> InvalidMcpUrlError: + return InvalidMcpUrlError(f"Invalid MCP URL for '{server_name}': {detail}") + if not isinstance(url, str): - raise InvalidMcpUrlError( - f"Invalid MCP URL for '{server_name}': expected a string, got " - f"{type(url).__name__}" - ) + raise _bad(f"expected a string, got {type(url).__name__}") stripped = url.strip() if not stripped: - raise InvalidMcpUrlError( - f"Invalid MCP URL for '{server_name}': empty url" - ) + raise _bad("empty url") try: parsed = urlparse(stripped) except Exception as exc: # urlparse is very permissive — belt and braces - raise InvalidMcpUrlError( - f"Invalid MCP URL for '{server_name}': {stripped!r} ({exc})" - ) from exc + raise _bad(f"{stripped!r} ({exc})") from exc if parsed.scheme.lower() not in {"http", "https"}: - raise InvalidMcpUrlError( - f"Invalid MCP URL for '{server_name}': scheme must be http or " - f"https, got {parsed.scheme!r} ({stripped!r})" - ) + raise _bad(f"scheme must be http or https, got {parsed.scheme!r} ({stripped!r})") if not parsed.netloc: - raise InvalidMcpUrlError( - f"Invalid MCP URL for '{server_name}': missing host ({stripped!r})" - ) + raise _bad(f"missing host ({stripped!r})") # ``urlparse`` accepts ``http://:8080`` (empty host, explicit port) — reject it. if not parsed.hostname: - raise InvalidMcpUrlError( - f"Invalid MCP URL for '{server_name}': missing hostname " - f"({stripped!r})" - ) + raise _bad(f"missing hostname ({stripped!r})") return stripped @@ -183,46 +168,37 @@ def _resolve_client_cert(server_name: str, config: dict): """ raw_cert = config.get("client_cert") raw_key = config.get("client_key") - if raw_cert is None and raw_key is None: return None + prefix = f"MCP server '{server_name}': " def _expand(path: Any, label: str) -> str: if not isinstance(path, str) or not path.strip(): raise ValueError( - f"MCP server '{server_name}': {label} must be a non-empty " - f"string path (got {type(path).__name__})" + f"{prefix}{label} must be a non-empty string path (got {type(path).__name__})" ) expanded = os.path.expanduser(path.strip()) if not os.path.isfile(expanded): - raise FileNotFoundError( - f"MCP server '{server_name}': {label} not found at " - f"{expanded!r}" - ) + 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"MCP server '{server_name}': specify either client_cert as " - f"a list [cert, key] OR client_cert + client_key, not both" + 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 (_expand(raw_cert[0], "client_cert[0]"), _expand(raw_cert[1], "client_cert[1]")) - if len(raw_cert) == 3: - cert_path = _expand(raw_cert[0], "client_cert[0]") - key_path = _expand(raw_cert[1], "client_cert[1]") - password = raw_cert[2] - if not isinstance(password, str): - raise ValueError( - f"MCP server '{server_name}': client_cert[2] (key " - f"passphrase) must be a string" - ) - return (cert_path, key_path, password) - raise ValueError( - f"MCP server '{server_name}': client_cert list form must have 2 " - f"or 3 elements (got {len(raw_cert)})" - ) + return pair + password = raw_cert[2] + if not isinstance(password, str): + raise ValueError(f"{prefix}client_cert[2] (key passphrase) must be a string") + return (*pair, password) cert_path = _expand(raw_cert, "client_cert") if raw_key is not None: @@ -241,39 +217,26 @@ def _resolve_identity_header(server_name: str, config: dict): raw = config.get("identity_header") if raw is None: return None - if not isinstance(raw, dict): - logger.warning( - "MCP server '%s': identity_header must be a mapping with " - "'name' and 'value'/'value_from' keys (got %s) — ignoring", - server_name, type(raw).__name__, - ) + + def _ignore(detail: str, *args): + logger.warning("MCP server '%s': identity_header " + detail + " — ignoring", server_name, *args) return None + + if not isinstance(raw, dict): + return _ignore("must be a mapping with 'name' and 'value'/'value_from' keys (got %s)", type(raw).__name__) name = raw.get("name") if not isinstance(name, str) or not name.strip(): - logger.warning( - "MCP server '%s': identity_header requires a non-empty " - "'name' — ignoring", server_name, - ) - return None + 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(): - logger.warning( - "MCP server '%s': identity_header with value_from: static " - "requires a non-empty string 'value' — ignoring", - server_name, - ) - return None + 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()) - logger.warning( - "MCP server '%s': identity_header value_from must be 'static' or " - "'profile' (got %r) — ignoring", server_name, value_from, - ) - return None + return _ignore("value_from must be 'static' or 'profile' (got %r)", value_from) def _apply_identity_header(server_name: str, config: dict, headers: dict) -> dict: @@ -309,63 +272,56 @@ def _make_redirect_header_stripper( header (lowercase names in *configured_header_names*) is stripped too — the v1 spec forbids forwarding package-configured headers cross-origin. """ + origin = (original_url.scheme, original_url.host, original_url.port) async def _strip_on_cross_origin_redirect(response): - if response.is_redirect and response.next_request: - target = response.next_request.url - if (target.scheme, target.host, target.port) != ( - original_url.scheme, original_url.host, original_url.port, - ): - response.next_request.headers.pop("authorization", None) - response.next_request.headers.pop("Authorization", None) - if strict: - for _name in configured_header_names: - while _name in response.next_request.headers: - del response.next_request.headers[_name] + if not (response.is_redirect and response.next_request): + return + target = response.next_request.url + if (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] return _strip_on_cross_origin_redirect +def _exc_causes(exc: BaseException) -> List[BaseException]: + """``__cause__`` then ``__context__`` of *exc*, when they are exceptions.""" + return [ + nested for nested in (exc.__cause__, exc.__context__) + if isinstance(nested, 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]: nested = getattr(current, "exceptions", None) if nested: - for child in nested: - missing = _find_missing(child) - if missing: - return missing - return None + return next(filter(None, map(_find_missing, nested)), None) if isinstance(current, FileNotFoundError): if getattr(current, "filename", None): return str(current.filename) match = re.search(r"No such file or directory: '([^']+)'", str(current)) if match: return match.group(1) - for attr in ("__cause__", "__context__"): - nested_exc = getattr(current, attr, None) - if isinstance(nested_exc, BaseException): - missing = _find_missing(nested_exc) - if missing: - return missing - return None + return next(filter(None, map(_find_missing, _exc_causes(current))), None) def _flatten_messages(current: BaseException) -> List[str]: nested = getattr(current, "exceptions", None) if nested: - flattened: List[str] = [] - for child in nested: - flattened.extend(_flatten_messages(child)) - return flattened - messages = [] + return [m for child in nested for m in _flatten_messages(child)] text = str(current).strip() - if text: - messages.append(text) - for attr in ("__cause__", "__context__"): - nested_exc = getattr(current, attr, None) - if isinstance(nested_exc, BaseException): - messages.extend(_flatten_messages(nested_exc)) + messages = [text] if text else [] + for nested_exc in _exc_causes(current): + messages.extend(_flatten_messages(nested_exc)) return messages or [current.__class__.__name__] missing = _find_missing(exc) @@ -378,11 +334,7 @@ def _format_connect_error(exc: BaseException) -> str: "that directory in mcp_servers..env.PATH)" ) return _sanitize_error(message) - - deduped: List[str] = [] - for item in _flatten_messages(exc): - if item not in deduped: - deduped.append(item) + deduped = list(dict.fromkeys(_flatten_messages(exc))) return _sanitize_error("; ".join(deduped[:3])) @@ -391,6 +343,15 @@ _AUTH_ERROR_TYPES: tuple = () _HTTP_STATUS_ERROR_TYPES: Optional[tuple] = None +def _optional_types(module: str, *names: str) -> list: + """``[module.name, ...]`` or ``[]`` when the module/attribute is unavailable.""" + try: + mod = importlib.import_module(module) + return [getattr(mod, name) for name in names] + except (ImportError, AttributeError): + return [] + + def _http_status_error_types() -> tuple: """``HTTPStatusError`` classes from both httpx flavours. @@ -398,19 +359,15 @@ def _http_status_error_types() -> tuple: Hermes' pinned ``httpx``; the classes are unrelated, so both go in the tuple. """ global _HTTP_STATUS_ERROR_TYPES - if _HTTP_STATUS_ERROR_TYPES is not None: - return _HTTP_STATUS_ERROR_TYPES - found: list = [] - sdk_mod = _core.sdk_httpx() - if sdk_mod is not None: - found.append(sdk_mod.HTTPStatusError) - try: - import httpx - if httpx.HTTPStatusError not in found: - found.append(httpx.HTTPStatusError) - except ImportError: - pass - _HTTP_STATUS_ERROR_TYPES = tuple(found) + if _HTTP_STATUS_ERROR_TYPES is None: + found: list = [] + sdk_mod = _core.sdk_httpx() + if sdk_mod is not None: + found.append(sdk_mod.HTTPStatusError) + for cls in _optional_types("httpx", "HTTPStatusError"): + if cls not in found: + found.append(cls) + _HTTP_STATUS_ERROR_TYPES = tuple(found) return _HTTP_STATUS_ERROR_TYPES @@ -422,26 +379,13 @@ def _get_auth_error_types() -> tuple: flavours — the latter needs the 401 status check in :func:`_is_auth_error`. """ global _AUTH_ERROR_TYPES - if _AUTH_ERROR_TYPES: - return _AUTH_ERROR_TYPES - types: list = [] - try: - from mcp.client.auth import OAuthFlowError, OAuthTokenError - types.extend([OAuthFlowError, OAuthTokenError]) - except ImportError: - pass - try: - from mcp.client.auth import UnauthorizedError # type: ignore # older SDKs - types.append(UnauthorizedError) - except ImportError: - pass - try: - from tools.mcp_oauth import OAuthNonInteractiveError - types.append(OAuthNonInteractiveError) - except ImportError: - pass - types.extend(_http_status_error_types()) - _AUTH_ERROR_TYPES = tuple(types) + 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()) + ) return _AUTH_ERROR_TYPES @@ -493,16 +437,9 @@ def _is_session_expired_error(exc: BaseException) -> bool: """ # AnyIO stream exceptions are often message-less (``str(ClosedResourceError()) == ""``), # so type checks are needed in addition to marker matching. - try: - from anyio import BrokenResourceError, ClosedResourceError, EndOfStream - - transport_error_types = ( - BrokenResourceError, - ClosedResourceError, - EndOfStream, - ) - except ImportError: # pragma: no cover - AnyIO is supplied by the MCP SDK - transport_error_types = () + 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 @@ -515,23 +452,19 @@ def _is_session_expired_error(exc: BaseException) -> bool: budget = _EXC_TRAVERSAL_MAX_NODES while stack and budget > 0: current = stack.pop() - if current is None: + if current is None or id(current) in seen: continue - identity = id(current) - if identity in seen: - continue - seen.add(identity) + seen.add(id(current)) budget -= 1 if isinstance(current, InterruptedError): return False - if isinstance(current, transport_error_types): - transport_error_found = True - # Messages vary across SDK versions and servers: match a narrow # allow-list of stable substrings, not exception type, to avoid false positives. msg = str(current).lower() - if msg and any(marker in msg for marker in _SESSION_EXPIRED_MARKERS): + if isinstance(current, transport_error_types) or ( + msg and any(marker in msg for marker in _SESSION_EXPIRED_MARKERS) + ): transport_error_found = True stack.extend(getattr(current, "exceptions", ()))