diff --git a/tools/mcp_oauth.py b/tools/mcp_oauth.py index 8ba9e6ae3f..d3c6ba8d3d 100644 --- a/tools/mcp_oauth.py +++ b/tools/mcp_oauth.py @@ -1,18 +1,17 @@ #!/usr/bin/env python3 """MCP OAuth 2.1 client support: browser authorization-code flow with PKCE. -The MCP SDK's ``OAuthClientProvider`` (an ``httpx.Auth``) does discovery, client -identification, PKCE, exchange, refresh and step-up; this module supplies -``HermesTokenStorage`` (on-disk persistence), the localhost callback listener, -and ``build_oauth_auth()`` (legacy entry point). Per the MCP 2026-07-28 spec the -client_id is Hermes' published Client ID Metadata Document URL (CIMD) when the -server advertises support, else RFC 7591 dynamic client registration (DCR). +The SDK's ``OAuthClientProvider`` (an ``httpx.Auth``) does discovery, client identification, +PKCE, exchange, refresh and step-up; this module supplies ``HermesTokenStorage`` (on-disk +persistence), the localhost callback listener and ``build_oauth_auth()`` (legacy entry point). +The client_id is Hermes' published Client ID Metadata Document URL (CIMD) when the server +advertises support, else RFC 7591 dynamic client registration (DCR). -``mcp_servers..oauth`` keys (all optional): ``client_id`` (skip DCR), -``client_secret`` (confidential clients), ``scope``, ``redirect_port`` (0 = -auto), ``redirect_uri`` (proxy callback), ``redirect_host`` (loopback hostname, -WAF-safe), ``client_name`` (default "Hermes Agent"), ``client_metadata_url`` -(self-hosted CIMD), ``cimd: false`` (force DCR), ``user_agent``, ``timeout``. +``mcp_servers..oauth`` keys (all optional): ``client_id`` (skip DCR), ``client_secret`` +(confidential clients), ``scope``, ``redirect_port`` (0 = auto), ``redirect_uri`` (proxy +callback), ``redirect_host`` (loopback hostname, WAF-safe), ``client_name`` (default "Hermes +Agent"), ``client_metadata_url`` (self-hosted CIMD), ``cimd: false`` (force DCR), +``user_agent``, ``timeout``. """ import asyncio @@ -40,11 +39,10 @@ from tools.mcp_dashboard_oauth import contextvar_set as _contextvar_set, get_das logger = logging.getLogger(__name__) -# Lazy SDK imports: availability is detected WITHOUT importing mcp (~170 ms). -# Module-level names stay None placeholders so tests can patch.object them; -# _ensure_sdk_loaded() binds the real classes on first use. _SDK_CLASSES caches -# them so a test that patches a name and restores it to None afterwards doesn't -# strand the module in a broken state. +# Lazy SDK imports: availability is detected WITHOUT importing mcp (~170 ms). Module-level +# names stay None placeholders so tests can patch.object them; _ensure_sdk_loaded() binds the +# real classes on first use and _SDK_CLASSES caches them so a test that patches a name and +# restores it to None afterwards doesn't strand the module in a broken state. _OAUTH_AVAILABLE = _importlib_util.find_spec("mcp") is not None if not _OAUTH_AVAILABLE: logger.debug("MCP OAuth types not available -- OAuth MCP auth disabled") @@ -60,8 +58,8 @@ _SDK_LOAD_FAILED = False def _ensure_sdk_loaded() -> bool: - """Bind the SDK OAuth classes into module globals; True when available. - Only names currently ``None`` are (re)bound, so test patches stay intact.""" + """Bind the SDK OAuth classes into module globals; True when available. Only names + currently ``None`` are (re)bound, so test patches stay intact.""" global _SDK_LOAD_FAILED, _OAUTH_AVAILABLE if _SDK_LOAD_FAILED: return False @@ -103,26 +101,26 @@ class OAuthNonInteractiveError(RuntimeError): """Raised when OAuth requires browser interaction in a non-interactive env.""" -# Port used by the most recent callback-port resolution. Legacy global; the -# per-flow closures are the real mechanism (concurrent flows must not share it). +# Port used by the most recent callback-port resolution. Legacy global; the per-flow +# closures are the real mechanism (concurrent flows must not share it). _oauth_port: int | None = None -# Interactivity gates for OAuth stdin prompts. ContextVars (NOT threading.local): -# background MCP discovery sets them on the discovery thread while connect+OAuth -# runs on the `mcp-event-loop` thread via run_coroutine_threadsafe, which copies -# the calling context into the coroutine — a threading.local would not cross. +# Interactivity gates for OAuth stdin prompts. ContextVars, NOT threading.local: background +# MCP discovery sets them on the discovery thread while connect+OAuth runs on the +# `mcp-event-loop` thread via run_coroutine_threadsafe, which copies the calling context +# into the coroutine — a threading.local would not cross. _oauth_interactive_enabled = contextvars.ContextVar("_oauth_interactive_enabled", default=True) -# Forces _is_interactive() past the stdin-TTY check for GUI-driven flows -# (dashboard/desktop REST): the browser + callback server do the work and the -# stdin paste fallback degrades harmlessly (EOF swallowed). Suppression wins — -# background discovery must never start a browser flow. +# Forces _is_interactive() past the stdin-TTY check for GUI-driven flows (dashboard/desktop +# REST): the browser + callback server do the work and the stdin paste fallback degrades +# harmlessly (EOF swallowed). Suppression wins — background discovery must never start a +# browser flow. _oauth_interactive_forced = contextvars.ContextVar("_oauth_interactive_forced", default=False) # Skip tokens accepted at the paste prompt — exit OAuth without auth. _SKIP_TOKENS = frozenset({"skip", "cancel", "s", "n", "no", "q", "quit"}) # Written to result["error"] on stdin skip; the waiter maps it to -# OAuthNonInteractiveError("user_skipped") so MCP setup treats it as a -# non-fatal "continue without this server". +# OAuthNonInteractiveError("user_skipped") so MCP setup treats it as a non-fatal +# "continue without this server". _USER_SKIPPED_SENTINEL = "__hermes_user_skipped__" @@ -139,29 +137,24 @@ def _safe_filename(name: str) -> str: # -- Callback-port reservation --------------------------------------------- -# Bound-but-not-listening sockets for pending callback flows, keyed by port. -# Holding the socket from port selection until the waiter adopts it closes the -# TOCTOU window where another process grabs the port in between. Bounded FIFO -# so reconnect loops cannot leak fds. +# Bound-but-not-listening sockets for pending callback flows, keyed by port. Holding the +# socket from port selection until the waiter adopts it closes the TOCTOU window where +# another process grabs the port in between. Bounded FIFO so reconnect loops cannot leak fds. _reserved_sockets: "dict[int, socket.socket]" = {} _MAX_RESERVED_SOCKETS = 8 def _park_reserved_socket(port: int, sock: socket.socket) -> None: - """Hold *sock* bound to *port* until the callback waiter adopts it. - - Pinned CIMD sockets are never evicted: the published document only declares - the pinned ports, so losing one mid-flow would reopen the exact race the - parking prevents. The FIFO cap applies to ephemeral ports only. - """ + """Hold *sock* bound to *port* until the callback waiter adopts it. Pinned CIMD sockets + are never evicted: the published document only declares the pinned ports, so losing one + mid-flow would reopen the race the parking prevents. The FIFO cap applies to ephemeral + ports only.""" while len(_reserved_sockets) >= _MAX_RESERVED_SOCKETS: stale_port = next((p for p in _reserved_sockets if p not in _CIMD_PORTS), None) if stale_port is None: break # only pinned sockets remain — never evict those - try: + with contextlib.suppress(OSError): _reserved_sockets.pop(stale_port).close() - except OSError: - pass _reserved_sockets[port] = sock @@ -202,12 +195,10 @@ def _cached_redirect_uris(storage: "HermesTokenStorage | None"): def _cached_redirect_port(storage: "HermesTokenStorage | None") -> int | None: - """Loopback callback port from the cached client registration. - - Providers bind a dynamically-registered ``client_id`` to the exact redirect - URI registered with it; a new random port on restart under the stored - ``client_id`` gets ``redirect_uri does not match any registered URIs``. - """ + """Loopback callback port from the cached client registration. Providers bind a + dynamically-registered ``client_id`` to the exact redirect URI registered with it; a new + random port under the stored ``client_id`` gets ``redirect_uri does not match any + registered URIs``.""" for _uri, parsed in _cached_redirect_uris(storage): is_loopback_callback = parsed.scheme == "http" and parsed.path == "/callback" and parsed.hostname in {"127.0.0.1", "localhost"} if is_loopback_callback and parsed.port is not None: @@ -237,26 +228,25 @@ def _is_interactive() -> bool: def _raise_if_non_interactive(lead: str) -> None: - """Raise ``OAuthNonInteractiveError`` unless interactive; *lead* is the - boundary-specific first sentence, the ``hermes mcp login`` next step is shared.""" + """Raise ``OAuthNonInteractiveError`` unless interactive; *lead* is the boundary-specific + first sentence, the ``hermes mcp login`` next step is shared.""" if not _is_interactive(): raise OAuthNonInteractiveError( - f"{lead} Run `hermes mcp login ` interactively to (re)authorize, " - "then restart or reload the gateway." + f"{lead} Run `hermes mcp login ` interactively to (re)authorize, then restart or reload the gateway." ) def force_interactive_oauth(): - """Treat the current context as interactive despite no TTY (GUI-driven auth): - the user IS present, just not on stdin. Crosses the MCP event-loop thread - like ``suppress_interactive_oauth``.""" + """Treat the current context as interactive despite no TTY (GUI-driven auth): the user IS + present, just not on stdin. Crosses the MCP event-loop thread like + ``suppress_interactive_oauth``.""" return _contextvar_set(_oauth_interactive_forced, True) def suppress_interactive_oauth(): - """Disable stdin-based OAuth prompts for the current execution context; - ContextVar-based so a background-discovery thread's suppression reaches the - coroutine scheduled on the MCP event-loop thread.""" + """Disable stdin-based OAuth prompts for the current execution context; ContextVar-based + so a background-discovery thread's suppression reaches the coroutine scheduled on the + MCP event-loop thread.""" return _contextvar_set(_oauth_interactive_enabled, False) @@ -282,14 +272,11 @@ def _read_json(path: Path) -> dict | None: def _write_json(path: Path, data: dict) -> None: - """Atomically write *data* as JSON created at 0o600. - - ``os.open`` with ``O_EXCL`` + explicit mode avoids the write-then-chmod TOCTOU - window where the file briefly inherits the umask (often world-readable). The - parent dir is tightened to 0o700 (no-op on Windows; ``secure_parent_dir`` - refuses /, top-level dirs and the install tree). The per-process random tmp - suffix avoids collisions between concurrent writers and stale crash leftovers. - """ + """Atomically write *data* as JSON created at 0o600. ``os.open`` with ``O_EXCL`` + explicit + mode avoids the write-then-chmod TOCTOU window where the file briefly inherits the umask + (often world-readable). The parent dir is tightened to 0o700 (no-op on Windows; + ``secure_parent_dir`` refuses /, top-level dirs and the install tree). The per-process + random tmp suffix avoids collisions between concurrent writers and stale crash leftovers.""" path.parent.mkdir(parents=True, exist_ok=True) secure_parent_dir(path) tmp = path.with_suffix(f".tmp.{os.getpid()}.{secrets.token_hex(4)}") @@ -306,6 +293,11 @@ def _write_json(path: Path, data: dict) -> None: raise +def _model_json(model: Any) -> dict: + """The on-disk JSON shape of an SDK pydantic model.""" + return model.model_dump(mode="json", exclude_none=True) + + # -- HermesTokenStorage -- persistent token/client-info on disk -------------- class HermesTokenStorage: """Persist OAuth tokens and client registration to JSON files. @@ -356,17 +348,13 @@ class HermesTokenStorage: logger.warning("Corrupt %s at %s -- ignoring: %s", label, path, exc) return None - # -- tokens ------------------------------------------------------------ - + # -- tokens -- def _rebase_expires_in(self, data: dict) -> None: - """Rewrite ``expires_in`` to the seconds remaining now. - - ``set_tokens`` stores an absolute ``expires_at`` (not an SDK field, so it - is stripped here); a relative value reloaded after restart would make - ``is_token_valid()`` True for tokens that expired while down. Legacy files - without it use the file mtime as a best-effort write time, clamped to - zero (self-heals on the next ``set_tokens``). - """ + """Rewrite ``expires_in`` to the seconds remaining now. ``set_tokens`` stores an + absolute ``expires_at`` (not an SDK field, so stripped here); a relative value + reloaded after restart would make ``is_token_valid()`` True for tokens that expired + while down. Legacy files without it use the file mtime as a best-effort write time, + clamped to zero (self-heals on the next ``set_tokens``).""" absolute_expiry = data.pop("expires_at", None) if absolute_expiry is not None: data["expires_in"] = int(max(absolute_expiry - time.time(), 0)) @@ -379,26 +367,23 @@ class HermesTokenStorage: return self._load_model(self._tokens_path(), "OAuthToken", "tokens", self._rebase_expires_in) async def set_tokens(self, tokens: "OAuthToken") -> None: - payload = tokens.model_dump(mode="json", exclude_none=True) - # Persist an absolute ``expires_at``: a relative ``expires_in`` reloaded - # after restart has no wall-clock reference, leaving the SDK's - # ``token_expiry_time=None`` and ``is_token_valid()`` falsely True. + payload = _model_json(tokens) + # Persist an absolute ``expires_at``: a relative ``expires_in`` reloaded after restart + # has no wall-clock reference, leaving the SDK's ``token_expiry_time=None`` and + # ``is_token_valid()`` falsely True. if payload.get("expires_in") is not None: with contextlib.suppress(TypeError, ValueError): # mock tokens / odd shapes: skip, don't fail persistence payload["expires_at"] = time.time() + int(payload["expires_in"]) _write_json(self._tokens_path(), payload) logger.debug("OAuth tokens saved for %s", self._server_name) - # -- client info ------------------------------------------------------- - + # -- client info -- @staticmethod def _coerce_secret_auth_method(data: dict) -> bool: - """Set ``client_secret_post`` when a secret is present but no method is. - - Some DCR providers (notably Supabase) return a ``client_secret`` but omit - ``token_endpoint_auth_method``; the SDK defaults it to ``none``, omits the - secret at the token endpoint and the exchange fails. - """ + """Set ``client_secret_post`` when a secret is present but no method is. Some DCR + providers (notably Supabase) return a ``client_secret`` but omit + ``token_endpoint_auth_method``; the SDK defaults it to ``none``, omits the secret at + the token endpoint and the exchange fails.""" if data.get("client_secret") and data.get("token_endpoint_auth_method") in (None, "none", ""): data["token_endpoint_auth_method"] = "client_secret_post" return True @@ -412,37 +397,32 @@ class HermesTokenStorage: ) if info is not None and coerced and coerced[0]: # Persist the effective method so later flows skip the coercion. - _write_json(self._client_info_path(), info.model_dump(mode="json", exclude_none=True)) + _write_json(self._client_info_path(), _model_json(info)) return info async def set_client_info(self, client_info: "OAuthClientInformationFull") -> None: - data = client_info.model_dump(mode="json", exclude_none=True) + data = _model_json(client_info) self._coerce_secret_auth_method(data) _write_json(self._client_info_path(), data) logger.debug("OAuth client info saved for %s", self._server_name) - # -- oauth server metadata -------------------------------------------- - # The SDK keeps discovered ``OAuthMetadata`` in memory only. Persisting it - # lets a restarted process refresh without re-discovery; otherwise the SDK - # guesses ``{server_url}/token`` (404 on most providers) and forces a full - # browser re-authorization. - + # -- oauth server metadata -- + # The SDK keeps discovered ``OAuthMetadata`` in memory only. Persisting it lets a restarted + # process refresh without re-discovery; otherwise the SDK guesses ``{server_url}/token`` + # (404 on most providers) and forces a full browser re-authorization. def save_oauth_metadata(self, metadata: "OAuthMetadata") -> None: - _write_json(self._meta_path(), metadata.model_dump(mode="json", exclude_none=True)) + _write_json(self._meta_path(), _model_json(metadata)) logger.debug("OAuth metadata saved for %s", self._server_name) def load_oauth_metadata(self) -> "OAuthMetadata | None": return self._load_model(self._meta_path(), "OAuthMetadata", "OAuth metadata") - # -- CIMD refusal ------------------------------------------------------ - + # -- CIMD refusal -- def mark_cimd_rejected(self) -> None: - """Durably record that this server refused our Client ID Metadata Document. - - The in-memory fallback only holds for one process; without a marker every - restart re-presents a refused client_id. Cleared by ``remove()`` - (``hermes mcp login`` / ``remove``) so a fixed document gets a retry. - """ + """Durably record that this server refused our Client ID Metadata Document. The + in-memory fallback only holds for one process; without a marker every restart + re-presents a refused client_id. Cleared by ``remove()`` (``hermes mcp login`` / + ``remove``) so a fixed document gets a retry.""" path = self._cimd_rejected_path() try: path.parent.mkdir(parents=True, exist_ok=True) @@ -454,16 +434,15 @@ class HermesTokenStorage: """True when this server has refused our metadata document before.""" return self._cimd_rejected_path().exists() - # -- cleanup ----------------------------------------------------------- - + # -- cleanup -- def remove(self) -> None: """Delete all stored OAuth state for this server.""" for p in (*self._state_paths(), self._cimd_rejected_path()): p.unlink(missing_ok=True) def snapshot(self) -> dict[str, bytes]: - """filename -> bytes for the existing state files; feed to ``restore()`` - to undo a ``remove()`` after a failed re-auth so a valid token survives.""" + """filename -> bytes for the existing state files; feed to ``restore()`` to undo a + ``remove()`` after a failed re-auth so a valid token survives.""" snap: dict[str, bytes] = {} for p in self._state_paths(): with contextlib.suppress(OSError): @@ -489,12 +468,11 @@ class HermesTokenStorage: logger.warning("Failed to restore OAuth state %s: %s", fname, exc) def poison_client_registration(self) -> bool: - """Discard a dead dynamically-registered client (``invalid_client`` at the - token endpoint) so the SDK re-runs DCR next flow; stale ``meta.json`` goes - too. Tokens are kept: re-auth overwrites them and a valid refresh token - survives if re-registration never completes. Keeps one ``.bak``. - Returns True if a client file was present and removed. - """ + """Discard a dead dynamically-registered client (``invalid_client`` at the token + endpoint) so the SDK re-runs DCR next flow; stale ``meta.json`` goes too. Tokens are + kept: re-auth overwrites them and a valid refresh token survives if re-registration + never completes. Keeps one ``.bak``. Returns True if a client file was present and + removed.""" client_path = self._client_info_path() if not client_path.exists(): return False @@ -520,8 +498,8 @@ class HermesTokenStorage: # -- Callback capture -- HTTP listener and stdin paste share one result dict -- def _authorization_code_result(code: str, state: "str | None", iss: "str | None" = None): """Package redirect parameters in the shape the installed SDK expects: mcp 2.0's - ``callback_handler`` returns an ``AuthorizationCodeResult`` (the SDK reads - ``.state`` / ``.iss`` off it); older SDKs take a tuple.""" + ``callback_handler`` returns an ``AuthorizationCodeResult`` (the SDK reads ``.state`` / + ``.iss`` off it); older SDKs take a tuple.""" try: from mcp.shared.auth import AuthorizationCodeResult except ImportError: # mcp < 2.0 @@ -530,9 +508,9 @@ def _authorization_code_result(code: str, state: "str | None", iss: "str | None" def _parse_redirect_query(query: str) -> dict[str, Any]: - """Extract code/state/error/iss from a redirect query string. ``iss`` is the - RFC 9207 issuer: mcp 2.0 rejects a response that omits it when the server - advertised ``authorization_response_iss_parameter_supported``, so keep it.""" + """Extract code/state/error/iss from a redirect query string. ``iss`` is the RFC 9207 + issuer: mcp 2.0 rejects a response that omits it when the server advertised + ``authorization_response_iss_parameter_supported``, so keep it.""" params = parse_qs(query) return {k: params.get(k, [None])[0] for k in ("code", "state", "error", "iss")} @@ -546,8 +524,7 @@ def _fill_result(result: dict, parsed: dict[str, Any]) -> None: def _make_callback_handler() -> tuple[type, dict]: - """Fresh ``(HandlerClass, result_dict)`` per flow so concurrent flows don't - stomp on each other.""" + """Fresh ``(HandlerClass, result_dict)`` per flow so concurrent flows don't stomp on each other.""" result: dict[str, Any] = {"auth_code": None, "state": None, "error": None, "iss": None} class _Handler(BaseHTTPRequestHandler): @@ -570,14 +547,12 @@ def _make_callback_handler() -> tuple[type, dict]: def _paste_callback_reader(result: dict) -> None: - """Read one stdin line, parse it as an OAuth redirect, write to *result*. - - Accepts a full redirect URL, the provider's own callback URL, just the query - string (``?code=...&state=...`` or ``code=...``), or a skip token - (``skip``/``cancel``/``s``/``n``/``no``/``q``/``quit``) which exits the flow - without auth. Parse failures, EOF and interrupts are swallowed — this is a - best-effort fallback racing the HTTP listener, which stays primary. - """ + """Read one stdin line, parse it as an OAuth redirect, write to *result*. Accepts a full + redirect URL, the provider's own callback URL, just the query string + (``?code=...&state=...`` or ``code=...``), or a skip token + (``skip``/``cancel``/``s``/``n``/``no``/``q``/``quit``) which exits the flow without auth. + Parse failures, EOF and interrupts are swallowed — this is a best-effort fallback racing + the HTTP listener, which stays primary.""" try: line = sys.stdin.readline() except (KeyboardInterrupt, OSError, ValueError): @@ -585,7 +560,6 @@ def _paste_callback_reader(result: dict) -> None: line = line.strip() if line else "" if not line or _result_taken(result): return # EOF / blank, or the HTTP listener already won - if line.lower() in _SKIP_TOKENS: result["error"] = _USER_SKIPPED_SENTINEL print( @@ -594,7 +568,6 @@ def _paste_callback_reader(result: dict) -> None: file=sys.stderr, ) return - # Full URL or "?code=...": take everything after the first "?". query = line.split("?", 1)[1] if "?" in line else line try: @@ -613,7 +586,10 @@ def _paste_callback_reader(result: dict) -> None: # -- Async redirect + callback handlers for OAuthClientProvider -------------- - +# Remote-session guidance printed under the authorization URL. With a proxy callback (e.g. +# Tailscale Funnel) the redirect is forwarded to this machine's listener, so no tunnel/paste +# is needed. On the loopback default the redirect reaches the *remote* machine's listener, +# not the browser's, so the user must paste the redirect URL back or SSH-forward the port. _SSH_HINT_PROXY = ( " Remote session detected. After you authorize, the provider redirects to\n" " {redirect_uri}\n" @@ -637,26 +613,14 @@ _SSH_HINT_LOOPBACK = ( ) -def _print_ssh_hint(port: int, redirect_uri: str | None) -> None: - """Remote-session guidance printed under the authorization URL. - - With a proxy callback (e.g. Tailscale Funnel) the redirect is forwarded to - this machine's listener, so no tunnel/paste is needed. On the loopback default - the redirect reaches the *remote* machine's listener, not the browser's, so - the user must paste the redirect URL back or SSH-forward the port. - """ - if redirect_uri: - print(_SSH_HINT_PROXY.format(redirect_uri=redirect_uri), file=sys.stderr) - elif port: - print(_SSH_HINT_LOOPBACK.format(port=port), file=sys.stderr) - - def _announce_authorization_url(authorization_url: str, port: int, redirect_uri: str | None) -> None: """Print the URL (always, as the fallback) and open the browser when possible.""" print(f"\n MCP OAuth: authorization required.\n Open this URL in your browser:\n\n {authorization_url}\n", file=sys.stderr) if os.getenv("SSH_CLIENT") or os.getenv("SSH_TTY"): - _print_ssh_hint(port, redirect_uri) - + if redirect_uri: + print(_SSH_HINT_PROXY.format(redirect_uri=redirect_uri), file=sys.stderr) + elif port: + print(_SSH_HINT_LOOPBACK.format(port=port), file=sys.stderr) if not _can_open_browser(): note = "Headless environment detected — open the URL manually." else: @@ -669,23 +633,20 @@ def _announce_authorization_url(authorization_url: str, port: int, redirect_uri: def _make_redirect_handler(port: int, redirect_uri: str | None = None): - """Return a redirect handler closing over this flow's port (a closure, not the - module-level ``_oauth_port``, keeps concurrent server flows isolated). - ``redirect_uri`` is a configured proxy callback (None for loopback) and only - tailors the remote-session hint.""" + """Return a redirect handler closing over this flow's port (a closure, not the module-level + ``_oauth_port``, keeps concurrent server flows isolated). ``redirect_uri`` is a configured + proxy callback (None for loopback) and only tailors the remote-session hint.""" async def _redirect_handler(authorization_url: str) -> None: dashboard_flow = get_dashboard_oauth_flow() if dashboard_flow is not None: await dashboard_flow.publish_authorization_url(authorization_url) return - # Fail fast in non-interactive contexts (systemd gateway, cron, background - # discovery): a cached-but-unusable token makes the SDK fall through to - # the authorization-code flow even though the token-file guard passed, and - # we would otherwise block in the waiter for the full timeout. - # Deliberately re-checks interactivity here. + # Fail fast in non-interactive contexts (systemd gateway, cron, background discovery): + # a cached-but-unusable token makes the SDK fall through to the authorization-code + # flow even though the token-file guard passed, and we would otherwise block in the + # waiter for the full timeout. Deliberately re-checks interactivity here. _raise_if_non_interactive( - "MCP OAuth requires browser authorization but no interactive " - "session is available (non-interactive/background context)." + "MCP OAuth requires browser authorization but no interactive session is available (non-interactive/background context)." ) _announce_authorization_url(authorization_url, port, redirect_uri) @@ -693,12 +654,9 @@ def _make_redirect_handler(port: int, redirect_uri: str | None = None): def _start_callback_server(port: int, handler_cls: type) -> HTTPServer: - """Bind the callback listener on *port*, adopting a parked reserved socket. - - Adopting the bound socket closes the select→bind TOCTOU window. - ``allow_reuse_address`` is set BEFORE binding (a no-op afterwards) so a - lingering TIME_WAIT socket from a previous flow cannot block the next. - """ + """Bind the callback listener on *port*, adopting a parked reserved socket (closes the + select→bind TOCTOU window). ``allow_reuse_address`` is set BEFORE binding (a no-op + afterwards) so a lingering TIME_WAIT socket from a previous flow cannot block the next.""" try: server = HTTPServer(("127.0.0.1", port), handler_cls, bind_and_activate=False) reserved = _reserved_sockets.pop(port, None) @@ -710,12 +668,11 @@ def _start_callback_server(port: int, handler_cls: type) -> HTTPServer: server.server_bind() server.server_activate() except OSError as exc: - # Genuinely in use: a concurrent flow, leftover listener, or colliding - # fixed `oauth.redirect_port`. Nothing to poll — say so, not "timed out". + # Genuinely in use: a concurrent flow, leftover listener, or colliding fixed + # `oauth.redirect_port`. Nothing to poll — say so, not "timed out". raise OAuthNonInteractiveError( - f"OAuth callback port {port} is already in use ({exc}). " - "Close any other in-progress login, or set a free `oauth.redirect_port` " - "in the server config, then retry." + f"OAuth callback port {port} is already in use ({exc}). Close any other in-progress login, " + "or set a free `oauth.redirect_port` in the server config, then retry." ) from exc return server @@ -739,30 +696,23 @@ def _callback_outcome(result: dict, cimd_url: str | None): raise RuntimeError(f"OAuth authorization failed: {result['error']}") if result["auth_code"] is None: hint = ( - " If the browser showed an invalid-client error instead of " - "an approval prompt, the authorization server rejected " - f"Hermes' Client ID Metadata Document ({cimd_url}); set " - "``cimd: false`` under that server's ``oauth:`` block in " - "config.yaml to authorize via dynamic client registration " - "instead." + " If the browser showed an invalid-client error instead of an approval prompt, the authorization " + f"server rejected Hermes' Client ID Metadata Document ({cimd_url}); set ``cimd: false`` under that " + "server's ``oauth:`` block in config.yaml to authorize via dynamic client registration instead." ) if cimd_url else "" raise OAuthNonInteractiveError( - "OAuth callback timed out — no authorization code received. " - "Ensure you completed the browser authorization flow." + hint + "OAuth callback timed out — no authorization code received. Ensure you completed the browser authorization flow." + hint ) return _authorization_code_result(result["auth_code"], result["state"], result.get("iss")) def _make_callback_waiter(port: int, cimd_url: str | None = None, timeout: float = 300.0): """Return a callback waiter bound to one flow's port (isolating concurrent flows). - - ``timeout`` is the only place ``oauth.timeout`` applies (mcp 2.0 dropped the - provider's own). ``cimd_url`` only tailors the timeout message: a server that - refuses the document aborts at the *authorization* endpoint, so no redirect - arrives and a bare "timed out" would hide the cause. On a TTY the HTTP - listener races a stdin paste fallback. Raises ``OAuthNonInteractiveError`` on - timeout or when non-interactive. - """ + ``timeout`` is the only place ``oauth.timeout`` applies (mcp 2.0 dropped the provider's + own). ``cimd_url`` only tailors the timeout message: a server that refuses the document + aborts at the *authorization* endpoint, so no redirect arrives and a bare "timed out" + would hide the cause. On a TTY the HTTP listener races a stdin paste fallback. Raises + ``OAuthNonInteractiveError`` on timeout or when non-interactive.""" async def _wait(): dashboard_flow = get_dashboard_oauth_flow() @@ -770,31 +720,25 @@ def _make_callback_waiter(port: int, cimd_url: str | None = None, timeout: float # Dashboard flow speaks the legacy tuple; normalize to one shape. dash_code, dash_state = await dashboard_flow.wait_for_callback() return _authorization_code_result(dash_code, dash_state) - - # Reaching here means the SDK entered the authorization-code flow, so any - # cached token is unusable. Reject BEFORE binding: binding would block for - # the full timeout and collide with the TIME_WAIT port on retry - # (``Address already in use``). Holds regardless of token files. + # Reaching here means the SDK entered the authorization-code flow, so any cached token + # is unusable. Reject BEFORE binding: binding would block for the full timeout and + # collide with the TIME_WAIT port on retry (``Address already in use``). Holds + # regardless of token files. _raise_if_non_interactive( - "OAuth callback requires an interactive session but none is " - "available (non-interactive/background context); skipping browser " - "authorization without binding a callback listener." + "OAuth callback requires an interactive session but none is available (non-interactive/background " + "context); skipping browser authorization without binding a callback listener." ) - handler_cls, result = _make_callback_handler() server = _start_callback_server(port, handler_cls) threading.Thread(target=server.handle_request, daemon=True).start() - # Paste fallback races the HTTP listener; whichever fills result first wins. if _is_interactive(): print( - "\n Or paste the redirect URL here (or the ``?code=...&state=...`` " - "portion) and press Enter. Type ``skip`` + Enter to continue " - "without this server:", + "\n Or paste the redirect URL here (or the ``?code=...&state=...`` portion) and press Enter. " + "Type ``skip`` + Enter to continue without this server:", file=sys.stderr, flush=True, ) threading.Thread(target=_paste_callback_reader, args=(result,), daemon=True).start() - await _poll_callback_result(server, result, timeout) return _callback_outcome(result, cimd_url) @@ -811,15 +755,11 @@ def _get_hermes_oauth_provider_class() -> type | None: if HermesOAuthClientProvider is None and _ensure_sdk_loaded(): from tools.mcp_oauth_provider import HermesProviderMixin - HermesOAuthClientProvider = type( - "HermesOAuthClientProvider", - (HermesProviderMixin, OAuthClientProvider), - { - "__doc__": "SDK provider plus Hermes' token-endpoint fixes (see ``HermesProviderMixin``).", - "__module__": __name__, - "_hermes_logger": logger, - }, - ) + HermesOAuthClientProvider = type("HermesOAuthClientProvider", (HermesProviderMixin, OAuthClientProvider), { + "__doc__": "SDK provider plus Hermes' token-endpoint fixes (see ``HermesProviderMixin``).", + "__module__": __name__, + "_hermes_logger": logger, + }) return HermesOAuthClientProvider @@ -830,20 +770,19 @@ def remove_oauth_tokens(server_name: str, *, hermes_home: str | Path | None = No # -- CIMD -- OAuth Client ID Metadata Documents ------------------------------- -# Under CIMD the client_id IS an HTTPS URL the authorization server fetches to -# learn our name, logo and permitted redirect URIs, replacing per-install DCR. -# The SDK does the protocol work; Hermes only decides whether a flow is -# eligible and hands the URL to ``OAuthClientProvider``. +# Under CIMD the client_id IS an HTTPS URL the authorization server fetches to learn our +# name, logo and permitted redirect URIs, replacing per-install DCR. The SDK does the protocol +# work; Hermes only decides whether a flow is eligible and hands the URL to +# ``OAuthClientProvider``. -# Published from ``website/static/oauth/client-metadata.json`` by the docs -# deploy. The github.io origin is deliberate: an authorization server MUST NOT -# follow HTTP redirects when fetching the document, and -# hermes-agent.nousresearch.com/docs/* 301s here. +# Published from ``website/static/oauth/client-metadata.json`` by the docs deploy. The +# github.io origin is deliberate: an authorization server MUST NOT follow HTTP redirects when +# fetching the document, and hermes-agent.nousresearch.com/docs/* 301s here. _CIMD_CLIENT_METADATA_URL = "https://nousresearch.github.io/hermes-agent/docs/oauth/client-metadata.json" -# Loopback callback ports declared in that document (exact string match, so a -# CIMD flow cannot use an ephemeral port). Below Linux's 32768 ephemeral floor so -# the kernel never hands one to another process. Keep in sync with the document -# (tests/tools/test_mcp_cimd.py enforces it). +# Loopback callback ports declared in that document (exact string match, so a CIMD flow +# cannot use an ephemeral port). Below Linux's 32768 ephemeral floor so the kernel never +# hands one to another process. Keep in sync with the document (tests/tools/test_mcp_cimd.py +# enforces it). _CIMD_PORTS = (27890, 27891, 27892, 27893, 27894) # Loopback hostnames the document lists alongside each port, so the # ``oauth.redirect_host: localhost`` WAF workaround still works under CIMD. @@ -851,13 +790,11 @@ _CIMD_REDIRECT_HOSTS = frozenset({"127.0.0.1", "localhost"}) def _is_valid_cimd_url(url: str) -> bool: - """True when *url* is usable as a CIMD client_id on the installed SDK. - - Delegates to the SDK's validator (ImportError = SDK predates CIMD → DCR - only). The SDK checks only https-scheme and non-root-path; userinfo, - fragments and dot segments are rejected here because they fail at the - authorization server mid-browser-flow as an opaque invalid-client page. - """ + """True when *url* is usable as a CIMD client_id on the installed SDK. Delegates to the + SDK's validator (ImportError = SDK predates CIMD → DCR only). The SDK checks only + https-scheme and non-root-path; userinfo, fragments and dot segments are rejected here + because they fail at the authorization server mid-browser-flow as an opaque + invalid-client page.""" try: from mcp.client.auth.utils import is_valid_client_metadata_url except ImportError: @@ -872,30 +809,21 @@ def _is_valid_cimd_url(url: str) -> bool: return not (has_userinfo or parsed.fragment or any(seg in {".", ".."} for seg in parsed.path.split("/"))) -# Pinned ports this process has committed to, in order taken. A provider is -# built once per server and keeps its port for the process lifetime, so -# assignments are never released. Includes ports restored from a cached -# registration so a sibling server is never handed one already in use. +# Pinned ports this process has committed to, in order taken. A provider is built once per +# server and keeps its port for the process lifetime, so assignments are never released. +# Includes ports restored from a cached registration so a sibling server is never handed one +# already in use. _assigned_cimd_ports: "list[int]" = [] -def _note_assigned_cimd_port(port: int) -> None: - """Claim *port* for this process when it belongs to the pinned range.""" - if port in _CIMD_PORTS and port not in _assigned_cimd_ports: - _assigned_cimd_ports.append(port) - - def _pick_cimd_port() -> int | None: - """Reserve a pinned CIMD callback port, or None when none is usable. - - Holding the bound socket until adoption prevents the same steal window as - for ephemeral ports, and makes contention cooperative: another profile or - sibling server finds the bind refused and moves down the range. Once every - pinned port belongs to this process the range wraps rather than falling - back to DCR — a reused port only bites if both servers authorize at the - same moment (reported clearly by the waiter), whereas DCR may be - unsupported by the server entirely. - """ + """Reserve a pinned CIMD callback port, or None when none is usable. Holding the bound + socket until adoption prevents the same steal window as for ephemeral ports and makes + contention cooperative: another profile or sibling server finds the bind refused and + moves down the range. Once every pinned port belongs to this process the range wraps + rather than falling back to DCR — a reused port only bites if both servers authorize at + the same moment (reported clearly by the waiter), whereas DCR may be unsupported by the + server entirely.""" for port in _CIMD_PORTS: if port not in _assigned_cimd_ports and _bind_reserved(port) is not None: _assigned_cimd_ports.append(port) @@ -904,12 +832,10 @@ def _pick_cimd_port() -> int | None: def _server_declined_cimd(storage: "HermesTokenStorage | None") -> bool: - """True when cached metadata shows this server doesn't advertise CIMD. - - The SDK decides CIMD vs DCR in its 401 branch — after Hermes must fix the - redirect URI. Cached authorization-server metadata closes the gap for every - server already reached: only a genuinely unknown server pays the optimistic pin. - """ + """True when cached metadata shows this server doesn't advertise CIMD. The SDK decides + CIMD vs DCR in its 401 branch — after Hermes must fix the redirect URI. Cached + authorization-server metadata closes the gap for every server already reached: only a + genuinely unknown server pays the optimistic pin.""" try: metadata = storage.load_oauth_metadata() if storage is not None else None except (AttributeError, TypeError, ValueError): @@ -918,30 +844,28 @@ def _server_declined_cimd(storage: "HermesTokenStorage | None") -> bool: def _maybe_use_cimd(cfg: dict, storage: "HermesTokenStorage | None" = None) -> "tuple[str, int] | None": - """Return ``(client_id URL, pinned callback port)``, or None to use DCR. - - Every ineligibility case below is one where the redirect URI Hermes would - send is not one the document declares, where the client identity is already - settled, or where the server is known not to want a document — passing a - metadata URL anyway would make the authorization server reject the flow. - """ + """Return ``(client_id URL, pinned callback port)``, or None to use DCR. Every + ineligibility case is one where the redirect URI Hermes would send is not one the + document declares, where the client identity is already settled, or where the server is + known not to want a document — passing a metadata URL anyway would make the + authorization server reject the flow.""" url = cfg.get("client_metadata_url") or _CIMD_CLIENT_METADATA_URL ineligible = ( cfg.get("cimd") is False or not _is_valid_cimd_url(url) - # A pinned client is the user's explicit choice; a secret means a - # confidential client, which the document forbids. + # A pinned client is the user's explicit choice; a secret means a confidential + # client, which the document forbids. or cfg.get("client_id") or cfg.get("client_secret") - # The document supplies name and auth method; a caller setting either - # asks for an identity CIMD cannot present (Figma's DCR name allowlist). + # The document supplies name and auth method; a caller setting either asks for an + # identity CIMD cannot present (Figma's DCR name allowlist). or cfg.get("client_name") or (cfg.get("token_endpoint_auth_method") or "none") != "none" - # Dashboard/desktop flows redirect to a deployment-specific server URL - # that can never appear in a static document. + # Dashboard/desktop flows redirect to a deployment-specific server URL that can + # never appear in a static document. or get_dashboard_oauth_flow() is not None or cfg.get("redirect_uri") or cfg.get("redirect_port") or (cfg.get("redirect_host") or "127.0.0.1") not in _CIMD_REDIRECT_HOSTS - # An existing registration is bound to its redirect URI; swapping in a - # CIMD client_id now would invalidate stored tokens. + # An existing registration is bound to its redirect URI; swapping in a CIMD + # client_id now would invalidate stored tokens. or _cached_client_info(storage) is not None or (storage is not None and storage.cimd_rejected()) or _server_declined_cimd(storage) @@ -953,33 +877,29 @@ def _maybe_use_cimd(cfg: dict, storage: "HermesTokenStorage | None" = None) -> " def cimd_provider_kwargs(cfg: dict) -> dict[str, Any]: - """``client_metadata_url=`` for ``OAuthClientProvider``, when CIMD applies. - Returned as kwargs so the argument is omitted entirely on a DCR flow: an SDK - too old for CIMD rejects the keyword outright.""" + """``client_metadata_url=`` for ``OAuthClientProvider``, when CIMD applies. Returned as + kwargs so the argument is omitted entirely on a DCR flow: an SDK too old for CIMD rejects + the keyword outright.""" url = cfg.get("_cimd_url") return {"client_metadata_url": url} if url else {} def token_request_user_agent(cfg: dict) -> str | None: - """Configured ``oauth.user_agent`` for token-endpoint requests, or None. - - Opt-in and per-server; anything but a non-empty string is unset so a - null/empty YAML value never sends a blank header. Applied ONLY to - authorization-code exchange and refresh — never MCP traffic or discovery; - no other headers are configurable (secrets would land in config.yaml). - """ + """Configured ``oauth.user_agent`` for token-endpoint requests, or None. Opt-in and + per-server; anything but a non-empty string is unset so a null/empty YAML value never + sends a blank header. Applied ONLY to authorization-code exchange and refresh — never MCP + traffic or discovery; no other headers are configurable (secrets would land in + config.yaml).""" ua = cfg.get("user_agent") return ua.strip() if isinstance(ua, str) and ua.strip() else None def _configure_callback_port(cfg: dict, storage: "HermesTokenStorage | None" = None) -> int: """Resolve the callback port into ``cfg['_resolved_port']`` (0 = non-loopback URI). - - Precedence: dashboard flow / cached https redirect URI → CIMD pinned port - (also sets ``cfg['_cimd_url']``) → ``oauth.redirect_port`` → cached - registration port → fresh ephemeral port. Only the fresh pick is parked; - fixed ports bind via reuse_address. Also sets the legacy ``_oauth_port``. - """ + Precedence: dashboard flow / cached https redirect URI → CIMD pinned port (also sets + ``cfg['_cimd_url']``) → ``oauth.redirect_port`` → cached registration port → fresh + ephemeral port. Only the fresh pick is parked; fixed ports bind via reuse_address. Also + sets the legacy ``_oauth_port``.""" global _oauth_port dashboard_flow = get_dashboard_oauth_flow() if dashboard_flow is not None: @@ -996,9 +916,10 @@ def _configure_callback_port(cfg: dict, storage: "HermesTokenStorage | None" = N cfg["_cimd_url"], port = cimd else: port = int(cfg.get("redirect_port", 0)) or _cached_redirect_port(storage) or _reserve_callback_port() - # A cached port may be a pinned CIMD port left by an earlier CIMD - # login; claim it so a sibling's _pick_cimd_port doesn't reuse it. - _note_assigned_cimd_port(port) + # A cached port may be a pinned CIMD port left by an earlier CIMD login; claim it so a + # sibling's _pick_cimd_port doesn't reuse it. + if port in _CIMD_PORTS and port not in _assigned_cimd_ports: + _assigned_cimd_ports.append(port) cfg["_resolved_port"] = port _oauth_port = port return port @@ -1006,20 +927,16 @@ def _configure_callback_port(cfg: dict, storage: "HermesTokenStorage | None" = N def _resolve_redirect_uri(cfg: dict, port: int) -> str: """Configured ``redirect_uri`` (proxy, e.g. Tailscale Funnel) or - ``http://:/callback``. - - Client metadata and pre-registered client info must both derive the URI here - so they stay identical. ``redirect_host`` (default ``127.0.0.1``) only - changes the hostname: some WAFs reject a literal ``127.0.0.1`` in the - authorize query, and ``localhost`` works around it. The listener still binds - ``127.0.0.1``. - """ + ``http://:/callback``. Client metadata and pre-registered client info + must both derive the URI here so they stay identical. ``redirect_host`` (default + ``127.0.0.1``) only changes the hostname: some WAFs reject a literal ``127.0.0.1`` in the + authorize query, and ``localhost`` works around it. The listener still binds ``127.0.0.1``.""" return cfg.get("redirect_uri") or f"http://{cfg.get('redirect_host') or '127.0.0.1'}:{port}/callback" -# Figma's remote MCP implements DCR as a client_name *allowlist*: "Claude Code" -# and "Codex" register (200); "Hermes Agent"/"Cursor"/… get 403. Register under -# an allowlisted name so the browser flow can start; oauth.client_name overrides. +# Figma's remote MCP implements DCR as a client_name *allowlist*: "Claude Code" and "Codex" +# register (200); "Hermes Agent"/"Cursor"/… get 403. Register under an allowlisted name so +# the browser flow can start; oauth.client_name overrides. _FIGMA_DCR_CLIENT_NAME = "Claude Code" _FIGMA_DEFAULT_SCOPE = "mcp:connect" @@ -1036,24 +953,21 @@ def _is_figma_remote_mcp(server_name: str | None = None, server_url: str | None def apply_oauth_provider_defaults(cfg: dict, *, server_name: str = "", server_url: str | None = None) -> dict: - """Mutate *cfg* with provider-specific OAuth workarounds. Returns *cfg*. - - Call before building client metadata / pre-registering. Only fills keys - the user left unset — explicit ``oauth.client_name`` / ``oauth.scope`` win. - """ + """Mutate *cfg* with provider-specific OAuth workarounds; returns *cfg*. Call before + building client metadata / pre-registering. Only fills keys the user left unset — + explicit ``oauth.client_name`` / ``oauth.scope`` win.""" if _is_figma_remote_mcp(server_name, server_url): if not cfg.get("client_name"): cfg["client_name"] = _FIGMA_DCR_CLIENT_NAME logger.info( - "MCP OAuth '%s': Figma DCR allowlist — registering as " - "client_name=%r (override via oauth.client_name)", + "MCP OAuth '%s': Figma DCR allowlist — registering as client_name=%r (override via oauth.client_name)", server_name or server_url, _FIGMA_DCR_CLIENT_NAME, ) if not cfg.get("scope"): cfg["scope"] = _FIGMA_DEFAULT_SCOPE - # Figma advertises token_endpoint_auth_method=none yet returns a - # client_secret and then demands it at the token endpoint; request a - # confidential-client registration so the SDK posts the secret. + # Figma advertises token_endpoint_auth_method=none yet returns a client_secret and + # then demands it at the token endpoint; request a confidential-client registration + # so the SDK posts the secret. cfg["token_endpoint_auth_method"] = cfg.get("token_endpoint_auth_method") or "client_secret_post" return cfg @@ -1063,10 +977,9 @@ def _build_client_metadata(cfg: dict) -> "OAuthClientMetadata": port = cfg.get("_resolved_port") if port is None: raise ValueError("_configure_callback_port() must be called before _build_client_metadata()") - if OAuthClientMetadata is None: - _ensure_sdk_loaded() - # Public client by default; confidential only with a known secret or a - # provider (e.g. Figma) that needs confidential-style token posts. + metadata_cls = _sdk_class("OAuthClientMetadata") + # Public client by default; confidential only with a known secret or a provider (e.g. + # Figma) that needs confidential-style token posts. auth_method = cfg.get("token_endpoint_auth_method") or ("client_secret_post" if cfg.get("client_secret") else "none") metadata_kwargs: dict[str, Any] = { "client_name": cfg.get("client_name", "Hermes Agent"), @@ -1074,34 +987,31 @@ def _build_client_metadata(cfg: dict) -> "OAuthClientMetadata": "grant_types": ["authorization_code", "refresh_token"], "response_types": ["code"], "token_endpoint_auth_method": auth_method, - # SEP-837: clients MUST declare application_type so OIDC-strict servers - # accept loopback redirect_uris; Hermes is a CLI/desktop app → "native". - # Overridable for a hosted dashboard fronting a real https redirect. + # SEP-837: clients MUST declare application_type so OIDC-strict servers accept + # loopback redirect_uris; Hermes is a CLI/desktop app → "native". Overridable for a + # hosted dashboard fronting a real https redirect. "application_type": cfg.get("application_type", "native"), } if cfg.get("scope"): metadata_kwargs["scope"] = cfg["scope"] try: - return OAuthClientMetadata.model_validate(metadata_kwargs) + return metadata_cls.model_validate(metadata_kwargs) except Exception: # mcp 1.x metadata models predate SEP-837 and reject the unknown field. metadata_kwargs.pop("application_type", None) - return OAuthClientMetadata.model_validate(metadata_kwargs) + return metadata_cls.model_validate(metadata_kwargs) def _invalidate_tokens_on_client_change( storage: "HermesTokenStorage", new_client_id: str, new_client_secret: str | None ) -> None: - """Drop cached tokens when the configured OAuth client identity changes. - - Tokens are minted for a specific ``client_id``; after editing - ``oauth.client_id`` / ``oauth.client_secret`` (or switching from DCR to a - pre-registered client) the old tokens fail refresh with ``invalid_client``. - Pre-registered clients are exempt from the auto-poison path, so without this - check stale tokens wedge every request until a manual wipe. Compares on-disk - ``client.json`` against the incoming identity BEFORE it is overwritten; a - matching identity is a no-op. - """ + """Drop cached tokens when the configured OAuth client identity changes. Tokens are minted + for a specific ``client_id``; after editing ``oauth.client_id`` / ``oauth.client_secret`` + (or switching from DCR to a pre-registered client) the old tokens fail refresh with + ``invalid_client``. Pre-registered clients are exempt from the auto-poison path, so + without this check stale tokens wedge every request until a manual wipe. Compares on-disk + ``client.json`` against the incoming identity BEFORE it is overwritten; a matching + identity is a no-op.""" existing = _read_json(storage._client_info_path()) old_client_id = existing.get("client_id") if isinstance(existing, dict) else None if not old_client_id: @@ -1110,22 +1020,20 @@ def _invalidate_tokens_on_client_change( return removed = False for path in (storage._tokens_path(), storage._meta_path()): + if not path.exists(): + continue try: - if path.exists(): - path.unlink() - removed = True + path.unlink() + removed = True except OSError as exc: # non-fatal — stale tokens fail later anyway logger.warning( - "MCP OAuth '%s': could not remove stale %s after client " - "change: %s", storage._server_name, path.name, exc, + "MCP OAuth '%s': could not remove stale %s after client change: %s", storage._server_name, path.name, exc, ) if removed: logger.warning( - "MCP OAuth '%s': configured OAuth client changed (client_id %r " - "-> %r); discarded tokens minted under the previous client. " - "Re-authorize with: hermes mcp login %s", - storage._server_name, old_client_id, new_client_id, - storage._server_name, + "MCP OAuth '%s': configured OAuth client changed (client_id %r -> %r); discarded tokens minted under " + "the previous client. Re-authorize with: hermes mcp login %s", + storage._server_name, old_client_id, new_client_id, storage._server_name, ) @@ -1134,8 +1042,7 @@ def _maybe_preregister_client(storage: "HermesTokenStorage", cfg: dict, client_m client_id = cfg.get("client_id") if not client_id: return - if OAuthClientInformationFull is None: - _ensure_sdk_loaded() + info_cls = _sdk_class("OAuthClientInformationFull") _invalidate_tokens_on_client_change(storage, client_id, cfg.get("client_secret")) info_dict: dict[str, Any] = { "client_id": client_id, @@ -1145,20 +1052,17 @@ def _maybe_preregister_client(storage: "HermesTokenStorage", cfg: dict, client_m "token_endpoint_auth_method": client_metadata.token_endpoint_auth_method, **{key: cfg[key] for key in ("client_secret", "client_name", "scope") if cfg.get(key)}, } - client_info = OAuthClientInformationFull.model_validate(info_dict) - _write_json(storage._client_info_path(), client_info.model_dump(mode="json", exclude_none=True)) + _write_json(storage._client_info_path(), _model_json(info_cls.model_validate(info_dict))) logger.debug("Pre-registered client_id=%s for '%s'", client_id, storage._server_name) def humanize_oauth_registration_error( server_name: str, exc: BaseException | str, *, server_url: str | None = None ) -> str | None: - """Turn a Dynamic Client Registration 403/Forbidden into a useful next step. - - Returns None for anything else so the caller keeps the original text. - Figma gates DCR on exact ``client_name``; Hermes auto-sets ``Claude Code``, - so this fires when the user overrode it or an older Hermes is running. - """ + """Turn a Dynamic Client Registration 403/Forbidden into a useful next step; None for + anything else so the caller keeps the original text. Figma gates DCR on exact + ``client_name``; Hermes auto-sets ``Claude Code``, so this fires when the user overrode + it or an older Hermes is running.""" msg = str(exc) lowered = msg.lower() if "403" not in msg and "forbidden" not in lowered: @@ -1170,55 +1074,41 @@ def humanize_oauth_registration_error( ) if not looks_like_registration: return None - if _is_figma_remote_mcp(server_name, server_url): return ( - f"'{server_name}' is Figma's remote MCP — DCR is allowlisted by " - f"exact client_name (\"{_FIGMA_DCR_CLIENT_NAME}\" and \"Codex\" " - "work; most other names 403). Hermes defaults to " - f"client_name: {_FIGMA_DCR_CLIENT_NAME!r} automatically. If you " - "set oauth.client_name yourself, change it to one of those, or " - "clear it and re-run:\n" - f" hermes mcp login {server_name}" + f"'{server_name}' is Figma's remote MCP — DCR is allowlisted by exact client_name " + f"(\"{_FIGMA_DCR_CLIENT_NAME}\" and \"Codex\" work; most other names 403). Hermes defaults to " + f"client_name: {_FIGMA_DCR_CLIENT_NAME!r} automatically. If you set oauth.client_name yourself, " + f"change it to one of those, or clear it and re-run:\n hermes mcp login {server_name}" ) - return ( - f"'{server_name}' only allows pre-approved OAuth clients — it rejected " - "client registration (403), so no browser flow can start. Options: " - "set oauth.client_name to a name the provider allowlists, add a " - "pre-registered client (oauth: {client_id: ..., client_secret: ...}), " - "or use the provider's stdio / API-key / local server instead." + f"'{server_name}' only allows pre-approved OAuth clients — it rejected client registration (403), so no " + "browser flow can start. Options: set oauth.client_name to a name the provider allowlists, add a " + "pre-registered client (oauth: {client_id: ..., client_secret: ...}), or use the provider's stdio / " + "API-key / local server instead." ) def build_oauth_auth(server_name: str, server_url: str, oauth_config: dict | None = None) -> "OAuthClientProvider | None": - """Build an ``httpx.Auth``-compatible OAuth handler for an MCP server. - - Legacy public API; new code should use - :func:`tools.mcp_oauth_manager.get_manager` so OAuth state is shared across - config-time, runtime, and reconnect paths. Returns None if the MCP SDK lacks - OAuth support. - """ + """Build an ``httpx.Auth``-compatible OAuth handler for an MCP server. Legacy public API; + new code should use :func:`tools.mcp_oauth_manager.get_manager` so OAuth state is shared + across config-time, runtime, and reconnect paths. Returns None if the MCP SDK lacks OAuth + support.""" if not _OAUTH_AVAILABLE or (OAuthClientProvider is None and not _ensure_sdk_loaded()): logger.warning( - "MCP OAuth requested for '%s' but SDK auth types are not available. " - "Install with: pip install 'mcp>=1.26.0'", + "MCP OAuth requested for '%s' but SDK auth types are not available. Install with: pip install 'mcp>=1.26.0'", server_name, ) return None - from tools.mcp_oauth_provider import build_provider_kwargs, prepare_oauth_config cfg, storage = prepare_oauth_config(server_name, server_url, oauth_config) if 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. The OAuth flow requires browser authorization. Run " - f"`hermes mcp login {server_name}` interactively first to complete " + f"MCP OAuth for '{server_name}': non-interactive environment and no cached tokens found. The OAuth flow " + f"requires browser authorization. Run `hermes mcp login {server_name}` interactively first to complete " "initial authorization, then cached tokens will be reused." ) - kwargs = build_provider_kwargs(cfg, storage, ssh_proxy_hint=True) provider_class = _get_hermes_oauth_provider_class() if provider_class is None: diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index a0d0d34b18..99891d21e3 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -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 `` and (indirectly) by - ``hermes mcp login `` 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 "" 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