refactor(mcp): dedupe error classification ladders, optional-type import helper, cause-chain walker
This commit is contained in:
@@ -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: <name>", 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: <name>",
|
||||
which some servers use; without it the ping→list_tools keepalive fallback
|
||||
never latches and reconnect-loops.
|
||||
"""
|
||||
err = getattr(exc, "error", None)
|
||||
code = getattr(err, "code", None)
|
||||
if code == _core._JSONRPC_METHOD_NOT_FOUND:
|
||||
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.<name>.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", ()))
|
||||
|
||||
Reference in New Issue
Block a user