Merge branch 'simp/r2a-mcp-oauth' into simp/r2a-mcp

This commit is contained in:
Teknium
2026-09-02 17:14:56 -07:00
2 changed files with 385 additions and 575 deletions
+269 -379
View File
File diff suppressed because it is too large Load Diff
+116 -196
View File
@@ -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