Merge branch 'simp/r2-tools-a' into simp/integration2
This commit is contained in:
@@ -1,18 +1,16 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Central manager for per-server MCP OAuth state (one instance per process).
|
||||
|
||||
Holds per-server provider instances and coordinates cross-process token reload
|
||||
(mtime-based disk watch, so tokens refreshed by cron/another CLI are picked up
|
||||
without a restart — Claude Code's ``invalidateOAuthCacheIfDiskChanged`` bug
|
||||
class), 401 deduplication (N concurrent tool calls hitting 401 with the same
|
||||
access_token trigger one recovery attempt) and reconnect signalling
|
||||
(``MCPServerTask`` in ``mcp_tool.py`` drives the reconnect; the manager decides
|
||||
when it is warranted).
|
||||
Holds per-server provider instances and coordinates cross-process token reload (mtime-based
|
||||
disk watch, so tokens refreshed by cron/another CLI are picked up without a restart), 401
|
||||
deduplication (N concurrent tool calls hitting 401 with the same access_token trigger one
|
||||
recovery attempt) and reconnect signalling (``MCPServerTask`` in ``mcp_tool.py`` drives the
|
||||
reconnect; the manager decides when it is warranted).
|
||||
|
||||
This module is the ONLY place that instantiates the SDK's ``OAuthClientProvider``
|
||||
for runtime use; other code paths go through ``get_manager()``. We lean on the
|
||||
SDK's lazy refresh rather than refreshing before every op: one ``stat()`` per
|
||||
tool call is cheaper than an await + refresh round-trip.
|
||||
This module is the ONLY place that instantiates the SDK's ``OAuthClientProvider`` for runtime
|
||||
use; other code paths go through ``get_manager()``. We lean on the SDK's lazy refresh rather
|
||||
than refreshing before every op: one ``stat()`` per tool call is cheaper than an await +
|
||||
refresh round-trip.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -31,28 +29,23 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _same_endpoint(a: str, b: str) -> bool:
|
||||
"""True if two URLs target the same endpoint: scheme, host (case-insensitive)
|
||||
and path, ignoring query/fragment. Confirms a rejected response actually came
|
||||
from the OAuth token endpoint before we act on an ``invalid_client`` body."""
|
||||
"""True if two URLs target the same endpoint: scheme, host (case-insensitive) and path,
|
||||
ignoring query/fragment. Confirms a rejected response actually came from the OAuth token
|
||||
endpoint before we act on an ``invalid_client`` body."""
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
try:
|
||||
pa, pb = urlsplit(a), urlsplit(b)
|
||||
except ValueError: # pragma: no cover — malformed URL
|
||||
return False
|
||||
return (
|
||||
pa.scheme == pb.scheme
|
||||
and pa.netloc.lower() == pb.netloc.lower()
|
||||
and pa.path.rstrip("/") == pb.path.rstrip("/")
|
||||
)
|
||||
return pa.scheme == pb.scheme and pa.netloc.lower() == pb.netloc.lower() and pa.path.rstrip("/") == pb.path.rstrip("/")
|
||||
|
||||
|
||||
@dataclass
|
||||
class _ProviderEntry:
|
||||
"""Per-server OAuth state. ``last_mtime_ns`` is the last-seen tokens-file
|
||||
mtime (0 = never read) for external-refresh detection; ``lock`` binds to
|
||||
whichever asyncio loop first awaits it (the MCP event loop);
|
||||
``pending_401`` dedupes thundering-herd 401s by failed access_token."""
|
||||
"""Per-server OAuth state. ``last_mtime_ns`` is the last-seen tokens-file mtime (0 = never
|
||||
read) for external-refresh detection; ``lock`` binds to whichever asyncio loop first awaits
|
||||
it (the MCP event loop); ``pending_401`` dedupes thundering-herd 401s by failed access_token."""
|
||||
|
||||
server_url: str
|
||||
oauth_config: Optional[dict]
|
||||
@@ -64,29 +57,27 @@ class _ProviderEntry:
|
||||
|
||||
# -- HermesMCPOAuthProvider — OAuthClientProvider subclass with disk-watch ----
|
||||
class _HermesRuntimeProviderMixin:
|
||||
"""Runtime-only provider behaviour layered over ``HermesProviderMixin``:
|
||||
pre-flow disk-mtime reload, expiry seeding on cold load, pre-flight metadata
|
||||
discovery, dead-client-registration detection and the bidirectional
|
||||
``async_auth_flow`` bridge. Must precede the SDK class in the MRO.
|
||||
"""
|
||||
"""Runtime-only provider behaviour layered over ``HermesProviderMixin``: pre-flow disk-mtime
|
||||
reload, expiry seeding on cold load, pre-flight metadata discovery, dead-client-registration
|
||||
detection and the bidirectional ``async_auth_flow`` bridge. Must precede the SDK class in
|
||||
the MRO."""
|
||||
|
||||
_hermes_logger = logger
|
||||
|
||||
def __init__(self, *args: Any, server_name: str = "", preregistered: bool = False, **kwargs: Any):
|
||||
super().__init__(*args, **kwargs)
|
||||
# mcp 2.0 uses a task-owned anyio.Lock held across the yielded resource
|
||||
# request: a session-long GET blocks every concurrent POST, and HTTPX may
|
||||
# close the auth-flow generator from another task. A binary semaphore
|
||||
# keeps mutual exclusion without task ownership; async_auth_flow narrows
|
||||
# its scope around resource I/O.
|
||||
# mcp 2.0 uses a task-owned anyio.Lock held across the yielded resource request: a
|
||||
# session-long GET blocks every concurrent POST, and HTTPX may close the auth-flow
|
||||
# generator from another task. A binary semaphore keeps mutual exclusion without task
|
||||
# ownership; async_auth_flow narrows its scope around resource I/O.
|
||||
import anyio
|
||||
|
||||
self.context.lock = anyio.Semaphore(1, max_value=1)
|
||||
self._hermes_server_name = server_name
|
||||
self._hermes_home = ""
|
||||
# A config-supplied (pre-registered) client_id rejected as invalid_client
|
||||
# means the *config* is wrong — re-registration can't help, so only
|
||||
# dynamically-registered clients auto-heal.
|
||||
# A config-supplied (pre-registered) client_id rejected as invalid_client means the
|
||||
# *config* is wrong — re-registration can't help, so only dynamically-registered
|
||||
# clients auto-heal.
|
||||
self._hermes_preregistered = preregistered
|
||||
|
||||
def _hermes_storage(self):
|
||||
@@ -103,15 +94,14 @@ class _HermesRuntimeProviderMixin:
|
||||
"""Load stored state, seed ``token_expiry_time``, restore/prefetch metadata.
|
||||
|
||||
The SDK's ``_initialize`` populates ``current_tokens`` but never calls
|
||||
``update_token_expiry``, so ``is_token_valid()`` is True for any loaded
|
||||
token regardless of age and a restarted process ships stale Bearer tokens
|
||||
(some providers answer 200 with an app-level auth error the transport
|
||||
can't see). Seeding the expiry makes the SDK take ``can_refresh_token()``
|
||||
and refresh first; ``HermesTokenStorage`` persists absolute ``expires_at``
|
||||
so the TTL reflects wall-clock age. Metadata is restored from disk, else
|
||||
discovered pre-flight when we hold tokens but no metadata: otherwise
|
||||
``_refresh_token`` guesses ``{server_url}/token`` (wrong for split-origin
|
||||
providers such as BetterStack), 404s, and we fall through to browser reauth.
|
||||
``update_token_expiry``, so ``is_token_valid()`` is True for any loaded token regardless
|
||||
of age and a restarted process ships stale Bearer tokens (some providers answer 200 with
|
||||
an app-level auth error the transport can't see). Seeding the expiry makes the SDK take
|
||||
``can_refresh_token()`` and refresh first; ``HermesTokenStorage`` persists absolute
|
||||
``expires_at`` so the TTL reflects wall-clock age. Metadata is restored from disk, else
|
||||
discovered pre-flight when we hold tokens but no metadata: otherwise ``_refresh_token``
|
||||
guesses ``{server_url}/token`` (wrong for split-origin providers such as BetterStack),
|
||||
404s, and we fall through to browser reauth.
|
||||
"""
|
||||
await super()._initialize()
|
||||
tokens = self.context.current_tokens
|
||||
@@ -131,20 +121,16 @@ class _HermesRuntimeProviderMixin:
|
||||
if tokens is not None and self.context.oauth_metadata is None:
|
||||
try:
|
||||
await self._prefetch_oauth_metadata()
|
||||
except Exception as exc: # pragma: no cover — defensive
|
||||
# Non-fatal: the SDK's 401-branch discovery runs next request.
|
||||
except Exception as exc: # pragma: no cover — non-fatal: the SDK's 401-branch discovery runs next request
|
||||
self._log_nonfatal("pre-flight metadata discovery", exc)
|
||||
|
||||
async def _prefetch_oauth_metadata(self) -> None:
|
||||
"""Fetch PRM + ASM from the well-known endpoints and cache on context.
|
||||
|
||||
Mirrors the SDK's 401-branch discovery but runs before the first request.
|
||||
Uses the SDK's own URL builders/response handlers so we track whatever
|
||||
the pinned SDK version expects.
|
||||
"""
|
||||
"""Fetch PRM + ASM from the well-known endpoints and cache on context. Mirrors the SDK's
|
||||
401-branch discovery but runs before the first request, using the SDK's own URL
|
||||
builders/response handlers so we track whatever the pinned SDK version expects."""
|
||||
# The SDK's httpx flavour, not Hermes' — mcp 2.0 builds on httpx2 and
|
||||
# `create_oauth_metadata_request` returns *its* Request objects, which
|
||||
# only its own AsyncClient can send (tools.mcp_tool.sdk_httpx).
|
||||
# `create_oauth_metadata_request` returns *its* Request objects, which only its own
|
||||
# AsyncClient can send (tools.mcp_tool.sdk_httpx).
|
||||
from tools.mcp_tool import sdk_httpx
|
||||
httpx = sdk_httpx()
|
||||
if httpx is None: # pragma: no cover — SDK import would have failed
|
||||
@@ -163,10 +149,7 @@ class _HermesRuntimeProviderMixin:
|
||||
try:
|
||||
return await client.send(create_oauth_metadata_request(url))
|
||||
except httpx.HTTPError as exc:
|
||||
logger.debug(
|
||||
"MCP OAuth '%s': %s discovery to %s failed: %s",
|
||||
self._hermes_server_name, label, url, exc,
|
||||
)
|
||||
logger.debug("MCP OAuth '%s': %s discovery to %s failed: %s", self._hermes_server_name, label, url, exc)
|
||||
return None
|
||||
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
@@ -180,8 +163,7 @@ class _HermesRuntimeProviderMixin:
|
||||
self.context.auth_server_url = str(prm.authorization_servers[0])
|
||||
break
|
||||
|
||||
# Step 2: ASM discovery against auth_server_url (server_url fallback
|
||||
# for legacy providers).
|
||||
# Step 2: ASM discovery against auth_server_url (server_url fallback for legacy providers).
|
||||
for url in build_oauth_authorization_server_metadata_discovery_urls(self.context.auth_server_url, server_url):
|
||||
resp = await _send(client, url, "ASM")
|
||||
if resp is None:
|
||||
@@ -202,8 +184,8 @@ class _HermesRuntimeProviderMixin:
|
||||
break
|
||||
|
||||
def _persist_oauth_metadata_if_changed(self) -> None:
|
||||
"""Save metadata the SDK discovered lazily (401 branch) for future
|
||||
restarts; no-op when absent, not our storage, or unchanged."""
|
||||
"""Save metadata the SDK discovered lazily (401 branch) for future restarts; no-op when
|
||||
absent, not our storage, or unchanged."""
|
||||
meta = self.context.oauth_metadata
|
||||
storage = self._hermes_storage()
|
||||
if meta is None or storage is None:
|
||||
@@ -213,16 +195,11 @@ class _HermesRuntimeProviderMixin:
|
||||
storage.save_oauth_metadata(meta)
|
||||
|
||||
async def _is_invalid_client_at_token_endpoint(self, response: Any) -> bool:
|
||||
"""True when *response* is the token endpoint rejecting our client_id
|
||||
with ``invalid_client`` (whole word, so RFC 7591's
|
||||
``invalid_client_metadata`` does not trip it). The body is read only
|
||||
after the endpoint matches."""
|
||||
"""True when *response* is the token endpoint rejecting our client_id with
|
||||
``invalid_client`` (whole word, so RFC 7591's ``invalid_client_metadata`` does not trip
|
||||
it). The body is read only after the endpoint matches."""
|
||||
meta = getattr(self.context, "oauth_metadata", None)
|
||||
token_endpoint = (
|
||||
str(meta.token_endpoint)
|
||||
if meta is not None and getattr(meta, "token_endpoint", None)
|
||||
else None
|
||||
)
|
||||
token_endpoint = str(meta.token_endpoint) if meta is not None and getattr(meta, "token_endpoint", None) else None
|
||||
req = getattr(response, "request", None)
|
||||
req_url = str(req.url) if req is not None else None
|
||||
if not token_endpoint or not req_url or not _same_endpoint(req_url, token_endpoint):
|
||||
@@ -233,17 +210,15 @@ class _HermesRuntimeProviderMixin:
|
||||
async def _maybe_flag_poisoned_client(self, response: Any) -> None:
|
||||
"""Detect a dead client registration and force re-registration.
|
||||
|
||||
An ``invalid_client`` rejection of our ``client_id`` at the token
|
||||
endpoint (exchange or refresh) proves the cached registration is dead
|
||||
server-side; delete ``client.json`` (+ stale metadata) so the SDK re-runs
|
||||
DCR next flow. The browser-side "Redirect URI Mismatch" case has no HTTP
|
||||
signal and is left to ``hermes mcp reauth``.
|
||||
An ``invalid_client`` rejection of our ``client_id`` at the token endpoint (exchange or
|
||||
refresh) proves the cached registration is dead server-side; delete ``client.json``
|
||||
(+ stale metadata) so the SDK re-runs DCR next flow. The browser-side "Redirect URI
|
||||
Mismatch" case has no HTTP signal and is left to ``hermes mcp reauth``.
|
||||
|
||||
Conservative by construction — acts ONLY when status is 400/401, the
|
||||
request hit the discovered ``token_endpoint`` (the only request carrying
|
||||
our ``client_id``), and the body carries ``invalid_client``.
|
||||
Pre-registered clients are never poisoned. Best-effort: any failure is
|
||||
swallowed so a miss never breaks the live flow. If ``token_endpoint`` was
|
||||
Conservative by construction — acts ONLY when status is 400/401, the request hit the
|
||||
discovered ``token_endpoint`` (the only request carrying our ``client_id``), and the body
|
||||
carries ``invalid_client``. Pre-registered clients are never poisoned. Best-effort: any
|
||||
failure is swallowed so a miss never breaks the live flow. If ``token_endpoint`` was
|
||||
never discovered the guard returns early.
|
||||
"""
|
||||
try:
|
||||
@@ -253,18 +228,16 @@ class _HermesRuntimeProviderMixin:
|
||||
return
|
||||
|
||||
storage = self._hermes_storage()
|
||||
# If the rejected client_id was our CIMD URL, re-presenting it would
|
||||
# loop (the server already fetched and refused it). Drop the URL so
|
||||
# the retry takes DCR, and mark it on disk so the next process
|
||||
# doesn't walk back into the same refusal (`hermes mcp login` clears
|
||||
# the marker).
|
||||
# If the rejected client_id was our CIMD URL, re-presenting it would loop (the
|
||||
# server already fetched and refused it). Drop the URL so the retry takes DCR, and
|
||||
# mark it on disk so the next process doesn't walk back into the same refusal
|
||||
# (`hermes mcp login` clears the marker).
|
||||
cimd_url = getattr(self.context, "client_metadata_url", None)
|
||||
rejected_id = getattr(self.context.client_info, "client_id", None)
|
||||
if cimd_url and rejected_id == cimd_url:
|
||||
logger.warning(
|
||||
"MCP OAuth '%s': authorization server rejected our "
|
||||
"Client ID Metadata Document (%s) with invalid_client "
|
||||
"— falling back to dynamic client registration.",
|
||||
"MCP OAuth '%s': authorization server rejected our Client ID Metadata Document (%s) "
|
||||
"with invalid_client — falling back to dynamic client registration.",
|
||||
self._hermes_server_name, cimd_url,
|
||||
)
|
||||
self.context.client_metadata_url = None
|
||||
@@ -282,17 +255,14 @@ class _HermesRuntimeProviderMixin:
|
||||
async def async_auth_flow(self, request): # type: ignore[override]
|
||||
# Pre-flow hook: reload from disk if it changed (non-fatal on error).
|
||||
try:
|
||||
await get_manager().invalidate_if_disk_changed(
|
||||
self._hermes_server_name, hermes_home=self._hermes_home
|
||||
)
|
||||
await get_manager().invalidate_if_disk_changed(self._hermes_server_name, hermes_home=self._hermes_home)
|
||||
except Exception as exc: # pragma: no cover — defensive
|
||||
self._log_nonfatal("pre-flow disk-watch", exc)
|
||||
|
||||
# Bridge the bidirectional generator protocol by hand: httpx feeds
|
||||
# responses back via ``auth_flow.asend(response)``. A naive
|
||||
# ``async for item in inner: yield item`` DISCARDS those values, so the
|
||||
# SDK's ``response = yield request`` sees None and crashes on
|
||||
# ``response.status_code`` (tests/tools/test_mcp_oauth_bidirectional.py).
|
||||
# Bridge the bidirectional generator protocol by hand: httpx feeds responses back via
|
||||
# ``auth_flow.asend(response)``. A naive ``async for item in inner: yield item``
|
||||
# DISCARDS those values, so the SDK's ``response = yield request`` sees None and
|
||||
# crashes on ``response.status_code``.
|
||||
inner = super().async_auth_flow(request)
|
||||
resource_lock_released = False
|
||||
sent_access_token = None
|
||||
@@ -300,10 +270,9 @@ class _HermesRuntimeProviderMixin:
|
||||
try:
|
||||
outgoing = await inner.__anext__()
|
||||
while True:
|
||||
# The SDK holds context.lock for its whole generator, even while
|
||||
# HTTPX waits on the MCP request. Release it for that request
|
||||
# only; discovery/refresh/registration/exchange stay serialized
|
||||
# exactly as the SDK implements them.
|
||||
# The SDK holds context.lock for its whole generator, even while HTTPX waits on
|
||||
# the MCP request. Release it for that request only; discovery/refresh/
|
||||
# registration/exchange stay serialized exactly as the SDK implements them.
|
||||
if outgoing is request:
|
||||
tokens = self.context.current_tokens
|
||||
sent_access_token = tokens.access_token if tokens is not None else None
|
||||
@@ -313,9 +282,9 @@ class _HermesRuntimeProviderMixin:
|
||||
if resource_lock_released:
|
||||
await self.context.lock.acquire()
|
||||
resource_lock_released = False
|
||||
# Another request may have refreshed/authorized while this one
|
||||
# was in flight: retry with that token instead of a duplicate
|
||||
# OAuth transition from the stale 401/403.
|
||||
# Another request may have refreshed/authorized while this one was in flight:
|
||||
# retry with that token instead of a duplicate OAuth transition from the stale
|
||||
# 401/403.
|
||||
tokens = self.context.current_tokens
|
||||
if (
|
||||
getattr(incoming, "status_code", None) in (401, 403)
|
||||
@@ -336,9 +305,8 @@ class _HermesRuntimeProviderMixin:
|
||||
return
|
||||
finally:
|
||||
if resource_lock_released:
|
||||
# Balance the SDK's surrounding ``async with`` even when HTTPX
|
||||
# cancels/closes the flow mid-request. Shield only this local
|
||||
# bookkeeping.
|
||||
# Balance the SDK's surrounding ``async with`` even when HTTPX cancels/closes
|
||||
# the flow mid-request. Shield only this local bookkeeping.
|
||||
import anyio
|
||||
|
||||
with anyio.CancelScope(shield=True):
|
||||
@@ -351,21 +319,18 @@ class _HermesRuntimeProviderMixin:
|
||||
|
||||
|
||||
def _make_hermes_provider_class() -> Optional[type]:
|
||||
"""Lazy-import the SDK base class and return our subclass (None if the
|
||||
SDK's OAuth module is unavailable, so this module still imports)."""
|
||||
"""Lazy-import the SDK base class and return our subclass (None if the SDK's OAuth module is
|
||||
unavailable, so this module still imports)."""
|
||||
try:
|
||||
from mcp.client.auth.oauth2 import OAuthClientProvider
|
||||
except ImportError: # pragma: no cover — SDK required in CI
|
||||
return None
|
||||
|
||||
class HermesMCPOAuthProvider(_HermesRuntimeProviderMixin, HermesProviderMixin, OAuthClientProvider):
|
||||
"""OAuthClientProvider with pre-flow disk-mtime reload.
|
||||
|
||||
Before every ``async_auth_flow`` the manager checks whether the tokens
|
||||
file changed on disk and, if so, resets ``_initialized`` so the next
|
||||
flow re-reads storage — making external refreshes visible to a running
|
||||
session. Token-endpoint fixes come from ``HermesProviderMixin``.
|
||||
"""
|
||||
"""OAuthClientProvider with pre-flow disk-mtime reload: before every ``async_auth_flow``
|
||||
the manager checks whether the tokens file changed on disk and, if so, resets
|
||||
``_initialized`` so the next flow re-reads storage — making external refreshes visible
|
||||
to a running session. Token-endpoint fixes come from ``HermesProviderMixin``."""
|
||||
|
||||
return HermesMCPOAuthProvider
|
||||
|
||||
@@ -376,50 +341,36 @@ _HERMES_PROVIDER_CLS: Optional[type] = _make_hermes_provider_class()
|
||||
|
||||
# -- Manager -----------------------------------------------------------------
|
||||
class MCPOAuthManager:
|
||||
"""Single source of truth for per-server MCP OAuth state.
|
||||
|
||||
Thread-safe: the ``_entries`` dict is guarded by ``_entries_lock`` for
|
||||
get-or-create semantics. Per-entry state is guarded by the entry's own
|
||||
``asyncio.Lock`` (used from the MCP event loop thread).
|
||||
"""
|
||||
"""Single source of truth for per-server MCP OAuth state. Thread-safe: the ``_entries`` dict
|
||||
is guarded by ``_entries_lock`` for get-or-create semantics; per-entry state is guarded by
|
||||
the entry's own ``asyncio.Lock`` (used from the MCP event loop thread)."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._entries: dict[tuple[str, str], _ProviderEntry] = {}
|
||||
self._entries_lock = threading.Lock()
|
||||
# Strong refs to in-flight 401 tasks so the loop's weak bookkeeping
|
||||
# cannot GC them mid-run and leave `await pending` hanging forever.
|
||||
# Strong refs to in-flight 401 tasks so the loop's weak bookkeeping cannot GC them
|
||||
# mid-run and leave `await pending` hanging forever.
|
||||
self._inflight_tasks: set[asyncio.Task] = set()
|
||||
|
||||
# -- Provider construction / caching -------------------------------------
|
||||
|
||||
# -- Provider construction / caching --
|
||||
def get_or_build_provider(self, server_name: str, server_url: str, oauth_config: Optional[dict]) -> Optional[Any]:
|
||||
"""Return a cached OAuth provider for ``server_name`` or build one.
|
||||
|
||||
Idempotent: repeat calls with the same name return the same instance.
|
||||
If ``server_url`` changes for a given name, the cached entry is
|
||||
discarded and a fresh provider is built.
|
||||
|
||||
Returns None if the MCP SDK's OAuth support is unavailable.
|
||||
"""
|
||||
"""Return a cached OAuth provider for ``server_name`` or build one. Idempotent: repeat
|
||||
calls with the same name return the same instance; if ``server_url`` changes for a
|
||||
given name the cached entry is discarded and a fresh provider is built. None if the MCP
|
||||
SDK's OAuth support is unavailable."""
|
||||
key = self._key(server_name)
|
||||
with self._entries_lock:
|
||||
entry = self._entries.get(key)
|
||||
if entry is not None and entry.server_url != server_url:
|
||||
logger.info(
|
||||
"MCP OAuth '%s': URL changed from %s to %s, discarding cache",
|
||||
server_name, entry.server_url, server_url,
|
||||
)
|
||||
logger.info("MCP OAuth '%s': URL changed from %s to %s, discarding cache", server_name, entry.server_url, server_url)
|
||||
entry = None
|
||||
|
||||
if entry is None:
|
||||
entry = _ProviderEntry(server_url=server_url, oauth_config=oauth_config)
|
||||
self._entries[key] = entry
|
||||
|
||||
if entry.provider is None:
|
||||
entry.provider = self._build_provider(server_name, entry)
|
||||
if entry.provider is not None:
|
||||
entry.provider._hermes_home = key[0]
|
||||
|
||||
return entry.provider
|
||||
|
||||
@staticmethod
|
||||
@@ -430,36 +381,24 @@ class MCPOAuthManager:
|
||||
return (str(home.expanduser().resolve(strict=False)), server_name)
|
||||
|
||||
def _build_provider(self, server_name: str, entry: _ProviderEntry) -> Optional[Any]:
|
||||
"""Build a :class:`HermesMCPOAuthProvider` from the shared
|
||||
``tools.mcp_oauth`` helpers; None if the SDK's OAuth support is unavailable."""
|
||||
"""Build a :class:`HermesMCPOAuthProvider` from the shared ``tools.mcp_oauth`` helpers;
|
||||
None if the SDK's OAuth support is unavailable."""
|
||||
if _HERMES_PROVIDER_CLS is None:
|
||||
logger.warning("MCP OAuth '%s': SDK auth module unavailable", server_name)
|
||||
return None
|
||||
|
||||
# Local imports avoid circular deps at module import time.
|
||||
from tools.mcp_dashboard_oauth import get_dashboard_oauth_flow
|
||||
from tools.mcp_oauth import _OAUTH_AVAILABLE, OAuthNonInteractiveError, _is_interactive
|
||||
from tools.mcp_oauth_provider import build_provider_kwargs, prepare_oauth_config
|
||||
|
||||
if not _OAUTH_AVAILABLE:
|
||||
return None
|
||||
|
||||
cfg, storage = prepare_oauth_config(server_name, entry.server_url, entry.oauth_config)
|
||||
|
||||
from tools.mcp_dashboard_oauth import get_dashboard_oauth_flow
|
||||
|
||||
if (
|
||||
get_dashboard_oauth_flow() is None
|
||||
and not _is_interactive()
|
||||
and not storage.has_cached_tokens()
|
||||
):
|
||||
if get_dashboard_oauth_flow() is None and not _is_interactive() and not storage.has_cached_tokens():
|
||||
raise OAuthNonInteractiveError(
|
||||
"MCP OAuth for "
|
||||
f"'{server_name}': non-interactive environment and no "
|
||||
"cached tokens found. Run `hermes mcp login "
|
||||
f"{server_name}` interactively first to complete initial "
|
||||
"authorization."
|
||||
f"MCP OAuth for '{server_name}': non-interactive environment and no cached tokens found. "
|
||||
f"Run `hermes mcp login {server_name}` interactively first to complete initial authorization."
|
||||
)
|
||||
|
||||
return _HERMES_PROVIDER_CLS(
|
||||
server_name=server_name,
|
||||
preregistered=bool(cfg.get("client_id")),
|
||||
@@ -468,11 +407,8 @@ class MCPOAuthManager:
|
||||
)
|
||||
|
||||
def remove(self, server_name: str, *, hermes_home: str | Path | None = None) -> _ProviderEntry | None:
|
||||
"""Evict the provider from cache AND delete tokens from disk.
|
||||
|
||||
Called by ``hermes mcp remove <name>`` and (indirectly) by
|
||||
``hermes mcp login <name>`` during forced re-auth.
|
||||
"""
|
||||
"""Evict the provider from cache AND delete tokens from disk (``hermes mcp remove`` and,
|
||||
indirectly, ``hermes mcp login`` during forced re-auth)."""
|
||||
entry = self.evict(server_name, hermes_home=hermes_home)
|
||||
from tools.mcp_oauth import remove_oauth_tokens
|
||||
remove_oauth_tokens(server_name, hermes_home=hermes_home)
|
||||
@@ -493,43 +429,34 @@ class MCPOAuthManager:
|
||||
with self._entries_lock:
|
||||
return self._entries.pop(self._key(server_name, hermes_home), None)
|
||||
|
||||
# -- Disk watch ----------------------------------------------------------
|
||||
|
||||
# -- Disk watch --
|
||||
async def invalidate_if_disk_changed(self, server_name: str, *, hermes_home: str | Path | None = None) -> bool:
|
||||
"""Force the SDK provider to reload when the tokens file mtime changed.
|
||||
|
||||
Returns True if invalidated. This is the external-refresh fix: a cron
|
||||
job writes fresh tokens and the next tool call picks them up.
|
||||
"""
|
||||
"""Force the SDK provider to reload when the tokens file mtime changed; True if
|
||||
invalidated. This is the external-refresh fix: a cron job writes fresh tokens and the
|
||||
next tool call picks them up."""
|
||||
from tools.mcp_oauth import _get_token_dir, _safe_filename
|
||||
|
||||
entry = self._entries.get(self._key(server_name, hermes_home))
|
||||
if entry is None or entry.provider is None:
|
||||
return False
|
||||
|
||||
async with entry.lock:
|
||||
tokens_path = _get_token_dir(hermes_home) / f"{_safe_filename(server_name)}.json"
|
||||
try:
|
||||
mtime_ns = tokens_path.stat().st_mtime_ns
|
||||
except (FileNotFoundError, OSError):
|
||||
return False
|
||||
|
||||
if mtime_ns == entry.last_mtime_ns:
|
||||
return False
|
||||
old = entry.last_mtime_ns
|
||||
entry.last_mtime_ns = mtime_ns
|
||||
# `_initialized` is private SDK API but stable across the versions
|
||||
# we pin (>=1.26.0); resetting it forces a reload.
|
||||
# `_initialized` is private SDK API but stable across the versions we pin
|
||||
# (>=1.26.0); resetting it forces a reload.
|
||||
if hasattr(entry.provider, "_initialized"):
|
||||
entry.provider._initialized = False # noqa: SLF001
|
||||
logger.info(
|
||||
"MCP OAuth '%s': tokens file changed (mtime %d -> %d), forcing reload",
|
||||
server_name, old, mtime_ns,
|
||||
)
|
||||
logger.info("MCP OAuth '%s': tokens file changed (mtime %d -> %d), forcing reload", server_name, old, mtime_ns)
|
||||
return True
|
||||
|
||||
# -- 401 handler (dedup'd) -----------------------------------------------
|
||||
|
||||
# -- 401 handler (dedup'd) --
|
||||
async def _recover_401(self, server_name: str, entry: _ProviderEntry, key: str, pending: asyncio.Future) -> None:
|
||||
"""Single recovery attempt behind *pending*; always clears the dedup slot."""
|
||||
try:
|
||||
@@ -538,9 +465,8 @@ class MCPOAuthManager:
|
||||
if not pending.done():
|
||||
pending.set_result(True)
|
||||
return
|
||||
|
||||
# Step 2: No disk change — if the SDK can refresh in place, let the
|
||||
# caller retry (the httpx.Auth flow refreshes on the next request).
|
||||
# Step 2: No disk change — if the SDK can refresh in place, let the caller retry
|
||||
# (the httpx.Auth flow refreshes on the next request).
|
||||
can_refresh_fn = getattr(getattr(entry.provider, "context", None), "can_refresh_token", None)
|
||||
try:
|
||||
can_refresh = bool(can_refresh_fn()) if callable(can_refresh_fn) else False
|
||||
@@ -556,21 +482,16 @@ class MCPOAuthManager:
|
||||
entry.pending_401.pop(key, None)
|
||||
|
||||
async def handle_401(self, server_name: str, failed_access_token: Optional[str] = None) -> bool:
|
||||
"""Handle a 401 from a tool call, deduplicated across concurrent callers.
|
||||
|
||||
True: a (possibly new) access token is available — caller should reconnect
|
||||
and retry. False: no recovery path — caller should surface a
|
||||
``needs_reauth`` error so the model stops hallucinating manual refreshes.
|
||||
Thundering-herd protection: N concurrent 401s with the same
|
||||
``failed_access_token`` fire one recovery attempt; the rest await its future.
|
||||
"""
|
||||
"""Handle a 401 from a tool call, deduplicated across concurrent callers. True: a
|
||||
(possibly new) access token is available — caller should reconnect and retry. False: no
|
||||
recovery path — caller should surface a ``needs_reauth`` error so the model stops
|
||||
hallucinating manual refreshes. Thundering-herd protection: N concurrent 401s with the
|
||||
same ``failed_access_token`` fire one recovery attempt; the rest await its future."""
|
||||
entry = self._entries.get(self._key(server_name))
|
||||
if entry is None or entry.provider is None:
|
||||
return False
|
||||
|
||||
key = failed_access_token or "<unknown>"
|
||||
loop = asyncio.get_running_loop()
|
||||
|
||||
async with entry.lock:
|
||||
pending = entry.pending_401.get(key)
|
||||
if pending is None:
|
||||
@@ -579,7 +500,6 @@ class MCPOAuthManager:
|
||||
task = asyncio.create_task(self._recover_401(server_name, entry, key, pending))
|
||||
self._inflight_tasks.add(task)
|
||||
task.add_done_callback(self._inflight_tasks.discard)
|
||||
|
||||
try:
|
||||
return await pending
|
||||
except Exception as exc: # pragma: no cover — defensive
|
||||
|
||||
@@ -23,68 +23,53 @@ def _jsonrpc_code(exc: BaseException):
|
||||
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).
|
||||
"""
|
||||
code = _jsonrpc_code(exc) or getattr(exc, "code", None)
|
||||
if code in (_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, _core._JSONRPC_METHOD_NOT_FOUND):
|
||||
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 (
|
||||
"unsupported protocol version" in msg
|
||||
or str(_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION) in msg
|
||||
or _is_method_not_found_error(exc)
|
||||
)
|
||||
return bool(msg) and any(marker in msg for marker in markers)
|
||||
|
||||
|
||||
def _handshake_rejected_as_modern(exc: BaseException) -> bool:
|
||||
"""True when a failed ``initialize`` signals a stateless-only (2026-07-28) server."""
|
||||
return _jsonrpc_matches(
|
||||
exc, _jsonrpc_code(exc) or getattr(exc, "code", None),
|
||||
(_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION, _core._JSONRPC_METHOD_NOT_FOUND),
|
||||
("unsupported protocol version", str(_JSONRPC_UNSUPPORTED_PROTOCOL_VERSION)),
|
||||
) or _is_method_not_found_error(exc)
|
||||
|
||||
|
||||
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
|
||||
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.
|
||||
"""
|
||||
if _jsonrpc_code(exc) == _core._JSONRPC_METHOD_NOT_FOUND:
|
||||
return True
|
||||
msg = str(exc).lower()
|
||||
return bool(msg) and (
|
||||
str(_core._JSONRPC_METHOD_NOT_FOUND) in msg
|
||||
or "method not found" in msg
|
||||
or "unknown method" in msg
|
||||
or "not found: ping" in msg
|
||||
"""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."""
|
||||
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"),
|
||||
)
|
||||
|
||||
|
||||
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 on every attempt."""
|
||||
|
||||
|
||||
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 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."""
|
||||
|
||||
|
||||
def _unwrap_exception_group(exc: BaseException) -> BaseException:
|
||||
"""Extract the root-cause leaf from anyio ``(Base)ExceptionGroup`` wrappers.
|
||||
|
||||
Group ``str()`` is opaque ("unhandled errors in a TaskGroup"), 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.
|
||||
"""
|
||||
"""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."""
|
||||
while isinstance(exc, BaseExceptionGroup) and exc.exceptions:
|
||||
fatal, _rest = exc.split((KeyboardInterrupt, SystemExit))
|
||||
if fatal is not None:
|
||||
@@ -104,19 +89,15 @@ def _contains_only_cancellation(exc: BaseException) -> bool:
|
||||
|
||||
|
||||
def _classify_mcp_failure(exc: BaseException) -> str:
|
||||
"""Classify a connection failure as ``'permanent'`` or ``'transient'``.
|
||||
|
||||
Permanent (deterministic — ``run()`` parks immediately instead of burning the
|
||||
retry ladder): auth 401/403, NonMcpEndpointError, InvalidMcpUrlError, missing
|
||||
stdio command (FileNotFoundError / ENOENT). Everything else keeps backoff retry.
|
||||
"""
|
||||
"""``'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)."""
|
||||
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 in this environment).
|
||||
# 401/403 HTTPStatusError that _is_auth_error's type-gate missed (auth types not importable here).
|
||||
or _response_status(root) in (401, 403)
|
||||
)
|
||||
return "permanent" if permanent else "transient"
|
||||
@@ -128,11 +109,9 @@ def _response_status(exc: BaseException):
|
||||
|
||||
|
||||
def _validate_remote_mcp_url(server_name: str, url: Any) -> str:
|
||||
"""Return 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) remote MCP URL. 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}")
|
||||
|
||||
@@ -149,20 +128,16 @@ def _validate_remote_mcp_url(server_name: str, url: Any) -> str:
|
||||
raise _bad(f"scheme must be http or https, got {parsed.scheme!r} ({stripped!r})")
|
||||
if not parsed.netloc:
|
||||
raise _bad(f"missing host ({stripped!r})")
|
||||
# ``urlparse`` accepts ``http://:8080`` (empty host, explicit port) — reject it.
|
||||
if not parsed.hostname:
|
||||
if not parsed.hostname: # ``urlparse`` accepts ``http://:8080`` (empty host, explicit port)
|
||||
raise _bad(f"missing hostname ({stripped!r})")
|
||||
return stripped
|
||||
|
||||
|
||||
def _resolve_client_cert(server_name: str, config: dict):
|
||||
"""Resolve ``client_cert`` / ``client_key`` into 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 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."""
|
||||
raw_cert = config.get("client_cert")
|
||||
raw_key = config.get("client_key")
|
||||
if raw_cert is None and raw_key is None:
|
||||
@@ -179,20 +154,16 @@ def _resolve_client_cert(server_name: str, config: dict):
|
||||
|
||||
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"
|
||||
)
|
||||
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
|
||||
password = raw_cert[2]
|
||||
if not isinstance(password, str):
|
||||
if not isinstance(raw_cert[2], str):
|
||||
raise ValueError(f"{prefix}client_cert[2] (key passphrase) must be a string")
|
||||
return (*pair, password)
|
||||
|
||||
return (*pair, raw_cert[2])
|
||||
cert_path = _expand(raw_cert, "client_cert")
|
||||
if raw_key is not None:
|
||||
return (cert_path, _expand(raw_key, "client_key"))
|
||||
@@ -200,13 +171,10 @@ def _resolve_client_cert(server_name: str, config: dict):
|
||||
|
||||
|
||||
def _resolve_identity_header(server_name: str, config: dict):
|
||||
"""Resolve the optional per-server ``identity_header`` config.
|
||||
|
||||
Shape: ``{name: "X-User-Id", value_from: "static"|"profile", value: "..."}``
|
||||
(``value`` required for static). Returns ``(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.
|
||||
"""
|
||||
"""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."""
|
||||
raw = config.get("identity_header")
|
||||
if raw is None:
|
||||
return None
|
||||
@@ -233,38 +201,28 @@ def _resolve_identity_header(server_name: str, config: dict):
|
||||
|
||||
|
||||
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.
|
||||
"""
|
||||
"""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:
|
||||
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,
|
||||
)
|
||||
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
|
||||
return headers
|
||||
|
||||
|
||||
def _make_redirect_header_stripper(
|
||||
original_url,
|
||||
*,
|
||||
strict: bool = False,
|
||||
configured_header_names: "set[str] | frozenset[str]" = frozenset(),
|
||||
):
|
||||
"""Build an httpx response hook that guards 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.
|
||||
"""
|
||||
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."""
|
||||
origin = (original_url.scheme, original_url.host, original_url.port)
|
||||
|
||||
async def _strip_on_cross_origin_redirect(response):
|
||||
@@ -284,45 +242,41 @@ def _make_redirect_header_stripper(
|
||||
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 _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)]
|
||||
|
||||
|
||||
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:
|
||||
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)
|
||||
return next(filter(None, map(_find_missing, _exc_causes(current))), None)
|
||||
return next(filter(None, map(_find_missing, _exc_children(current))), None)
|
||||
|
||||
def _flatten_messages(current: BaseException) -> List[str]:
|
||||
nested = getattr(current, "exceptions", None)
|
||||
if nested:
|
||||
return [m for child in nested for m in _flatten_messages(child)]
|
||||
text = str(current).strip()
|
||||
# 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 nested_exc in _exc_causes(current):
|
||||
messages.extend(_flatten_messages(nested_exc))
|
||||
for child in _exc_children(current):
|
||||
messages.extend(_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)"
|
||||
)
|
||||
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]))
|
||||
@@ -343,31 +297,21 @@ 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, so both go in the 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."""
|
||||
global _HTTP_STATUS_ERROR_TYPES
|
||||
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)
|
||||
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)
|
||||
return _HTTP_STATUS_ERROR_TYPES
|
||||
|
||||
|
||||
def _get_auth_error_types() -> tuple:
|
||||
"""Cached tuple of 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 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`."""
|
||||
global _AUTH_ERROR_TYPES
|
||||
if not _AUTH_ERROR_TYPES:
|
||||
_AUTH_ERROR_TYPES = tuple(
|
||||
@@ -380,11 +324,8 @@ def _get_auth_error_types() -> tuple:
|
||||
|
||||
|
||||
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.
|
||||
"""
|
||||
"""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):
|
||||
return False
|
||||
@@ -394,39 +335,33 @@ def _is_auth_error(exc: BaseException) -> bool:
|
||||
return 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 server-side transport session expired / was GC'd.
|
||||
# The OAuth token is still valid — only the transport needs rebuilding.
|
||||
_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",
|
||||
)
|
||||
|
||||
|
||||
# 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; the budget
|
||||
# bounds pathological acyclic graphs. Kept well above ``sys.getrecursionlimit()`` so deep
|
||||
# task-group nesting is still 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.
|
||||
"""
|
||||
"""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.
|
||||
transport_error_types = tuple(
|
||||
_optional_types("anyio", "BrokenResourceError", "ClosedResourceError", "EndOfStream")
|
||||
)
|
||||
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.
|
||||
# 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
|
||||
@@ -437,19 +372,13 @@ def _is_session_expired_error(exc: BaseException) -> bool:
|
||||
continue
|
||||
seen.add(id(current))
|
||||
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 and servers: match a narrow allow-list of stable
|
||||
# substrings, not exception type, to avoid false positives.
|
||||
msg = str(current).lower()
|
||||
if isinstance(current, transport_error_types) or (
|
||||
msg and any(marker in msg for marker in _SESSION_EXPIRED_MARKERS)
|
||||
):
|
||||
if isinstance(current, transport_error_types) or (msg and any(marker in msg for marker in _SESSION_EXPIRED_MARKERS)):
|
||||
transport_error_found = True
|
||||
|
||||
stack.extend(getattr(current, "exceptions", ()))
|
||||
stack.append(getattr(current, "__cause__", None))
|
||||
stack.append(getattr(current, "__context__", None))
|
||||
|
||||
stack.extend((getattr(current, "__cause__", None), getattr(current, "__context__", None)))
|
||||
return transport_error_found
|
||||
|
||||
@@ -115,9 +115,7 @@ def _lookup_reconnectable_server(server_name: str, require_loop: bool = False):
|
||||
With *require_loop*, also None unless the MCP loop is running (nothing to wait on)."""
|
||||
with _core._lock:
|
||||
srv = _core._servers.get(server_name)
|
||||
if srv is None or not hasattr(srv, "_reconnect_event"):
|
||||
return None
|
||||
if require_loop and not _mcp_loop_running():
|
||||
if srv is None or not hasattr(srv, "_reconnect_event") or (require_loop and not _mcp_loop_running()):
|
||||
return None
|
||||
return srv
|
||||
|
||||
@@ -481,12 +479,10 @@ def _make_utility_handler(server_name: str, tool_timeout: float, op: str, log_la
|
||||
result = await rpc(server.session, args)
|
||||
return json.dumps(render(result, server_name), ensure_ascii=False)
|
||||
|
||||
def _on_failure(exc):
|
||||
logger.error("MCP %s/%s failed: %s", server_name, log_label, exc)
|
||||
|
||||
return _invoke_with_recovery(
|
||||
server_name, lambda: _core._run_on_mcp_loop(_call, timeout=tool_timeout), op,
|
||||
(_handle_auth_error_and_retry, _handle_session_expired_and_retry), _on_failure,
|
||||
(_handle_auth_error_and_retry, _handle_session_expired_and_retry),
|
||||
lambda exc: logger.error("MCP %s/%s failed: %s", server_name, log_label, exc),
|
||||
)
|
||||
|
||||
return _handler
|
||||
|
||||
@@ -16,8 +16,8 @@ _KEEPALIVE_RPC_TIMEOUT = 30.0
|
||||
|
||||
|
||||
def _stdio_children_dead_impl(pids, is_http: bool) -> bool:
|
||||
"""True when every pid has exited. Best-effort: False (unknown → don't fail
|
||||
fast) for HTTP, no captured PIDs, missing psutil, or a failed probe."""
|
||||
"""True when every pid has exited. Best-effort: False (unknown → don't fail fast) for HTTP,
|
||||
no captured PIDs, missing psutil, or a failed probe."""
|
||||
if not pids or is_http:
|
||||
return False
|
||||
try:
|
||||
@@ -25,9 +25,8 @@ def _stdio_children_dead_impl(pids, is_http: bool) -> bool:
|
||||
except ImportError:
|
||||
return False
|
||||
for pid in pids:
|
||||
# pid_exists handles Windows without signal-permission noise.
|
||||
try:
|
||||
if psutil.pid_exists(pid):
|
||||
if psutil.pid_exists(pid): # handles Windows without signal-permission noise
|
||||
return False
|
||||
except Exception:
|
||||
return False
|
||||
@@ -40,11 +39,10 @@ class MCPServerHealthMixin:
|
||||
__slots__ = ()
|
||||
|
||||
def _is_http(self) -> bool:
|
||||
"""Check if this server uses HTTP transport."""
|
||||
return "url" in self._config
|
||||
|
||||
def _is_recycled_stdio(self) -> bool:
|
||||
"""Return True when a stdio server was intentionally recycled."""
|
||||
"""True when a stdio server was intentionally recycled."""
|
||||
return not self._is_http() and self._recycled_reason is not None
|
||||
|
||||
def mark_tool_call(self) -> None:
|
||||
@@ -52,16 +50,14 @@ class MCPServerHealthMixin:
|
||||
self._last_tool_call_at = time.monotonic()
|
||||
|
||||
def _mark_lifecycle_started(self) -> None:
|
||||
now = time.monotonic()
|
||||
self._lifecycle_started_at = now
|
||||
self._last_tool_call_at = now
|
||||
self._lifecycle_started_at = self._last_tool_call_at = time.monotonic()
|
||||
self._recycled_reason = None
|
||||
|
||||
# ------------------------------------------------------- stdio recycling
|
||||
|
||||
def _stdio_recycle_deadlines(self):
|
||||
"""``[(deadline, reason), ...]`` for the configured lifetime/idle limits;
|
||||
empty for HTTP servers or while an RPC holds the lock."""
|
||||
"""``[(deadline, reason), ...]`` for the configured lifetime/idle limits; empty for HTTP
|
||||
servers or while an RPC holds the lock."""
|
||||
if self._is_http() or self._rpc_lock.locked():
|
||||
return []
|
||||
deadlines = []
|
||||
@@ -72,15 +68,12 @@ class MCPServerHealthMixin:
|
||||
return deadlines
|
||||
|
||||
def _stdio_recycle_reason(self, now: Optional[float] = None) -> Optional[str]:
|
||||
"""Return the stdio recycle reason if idle/age limits have elapsed (lifetime wins)."""
|
||||
"""The stdio recycle reason if idle/age limits have elapsed (lifetime wins), else None."""
|
||||
now = time.monotonic() if now is None else now
|
||||
for deadline, reason in self._stdio_recycle_deadlines():
|
||||
if now >= deadline:
|
||||
return reason
|
||||
return None
|
||||
return next((reason for deadline, reason in self._stdio_recycle_deadlines() if now >= deadline), None)
|
||||
|
||||
def _next_stdio_recycle_deadline(self) -> Optional[float]:
|
||||
"""Return the next monotonic recycle deadline for stdio, if any."""
|
||||
"""The next monotonic recycle deadline for stdio, if any."""
|
||||
deadlines = self._stdio_recycle_deadlines()
|
||||
return min(d for d, _ in deadlines) if deadlines else None
|
||||
|
||||
@@ -106,13 +99,11 @@ class MCPServerHealthMixin:
|
||||
return task
|
||||
|
||||
def _make_logging_callback(self):
|
||||
"""Build a ``logging_callback`` that forwards server ``notifications/message``
|
||||
into Hermes logging tagged with the server name (the SDK default drops them)."""
|
||||
"""``logging_callback`` forwarding server ``notifications/message`` into Hermes logging
|
||||
tagged with the server name (the SDK default drops them)."""
|
||||
async def _on_log(params):
|
||||
try:
|
||||
level = _core._MCP_LOG_LEVEL_MAP.get(
|
||||
str(getattr(params, "level", "info")).lower(), logging.INFO,
|
||||
)
|
||||
level = _core._MCP_LOG_LEVEL_MAP.get(str(getattr(params, "level", "info")).lower(), logging.INFO)
|
||||
data = getattr(params, "data", None)
|
||||
if not isinstance(data, str):
|
||||
try:
|
||||
@@ -130,31 +121,27 @@ class MCPServerHealthMixin:
|
||||
return _on_log
|
||||
|
||||
def _make_message_handler(self):
|
||||
"""Build a ``message_handler`` for ``ClientSession``: only
|
||||
``ToolListChangedNotification`` triggers a refresh; prompt/resource changes are logged."""
|
||||
"""``message_handler`` for ``ClientSession``: only ``ToolListChangedNotification`` triggers
|
||||
a refresh; prompt/resource changes are logged."""
|
||||
async def _handler(message):
|
||||
try:
|
||||
if isinstance(message, Exception):
|
||||
logger.debug("MCP message handler (%s): exception: %s", self.name, message)
|
||||
return
|
||||
if _core._MCP_NOTIFICATION_TYPES and isinstance(message, _core.ServerNotification):
|
||||
# mcp 2.0 made ServerNotification a plain union (payload IS the
|
||||
# message) instead of a RootModel (payload under ``.root``).
|
||||
# ``isinstance`` accepts both; only the unwrap differs — without
|
||||
# it ``.root`` raises into the catch-all and refreshes stop.
|
||||
# mcp 2.0 made ServerNotification a plain union (payload IS the message)
|
||||
# instead of a RootModel (payload under ``.root``). ``isinstance`` accepts
|
||||
# both; only the unwrap differs — without it ``.root`` raises into the
|
||||
# catch-all and refreshes stop.
|
||||
match getattr(message, "root", message):
|
||||
case _core.ToolListChangedNotification():
|
||||
logger.info(
|
||||
"MCP server '%s': received tools/list_changed notification",
|
||||
self.name,
|
||||
)
|
||||
# Refresh in a separate task: some servers emit
|
||||
# list_changed right after initialize while another
|
||||
# request is in flight, and refreshing synchronously
|
||||
# inside the handler can wedge the stdio JSON-RPC stream.
|
||||
logger.info("MCP server '%s': received tools/list_changed notification", self.name)
|
||||
# Refresh in a separate task: some servers emit list_changed right
|
||||
# after initialize while another request is in flight, and refreshing
|
||||
# synchronously inside the handler can wedge the stdio JSON-RPC stream.
|
||||
self._schedule_tools_refresh()
|
||||
# Yield one tick so short-lived notification contexts
|
||||
# (and tests) can observe the scheduled refresh.
|
||||
# Yield one tick so short-lived notification contexts (and tests)
|
||||
# can observe the scheduled refresh.
|
||||
await asyncio.sleep(0)
|
||||
case _core.PromptListChangedNotification():
|
||||
logger.debug("MCP server '%s': prompts/list_changed (ignored)", self.name)
|
||||
@@ -167,8 +154,8 @@ class MCPServerHealthMixin:
|
||||
return _handler
|
||||
|
||||
def _deregister_owned(self, tool_names: Iterable[str]) -> None:
|
||||
"""Deregister *tool_names* that this server's toolset still owns.
|
||||
Never removes a colliding name currently owned by another server."""
|
||||
"""Deregister *tool_names* that this server's toolset still owns. Never removes a
|
||||
colliding name currently owned by another server."""
|
||||
from tools.registry import registry
|
||||
|
||||
toolset_name = f"mcp-{self.name}"
|
||||
@@ -179,69 +166,44 @@ class MCPServerHealthMixin:
|
||||
_forget_mcp_tool_server(tool_name)
|
||||
|
||||
async def _refresh_tools(self):
|
||||
"""Re-fetch tools on ``tools/list_changed`` and update the registry.
|
||||
|
||||
The lock serializes rapid-fire notifications. After the list_tools
|
||||
``await``, all mutations are synchronous — atomic on the event loop.
|
||||
"""
|
||||
"""Re-fetch tools on ``tools/list_changed`` and update the registry. The lock serializes
|
||||
rapid-fire notifications; after the list_tools ``await`` all mutations are synchronous —
|
||||
atomic on the event loop."""
|
||||
if not self._advertises_tools():
|
||||
# Shouldn't happen, but tools/list would raise MCPError(-32601).
|
||||
return
|
||||
|
||||
return # tools/list would raise MCPError(-32601)
|
||||
async with self._refresh_lock:
|
||||
old_tool_names = set(self._registered_tool_names)
|
||||
|
||||
# 1. Fetch the current tool list (follow nextCursor).
|
||||
async with self._rpc_lock:
|
||||
new_mcp_tools = await _core._paginate_full_list(self.session.list_tools, "tools", self.name)
|
||||
|
||||
# 2. Remove only stale names first — no nuke-and-repave: live agent
|
||||
# turns may hold tool-call IDs pointing at existing handlers, and
|
||||
# in-place replacement avoids transient "tool not connected" races.
|
||||
self._deregister_owned(old_tool_names - {
|
||||
mcp_prefixed_tool_name(self.name, tool.name) for tool in new_mcp_tools
|
||||
})
|
||||
|
||||
# 3. Re-register; the helper may skip names ambiguous after normalization.
|
||||
# Remove only stale names first — no nuke-and-repave: live agent turns may hold
|
||||
# tool-call IDs pointing at existing handlers, and in-place replacement avoids
|
||||
# transient "tool not connected" races.
|
||||
self._deregister_owned(old_tool_names - {mcp_prefixed_tool_name(self.name, tool.name) for tool in new_mcp_tools})
|
||||
# Re-register; the helper may skip names ambiguous after normalization. A raw name
|
||||
# can become ambiguous without changing its normalized name, so the pre-pass misses
|
||||
# it: drop any old entry the final collision-checked registration no longer owns.
|
||||
self._tools = new_mcp_tools
|
||||
registered_names = _core._register_server_tools(self.name, self, self._config)
|
||||
# A raw name can become ambiguous without changing its normalized
|
||||
# name, so the pre-pass misses it: drop any old entry the final
|
||||
# collision-checked registration no longer owns.
|
||||
self._deregister_owned(old_tool_names - set(registered_names))
|
||||
self._registered_tool_names = registered_names
|
||||
|
||||
# 4. Log what changed (user-visible).
|
||||
new_tool_names = set(self._registered_tool_names)
|
||||
added = new_tool_names - old_tool_names
|
||||
removed = old_tool_names - new_tool_names
|
||||
changes = []
|
||||
if added:
|
||||
changes.append(f"added: {', '.join(sorted(added))}")
|
||||
if removed:
|
||||
changes.append(f"removed: {', '.join(sorted(removed))}")
|
||||
# Log what changed (user-visible).
|
||||
new_tool_names = set(registered_names)
|
||||
changes = [f"{label}: {', '.join(sorted(names))}" for label, names in
|
||||
(("added", new_tool_names - old_tool_names), ("removed", old_tool_names - new_tool_names)) if names]
|
||||
if changes:
|
||||
logger.warning(
|
||||
"MCP server '%s': tools changed dynamically — %s. "
|
||||
"Verify these changes are expected.",
|
||||
self.name, "; ".join(changes),
|
||||
)
|
||||
logger.warning("MCP server '%s': tools changed dynamically — %s. "
|
||||
"Verify these changes are expected.", self.name, "; ".join(changes))
|
||||
else:
|
||||
logger.info(
|
||||
"MCP server '%s': dynamically refreshed %d tool(s) (no changes)",
|
||||
self.name, len(self._registered_tool_names),
|
||||
)
|
||||
logger.info("MCP server '%s': dynamically refreshed %d tool(s) (no changes)",
|
||||
self.name, len(self._registered_tool_names))
|
||||
|
||||
# ------------------------------------------------------ keepalive / health
|
||||
|
||||
async def _keepalive_probe(self) -> None:
|
||||
"""Exercise the session; raise on a genuine connection failure.
|
||||
|
||||
``ping`` first (cheap, OPTIONAL utility). On -32601 latch
|
||||
``_ping_unsupported`` and fall back to ``list_tools`` when the server
|
||||
advertises tools; otherwise the -32601 propagates (no liveness primitive
|
||||
left). The latch resets on each fresh transport connection.
|
||||
"""
|
||||
"""Exercise the session; raise on a genuine connection failure. ``ping`` first (cheap,
|
||||
OPTIONAL utility). On -32601 latch ``_ping_unsupported`` and fall back to ``list_tools``
|
||||
when the server advertises tools; otherwise the -32601 propagates (no liveness primitive
|
||||
left). The latch resets on each fresh transport connection."""
|
||||
if not self._ping_unsupported:
|
||||
try:
|
||||
await asyncio.wait_for(self.session.send_ping(), timeout=_KEEPALIVE_RPC_TIMEOUT)
|
||||
@@ -252,74 +214,55 @@ class MCPServerHealthMixin:
|
||||
if not self._advertises_tools():
|
||||
raise
|
||||
self._ping_unsupported = True
|
||||
logger.info(
|
||||
"MCP server '%s': does not implement the optional 'ping' utility (-32601); "
|
||||
"using 'list_tools' for keepalive on this connection.",
|
||||
self.name,
|
||||
)
|
||||
logger.info("MCP server '%s': does not implement the optional 'ping' utility (-32601); "
|
||||
"using 'list_tools' for keepalive on this connection.", self.name)
|
||||
elif isinstance(exc, (TimeoutError, asyncio.TimeoutError)) and self._advertises_tools():
|
||||
# A server that silently drops ping looks like a dead transport.
|
||||
# Confirm with list_tools before declaring it dead; if that
|
||||
# also fails, propagate the original failure.
|
||||
# A server that silently drops ping looks like a dead transport. Confirm with
|
||||
# list_tools before declaring it dead; if that also fails, propagate the
|
||||
# original failure.
|
||||
try:
|
||||
await asyncio.wait_for(self.session.list_tools(), timeout=_KEEPALIVE_RPC_TIMEOUT)
|
||||
except Exception:
|
||||
raise exc from None
|
||||
# Transport alive; latch so later keepalives skip the 30s wait.
|
||||
self._ping_unsupported = True
|
||||
logger.info(
|
||||
"MCP server '%s': ping timed out but list_tools succeeded — server "
|
||||
"silently drops ping; using 'list_tools' for keepalive on this connection.",
|
||||
self.name,
|
||||
)
|
||||
logger.info("MCP server '%s': ping timed out but list_tools succeeded — server "
|
||||
"silently drops ping; using 'list_tools' for keepalive on this connection.", self.name)
|
||||
return
|
||||
else:
|
||||
# Closed transport, expired session, etc. — real failure.
|
||||
raise
|
||||
|
||||
raise # closed transport, expired session, etc. — real failure
|
||||
# Fallback probe for servers without ping support.
|
||||
await asyncio.wait_for(self.session.list_tools(), timeout=_KEEPALIVE_RPC_TIMEOUT)
|
||||
|
||||
def _mark_session_proven(self) -> None:
|
||||
"""Record that the session demonstrated real health (keepalive or tool-call success).
|
||||
|
||||
Only then is the reconnect budget cleared: a handshake that drops moments
|
||||
later must keep consuming ``_reconnect_retries`` so a flapping transport
|
||||
still reaches the park instead of respawning forever.
|
||||
"""
|
||||
Only then is the reconnect budget cleared: a handshake that drops moments later must
|
||||
keep consuming ``_reconnect_retries`` so a flapping transport still reaches the park
|
||||
instead of respawning forever."""
|
||||
if self._session_proven:
|
||||
return
|
||||
self._session_proven = True
|
||||
self._reconnect_retries = 0
|
||||
if self._was_parked:
|
||||
self._was_parked = False
|
||||
logger.warning(
|
||||
"MCP server '%s': revived — session healthy again after "
|
||||
"parking (state: parked → connected)",
|
||||
self.name,
|
||||
)
|
||||
# A proven fresh transport clears the one-time permanent-failure
|
||||
# grace and any race bookkeeping.
|
||||
logger.warning("MCP server '%s': revived — session healthy again after "
|
||||
"parking (state: parked → connected)", self.name)
|
||||
# A proven fresh transport clears the one-time permanent-failure grace and any race bookkeeping.
|
||||
self._permanent_grace_used = False
|
||||
self._teardown_race = False
|
||||
|
||||
def mark_suspect(self, reason: str) -> None:
|
||||
"""Latch a suspicion (no I/O). The NEXT call verifies via
|
||||
:meth:`ensure_healthy` and recycles the transport if the probe fails."""
|
||||
"""Latch a suspicion (no I/O). The NEXT call verifies via :meth:`ensure_healthy` and
|
||||
recycles the transport if the probe fails."""
|
||||
if self._suspect_reason is None and reason:
|
||||
logger.warning(
|
||||
"MCP server '%s': connection marked suspect (%s); next call will health-check it",
|
||||
self.name, reason,
|
||||
)
|
||||
logger.warning("MCP server '%s': connection marked suspect (%s); next call will health-check it",
|
||||
self.name, reason)
|
||||
self._suspect_reason = reason or None
|
||||
|
||||
async def ensure_healthy(self, timeout: float = 5.0) -> bool:
|
||||
"""Verify a suspect connection before reuse; recycle if dead.
|
||||
|
||||
True when healthy (suspicion cleared). On failure requests a reconnect,
|
||||
drops the stale session so the caller's no-session path takes over, and
|
||||
returns False. Never raises.
|
||||
"""
|
||||
"""Verify a suspect connection before reuse; recycle if dead. True when healthy (suspicion
|
||||
cleared). On failure requests a reconnect, drops the stale session so the caller's
|
||||
no-session path takes over, and returns False. Never raises."""
|
||||
reason = self._suspect_reason
|
||||
if not reason:
|
||||
return True
|
||||
@@ -332,34 +275,27 @@ class MCPServerHealthMixin:
|
||||
await asyncio.wait_for(self._keepalive_probe(), timeout=timeout)
|
||||
except Exception as exc:
|
||||
root = _unwrap_exception_group(exc)
|
||||
logger.warning(
|
||||
"MCP server '%s': suspect connection (%s) failed health check (%s: %s) — "
|
||||
"requesting reconnect (state: suspect → degraded)",
|
||||
self.name, reason, type(root).__name__, root,
|
||||
)
|
||||
logger.warning("MCP server '%s': suspect connection (%s) failed health check (%s: %s) — "
|
||||
"requesting reconnect (state: suspect → degraded)",
|
||||
self.name, reason, type(root).__name__, root)
|
||||
self._suspect_reason = None
|
||||
self.mark_suspect(f"health check failed after {reason}")
|
||||
self.session = None
|
||||
self._ready.clear()
|
||||
self._reconnect_event.set()
|
||||
return False
|
||||
logger.info(
|
||||
"MCP server '%s': suspect connection passed health check (%s) — clearing suspicion",
|
||||
self.name, reason,
|
||||
)
|
||||
logger.info("MCP server '%s': suspect connection passed health check (%s) — clearing suspicion",
|
||||
self.name, reason)
|
||||
self._suspect_reason = None
|
||||
self._mark_session_proven()
|
||||
return True
|
||||
|
||||
def _fail_inflight_calls(self, reason: str) -> None:
|
||||
"""Cancel every in-flight RPC on this connection.
|
||||
|
||||
Called from lifecycle exits BEFORE the transport unwinds: the SDK does
|
||||
not always fail pending requests when streams close, so a call would
|
||||
otherwise wait out the full tool timeout. Cancelling anything flags
|
||||
``_teardown_race`` so run() treats the next reconnect as recovery
|
||||
rather than charging the rapid-drop budget.
|
||||
"""
|
||||
"""Cancel every in-flight RPC on this connection. Called from lifecycle exits BEFORE the
|
||||
transport unwinds: the SDK does not always fail pending requests when streams close, so
|
||||
a call would otherwise wait out the full tool timeout. Cancelling anything flags
|
||||
``_teardown_race`` so run() treats the next reconnect as recovery rather than charging
|
||||
the rapid-drop budget."""
|
||||
victims = [t for t in self._inflight_tasks if not t.done()]
|
||||
if not victims:
|
||||
return
|
||||
@@ -374,7 +310,7 @@ class MCPServerHealthMixin:
|
||||
return _stdio_children_dead_impl(getattr(self, "_stdio_child_pids", None), self._is_http())
|
||||
|
||||
async def _watch_stdio_children(self) -> None:
|
||||
"""Poll child liveness while a stdio RPC is in flight; resolves when a
|
||||
tracked child dies so the caller cancels the RPC instead of waiting out the timeout."""
|
||||
"""Poll child liveness while a stdio RPC is in flight; resolves when a tracked child dies
|
||||
so the caller cancels the RPC instead of waiting out the timeout."""
|
||||
while not self._stdio_children_dead():
|
||||
await asyncio.sleep(0.25)
|
||||
|
||||
Reference in New Issue
Block a user