diff --git a/tests/tools/test_mcp_oauth.py b/tests/tools/test_mcp_oauth.py index 8052a79135..bfb0cc5ebe 100644 --- a/tests/tools/test_mcp_oauth.py +++ b/tests/tools/test_mcp_oauth.py @@ -15,16 +15,27 @@ from tools.mcp_oauth import ( OAuthNonInteractiveError, build_oauth_auth, remove_oauth_tokens, - _find_free_port, _can_open_browser, _is_interactive, - _wait_for_callback, _make_callback_handler, _make_redirect_handler, _paste_callback_reader, ) +def _find_free_port() -> int: + import socket + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.bind(("127.0.0.1", 0)) + return s.getsockname()[1] + + +async def _wait_for_callback(): + """Await the per-flow waiter on the legacy module-level port (the removed shim).""" + import tools.mcp_oauth as mod + return await mod._make_callback_waiter(mod._oauth_port)() + + def _set_interactive_stdin(monkeypatch, *, is_tty: bool = True) -> None: mock_stdin = MagicMock() mock_stdin.isatty.return_value = is_tty @@ -503,7 +514,7 @@ class TestCallbackPortReservation: monkeypatch.setattr(mod, "_raise_if_non_interactive", lambda lead: None) async def drive(): - task = asyncio.create_task(mod._wait_for_callback()) + task = asyncio.create_task(_wait_for_callback()) threading.Thread( target=_hit_callback_when_ready, args=(f"http://127.0.0.1:{port}/callback?code=abc123&state=xyz",), @@ -814,7 +825,7 @@ class TestNonInteractiveFailFastAtCallbackBoundary: monkeypatch.setattr(mod.asyncio, "sleep", no_sleep) with pytest.raises(OAuthNonInteractiveError, match="interactive session"): - asyncio.run(mod._wait_for_callback()) + asyncio.run(_wait_for_callback()) fake_server.assert_not_called() def test_redirect_handler_rejects_and_does_not_open_browser(self, monkeypatch, capsys): @@ -1072,7 +1083,7 @@ def test_wait_for_callback_port_in_use_reports_clear_error(monkeypatch): mo, "HTTPServer", side_effect=OSError("address already in use") ): with pytest.raises(mo.OAuthNonInteractiveError) as excinfo: - asyncio.run(mo._wait_for_callback()) + asyncio.run(_wait_for_callback()) msg = str(excinfo.value) assert "54321" in msg diff --git a/tests/tools/test_mcp_oauth_metadata.py b/tests/tools/test_mcp_oauth_metadata.py index 57930fbfc4..335e08188b 100644 --- a/tests/tools/test_mcp_oauth_metadata.py +++ b/tests/tools/test_mcp_oauth_metadata.py @@ -108,7 +108,7 @@ class TestManagerOAuthProviderMetadata: provider = _manager_provider_with_context(storage, oauth_metadata=None) with patch.object( - _HERMES_PROVIDER_CLS.__bases__[0], "_initialize", new=AsyncMock() + _HERMES_PROVIDER_CLS.__bases__[-1], "_initialize", new=AsyncMock() ): asyncio.run(provider._initialize()) @@ -136,7 +136,7 @@ class TestManagerOAuthProviderMetadata: manager.invalidate_if_disk_changed = AsyncMock(return_value=False) with patch.object( - _HERMES_PROVIDER_CLS.__bases__[0], + _HERMES_PROVIDER_CLS.__bases__[-1], "async_auth_flow", new=fake_parent_flow, ), patch("tools.mcp_oauth_manager.get_manager", return_value=manager): diff --git a/tools/mcp_dashboard_oauth.py b/tools/mcp_dashboard_oauth.py index 31663e0f37..1d1447146a 100644 --- a/tools/mcp_dashboard_oauth.py +++ b/tools/mcp_dashboard_oauth.py @@ -40,6 +40,7 @@ class DashboardOAuthFlow: _lock: threading.Lock = field(default_factory=threading.Lock, init=False, repr=False) async def publish_authorization_url(self, url: str) -> None: + """Record the SDK's authorization URL (with its ``state``) for the dashboard to show.""" state = parse_qs(urlparse(url).query).get("state", [None])[0] if not state: raise ValueError("OAuth authorization URL did not include state") @@ -51,10 +52,15 @@ class DashboardOAuthFlow: self.status = "authorization_required" self._authorization_ready.set() + @staticmethod + async def _await_event(event: threading.Event, timeout: float, message: str) -> None: + if not await asyncio.to_thread(event.wait, timeout): + raise TimeoutError(message) + async def wait_for_authorization_url(self, timeout: float = 30.0) -> str: - ready = await asyncio.to_thread(self._authorization_ready.wait, timeout) - if not ready: - raise TimeoutError("Timed out waiting for MCP authorization URL") + await self._await_event( + self._authorization_ready, timeout, "Timed out waiting for MCP authorization URL" + ) if not self.authorization_url: raise RuntimeError(self.error or "MCP OAuth flow ended before authorization") return self.authorization_url @@ -66,6 +72,7 @@ class DashboardOAuthFlow: state: str | None, error: str | None, ) -> None: + """Hand the browser redirect to the waiting flow; ``state`` must match exactly.""" with self._lock: if self._callback_ready.is_set(): raise ValueError("OAuth callback already received") @@ -84,9 +91,9 @@ class DashboardOAuthFlow: self._callback_ready.set() async def wait_for_callback(self, timeout: float = 300.0) -> tuple[str, str | None]: - ready = await asyncio.to_thread(self._callback_ready.wait, timeout) - if not ready: - raise TimeoutError("Timed out waiting for MCP OAuth callback") + await self._await_event( + self._callback_ready, timeout, "Timed out waiting for MCP OAuth callback" + ) if self._callback_error: raise RuntimeError(f"OAuth authorization failed: {self._callback_error}") if self._callback is None: diff --git a/tools/mcp_oauth.py b/tools/mcp_oauth.py index 134b260a26..470586ceaf 100644 --- a/tools/mcp_oauth.py +++ b/tools/mcp_oauth.py @@ -2,27 +2,13 @@ """ MCP OAuth 2.1 Client Support -Implements the browser-based OAuth 2.1 authorization code flow with PKCE -for MCP servers that require OAuth authentication instead of static bearer -tokens. - -Uses the MCP Python SDK's ``OAuthClientProvider`` (an ``httpx.Auth`` subclass) -which handles discovery, client identification, PKCE, token exchange, -refresh, and step-up authorization automatically. - -Client identification follows the MCP 2026-07-28 spec: when the authorization -server advertises ``client_id_metadata_document_supported``, the SDK uses the -URL of Hermes' published Client ID Metadata Document (CIMD) as the -``client_id``; otherwise it falls back to RFC 7591 dynamic client registration, -which that spec revision deprecated. - -This module provides the glue: - - ``HermesTokenStorage``: persists tokens/client-info to disk so they - survive across process restarts. - - Callback server: ephemeral localhost HTTP server to capture the OAuth - redirect with the authorization code. - - ``build_oauth_auth()``: entry point called by ``mcp_tool.py`` that wires - everything together and returns the ``httpx.Auth`` object. +Browser-based OAuth 2.1 authorization-code flow with PKCE for MCP servers. 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). Configuration in config.yaml:: @@ -44,6 +30,7 @@ Configuration in config.yaml:: import asyncio import contextvars +import importlib.util as _importlib_util import json import logging import os @@ -64,43 +51,31 @@ from hermes_constants import secure_parent_dir logger = logging.getLogger(__name__) -# --------------------------------------------------------------------------- -# Lazy imports -- MCP SDK with OAuth support is optional -# --------------------------------------------------------------------------- - -# Availability is detected WITHOUT importing the mcp SDK (which costs -# ~170 ms at module load). The actual classes are imported lazily on first -# use via _ensure_sdk_loaded(); the module-level names below are kept as -# placeholders so tests can patch them (patch.object requires the attribute -# to exist on the module). -import importlib.util as _importlib_util +# 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. _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") -# Lazily-bound SDK names (rebound by _ensure_sdk_loaded on first use). -# Annotated ``Any`` so quoted type annotations elsewhere in the file remain -# valid for static checkers while the runtime value starts as None. OAuthClientProvider: Any = None OAuthClientInformationFull: Any = None OAuthClientMetadata: Any = None OAuthMetadata: Any = None OAuthToken: Any = None -# Cache of the real SDK classes so a test that temporarily patches one of the -# module-level names (and restores it to None afterwards) doesn't strand the -# module in a broken state. +# Cache of the real SDK classes so a test that patches a name and restores it +# to None afterwards doesn't strand the module in a broken state. _SDK_CLASSES: dict[str, Any] = {} _SDK_LOAD_FAILED = False def _ensure_sdk_loaded() -> bool: - """Import the MCP SDK OAuth classes on first use and bind module globals. + """Bind the SDK OAuth classes into module globals; True when available. - Returns True when the SDK classes are available. Module-level names that - have been replaced (e.g. patched by tests) are left untouched; only names - that are currently ``None`` are (re)bound to the real SDK classes. + Only names that are currently ``None`` are (re)bound, so test patches of + the module-level names are left untouched. """ global _SDK_LOAD_FAILED, _OAUTH_AVAILABLE if _SDK_LOAD_FAILED: @@ -132,69 +107,53 @@ def _ensure_sdk_loaded() -> bool: g[_name] = _cls return True + +def _sdk_class(name: str) -> Any: + """Return the (possibly test-patched) SDK class bound under *name*, or None.""" + if globals().get(name) is None and not _ensure_sdk_loaded(): + return None + return globals().get(name) + + try: from pydantic import AnyUrl except ImportError: AnyUrl = None # type: ignore[assignment, misc] -# --------------------------------------------------------------------------- -# Exceptions -# --------------------------------------------------------------------------- - - class OAuthNonInteractiveError(RuntimeError): """Raised when OAuth requires browser interaction in a non-interactive env.""" -# --------------------------------------------------------------------------- -# Module-level state -# --------------------------------------------------------------------------- - -# Port used by the most recent build_oauth_auth() call. Exposed so that -# tests can verify the callback server and the redirect_uri share a port. +# 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 gate for OAuth stdin prompts. A ContextVar (NOT threading.local) -# is required: background MCP discovery sets this on the discovery thread, but -# the actual connect+OAuth runs on the dedicated `mcp-event-loop` thread via -# run_coroutine_threadsafe. asyncio copies the *calling context* into the -# scheduled coroutine, so a ContextVar propagates across that boundary while a -# threading.local would not — see #35927. Default True (interactive allowed). + +# 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[bool]" = contextvars.ContextVar( "_oauth_interactive_enabled", default=True ) - -# Forces _is_interactive() past the stdin-TTY check for flows driven from a -# GUI (dashboard/desktop REST): the browser + localhost callback server do all -# the work there, and the stdin paste fallback degrades harmlessly (EOF is -# swallowed by _paste_callback_reader). Suppression still 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[bool]" = 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"}) - -# Sentinel value written to result["error"] when the user skipped via stdin. -# _wait_for_callback maps this to OAuthNonInteractiveError ("user_skipped") -# so the MCP setup path treats it as a non-fatal "continue without this -# server" rather than a hard failure. +# 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". _USER_SKIPPED_SENTINEL = "__hermes_user_skipped__" -# --------------------------------------------------------------------------- -# Helpers -# --------------------------------------------------------------------------- - - def _get_token_dir(hermes_home: str | Path | None = None) -> Path: - """Return the directory for MCP OAuth token files. - - Uses HERMES_HOME so each profile gets its own OAuth tokens. - Layout: ``HERMES_HOME/mcp-tokens/`` - """ + """``HERMES_HOME/mcp-tokens/`` — per-profile token directory.""" from hermes_constants import get_hermes_home base = Path(hermes_home) if hermes_home is not None else Path(get_hermes_home()) @@ -206,38 +165,23 @@ def _safe_filename(name: str) -> str: return re.sub(r"[^\w\-]", "_", name).strip("_")[:128] or "default" -def _find_free_port() -> int: - """Find an available TCP port on localhost.""" - with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: - s.bind(("127.0.0.1", 0)) - return s.getsockname()[1] - - -# Bound-but-not-listening sockets reserved for pending OAuth callback flows, -# keyed by port. Holding the socket from port-selection time until -# _wait_for_callback adopts it closes the TOCTOU window where another process -# could grab the port between _find_free_port() closing its probe socket and -# HTTPServer binding minutes later (#22161). Bounded FIFO so repeated -# build_oauth_auth calls (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 ``_wait_for_callback`` adopts it. + """Hold *sock* bound to *port* until the callback waiter adopts it. - Pinned CIMD sockets are never evicted: the published metadata document - only declares the pinned ports, so losing one mid-flow silently converts - a pinned reservation back into a stealable window — the exact race the - parking exists to prevent (#22161). The FIFO cap applies to ephemeral - reservations only; the pinned range is already bounded by ``_CIMD_PORTS``. + 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. """ - # Evict oldest ephemeral reservations past the cap (dict preserves - # insertion order). while len(_reserved_sockets) >= _MAX_RESERVED_SOCKETS: - stale_port = next( - (p for p in _reserved_sockets if p not in _CIMD_PORTS), None - ) + 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 stale = _reserved_sockets.pop(stale_port, None) @@ -250,50 +194,55 @@ def _park_reserved_socket(port: int, sock: socket.socket) -> None: _reserved_sockets[port] = sock -def _reserve_callback_port() -> int: - """Pick an ephemeral callback port and keep its socket bound. - - Returns the port. The bound (not yet listening) socket is parked in - ``_reserved_sockets`` so no other process can bind the port before - ``_wait_for_callback`` adopts it. Adoption (or ``server_close``) owns - the socket's lifetime from there. - """ - s = socket.socket(socket.AF_INET, socket.SOCK_STREAM) +def _bind_reserved(port: int) -> int | None: + """Bind ``127.0.0.1:port`` (0 = ephemeral) and park it; None if taken.""" + sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) try: - s.bind(("127.0.0.1", 0)) + sock.bind(("127.0.0.1", port)) except OSError: - s.close() + sock.close() + if port: + return None raise - port = s.getsockname()[1] - _park_reserved_socket(port, s) - return port + bound = sock.getsockname()[1] + _park_reserved_socket(bound, sock) + return bound -def _cached_redirect_port(storage: "HermesTokenStorage | None") -> int | None: - """Return the loopback callback port from cached client registration. +def _reserve_callback_port() -> int: + """Pick an ephemeral callback port and keep its socket bound (parked).""" + return _bind_reserved(0) # type: ignore[return-value] # port 0 never returns None - OAuth providers bind a dynamically-registered ``client_id`` to the exact - redirect URI that was registered with it. If Hermes restarts and chooses a - new random callback port while reusing the stored ``client_id``, providers - such as Summ reject the authorization request with ``redirect_uri does not - match any registered URIs``. Reusing the cached redirect port keeps the - authorization request consistent with the stored client registration. - """ + +def _cached_client_info(storage: "HermesTokenStorage | None") -> dict | None: + """The on-disk client registration for *storage*, or None.""" if storage is None: return None - try: - data = _read_json(storage._client_info_path()) + return _read_json(storage._client_info_path()) except (AttributeError, TypeError, ValueError): return None - if not data: - return None - for uri in data.get("redirect_uris") or []: + +def _cached_redirect_uris(storage: "HermesTokenStorage | None"): + """Yield ``(raw_uri, parsed)`` for each redirect URI in the cached registration.""" + for uri in (_cached_client_info(storage) or {}).get("redirect_uris") or []: try: parsed = urlparse(str(uri)) except (TypeError, ValueError): continue + yield str(uri), parsed + + +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; picking a new random port on restart + while reusing the stored ``client_id`` gets ``redirect_uri does not match + any registered URIs``. + """ + for _uri, parsed in _cached_redirect_uris(storage): if ( parsed.scheme == "http" and parsed.hostname in {"127.0.0.1", "localhost"} @@ -305,25 +254,15 @@ def _cached_redirect_port(storage: "HermesTokenStorage | None") -> int | None: def _cached_redirect_uri(storage: "HermesTokenStorage | None") -> str | None: - """Return a cached non-loopback redirect URI, if one was registered.""" - if storage is None: - return None - try: - data = _read_json(storage._client_info_path()) - except (AttributeError, TypeError, ValueError): - return None - for uri in (data or {}).get("redirect_uris") or []: - try: - parsed = urlparse(str(uri)) - except (TypeError, ValueError): - continue + """A cached non-loopback (https) redirect URI, if one was registered.""" + for uri, parsed in _cached_redirect_uris(storage): if parsed.scheme == "https" and parsed.netloc: - return str(uri) + return uri return None def _is_interactive() -> bool: - """Return True if we can reasonably expect to interact with a user.""" + """True if we can reasonably expect to interact with a user.""" if not _oauth_interactive_enabled.get(): return False if _oauth_interactive_forced.get(): @@ -335,11 +274,10 @@ def _is_interactive() -> bool: def _raise_if_non_interactive(lead: str) -> None: - """Raise ``OAuthNonInteractiveError`` unless an interactive session exists. + """Raise ``OAuthNonInteractiveError`` unless interactive. - ``lead`` is the boundary-specific first sentence; this helper appends the - shared, actionable ``hermes mcp login`` next-step so the guidance wording - lives in one place across every non-interactive OAuth boundary (#57836). + ``lead`` is the boundary-specific first sentence; the shared + ``hermes mcp login`` next-step wording lives here only. """ if not _is_interactive(): raise OAuthNonInteractiveError( @@ -350,43 +288,36 @@ def _raise_if_non_interactive(lead: str) -> None: @contextmanager -def force_interactive_oauth(): - """Treat the current execution context as interactive despite no TTY. - - For GUI-driven auth (dashboard/desktop REST endpoint): the user IS present - — just not on stdin. Opens the browser + localhost callback flow that the - TTY heuristic would otherwise refuse. Same ContextVar propagation story as - suppress_interactive_oauth() (#35927). - """ - token = _oauth_interactive_forced.set(True) +def _contextvar_set(var: "contextvars.ContextVar", value): + token = var.set(value) try: yield finally: - _oauth_interactive_forced.reset(token) + var.reset(token) + + +def force_interactive_oauth(): + """Treat the current context as interactive despite no TTY (GUI-driven auth). + + The user IS present — just not on stdin. Propagates across the MCP + event-loop thread boundary like ``suppress_interactive_oauth``. + """ + return _contextvar_set(_oauth_interactive_forced, True) -@contextmanager def suppress_interactive_oauth(): """Disable stdin-based OAuth prompts for the current execution context. - Uses a ContextVar so the suppression propagates from a background-discovery - thread onto the coroutine scheduled (via run_coroutine_threadsafe) on the - dedicated MCP event-loop thread — where the OAuth callback actually runs - (#35927). A threading.local would not cross that thread boundary. + ContextVar-based so suppression set on a background-discovery thread + reaches the coroutine scheduled on the MCP event-loop thread. """ - token = _oauth_interactive_enabled.set(False) - try: - yield - finally: - _oauth_interactive_enabled.reset(token) + return _contextvar_set(_oauth_interactive_enabled, False) def _can_open_browser() -> bool: - """Return True if opening a browser is likely to work.""" - # Explicit SSH session → no local display + """True if opening a browser is likely to work.""" if os.environ.get("SSH_CLIENT") or os.environ.get("SSH_TTY"): - return False - # macOS and Windows usually have a display + return False # explicit SSH session → no local display if os.name == "nt": return True try: @@ -394,10 +325,7 @@ def _can_open_browser() -> bool: return True except AttributeError: pass - # Linux/other posix: need DISPLAY or WAYLAND_DISPLAY - if os.environ.get("DISPLAY") or os.environ.get("WAYLAND_DISPLAY"): - return True - return False + return bool(os.environ.get("DISPLAY") or os.environ.get("WAYLAND_DISPLAY")) def _read_json(path: Path) -> dict | None: @@ -412,30 +340,20 @@ def _read_json(path: Path) -> dict | None: def _write_json(path: Path, data: dict) -> None: - """Write a dict as JSON with restricted permissions (0o600). + """Atomically write *data* as JSON created at 0o600. - Uses ``os.open`` with ``O_EXCL`` and an explicit mode so the file is - created atomically at 0o600. The previous ``write_text`` + post-write - ``chmod`` opened a TOCTOU window where the temp file briefly inherited - the process umask (commonly 0o644 = world-readable), exposing OAuth - tokens to other local users between create and chmod. Mirrors the fix - in ``agent/google_oauth.py`` (#19673). + ``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). """ path.parent.mkdir(parents=True, exist_ok=True) - # Tighten parent dir to 0o700 so siblings can't traverse to the creds. - # No-op on Windows (POSIX mode bits aren't enforced); ignore failures. - # secure_parent_dir refuses to chmod /, top-level dirs, or the - # hermes-agent install tree (#25821, #93050). secure_parent_dir(path) - # Per-process random suffix avoids collisions between concurrent - # writers and stale leftovers from a prior crashed write. + # Per-process random suffix avoids collisions between concurrent writers + # and stale leftovers from a crashed write. tmp = path.with_suffix(f".tmp.{os.getpid()}.{secrets.token_hex(4)}") try: - fd = os.open( - str(tmp), - os.O_WRONLY | os.O_CREAT | os.O_EXCL, - stat.S_IRUSR | stat.S_IWUSR, - ) + fd = os.open(str(tmp), os.O_WRONLY | os.O_CREAT | os.O_EXCL, stat.S_IRUSR | stat.S_IWUSR) with os.fdopen(fd, "w", encoding="utf-8") as fh: json.dump(data, fh, indent=2, default=str) fh.flush() @@ -469,152 +387,133 @@ class HermesTokenStorage: self._server_name = _safe_filename(server_name) self._hermes_home = Path(hermes_home) if hermes_home is not None else None + def _path(self, suffix: str) -> Path: + return _get_token_dir(self._hermes_home) / f"{self._server_name}{suffix}" + def _tokens_path(self) -> Path: - return _get_token_dir(self._hermes_home) / f"{self._server_name}.json" + return self._path(".json") def _client_info_path(self) -> Path: - return _get_token_dir(self._hermes_home) / f"{self._server_name}.client.json" + return self._path(".client.json") def _meta_path(self) -> Path: - return _get_token_dir(self._hermes_home) / f"{self._server_name}.meta.json" + return self._path(".meta.json") def _cimd_rejected_path(self) -> Path: - return _get_token_dir(self._hermes_home) / f"{self._server_name}.cimd-off" + return self._path(".cimd-off") + + def _state_paths(self) -> tuple[Path, Path, Path]: + return self._tokens_path(), self._client_info_path(), self._meta_path() + + @staticmethod + def _load_model(path: Path, sdk_name: str, label: str, fixup=None): + """Read *path* into SDK model *sdk_name*; None if absent, no SDK, or corrupt. + + ``fixup(data)`` may rewrite the raw dict before validation. + """ + data = _read_json(path) + cls = _sdk_class(sdk_name) if data is not None else None + if cls is None: + return None + if fixup is not None: + fixup(data) + try: + return cls.model_validate(data) + except (ValueError, TypeError, KeyError) as exc: + logger.warning("Corrupt %s at %s -- ignoring: %s", label, path, exc) + return None # -- tokens ------------------------------------------------------------ - async def get_tokens(self) -> "OAuthToken | None": - data = _read_json(self._tokens_path()) - if data is None: - return None - if OAuthToken is None and not _ensure_sdk_loaded(): - return None - # Hermes records an absolute wall-clock ``expires_at`` alongside the - # SDK's serialized token (see ``set_tokens``). On read we rewrite - # ``expires_in`` to the remaining seconds so the SDK's downstream - # ``update_token_expiry`` computes the correct absolute time and - # ``is_token_valid()`` correctly reports False for tokens that - # expired while the process was down. - # - # Legacy token files (pre-Fix-A) have ``expires_in`` but no - # ``expires_at``. We fall back to the file's mtime as a best-effort - # wall-clock proxy for when the token was written: if (mtime + - # expires_in) is in the past, clamp ``expires_in`` to zero so the - # SDK refreshes before the first request. This self-heals one-time - # on the next successful ``set_tokens``, which writes the new - # ``expires_at`` field. The stored ``expires_at`` is stripped before - # model_validate because it's not part of the SDK's OAuthToken schema. + 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``). + """ absolute_expiry = data.pop("expires_at", None) if absolute_expiry is not None: data["expires_in"] = int(max(absolute_expiry - time.time(), 0)) elif data.get("expires_in") is not None: try: - file_mtime = self._tokens_path().stat().st_mtime - except OSError: - file_mtime = None - if file_mtime is not None: - try: - implied_expiry = file_mtime + int(data["expires_in"]) - data["expires_in"] = int(max(implied_expiry - time.time(), 0)) - except (TypeError, ValueError): - pass - try: - return OAuthToken.model_validate(data) - except (ValueError, TypeError, KeyError) as exc: - logger.warning("Corrupt tokens at %s -- ignoring: %s", self._tokens_path(), exc) - return None + implied_expiry = self._tokens_path().stat().st_mtime + int(data["expires_in"]) + data["expires_in"] = int(max(implied_expiry - time.time(), 0)) + except (OSError, TypeError, ValueError): + pass + + async def get_tokens(self) -> "OAuthToken | None": + 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`` so a process restart can - # reconstruct the correct remaining TTL. Without this the MCP SDK's - # ``_initialize`` reloads a relative ``expires_in`` which has no - # wall-clock reference, leaving ``context.token_expiry_time=None`` - # and ``is_token_valid()`` falsely reporting True. See Fix A in - # ``mcp-oauth-token-diagnosis`` skill + Claude Code's - # ``OAuthTokens.expiresAt`` persistence (auth.ts ~180). + # 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. expires_in = payload.get("expires_in") if expires_in is not None: try: payload["expires_at"] = time.time() + int(expires_in) except (TypeError, ValueError): - # Mock tokens or unusual shapes: skip the expires_at write - # rather than fail persistence. - pass + pass # mock tokens / odd shapes: skip rather than fail persistence _write_json(self._tokens_path(), payload) logger.debug("OAuth tokens saved for %s", self._server_name) # -- 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. + """ + 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 + return False + async def get_client_info(self) -> "OAuthClientInformationFull | None": - data = _read_json(self._client_info_path()) - if data is None: - return None - if OAuthClientInformationFull is None and not _ensure_sdk_loaded(): - return None - try: - info = OAuthClientInformationFull.model_validate(data) - # Some dynamic registration providers (notably Supabase MCP) return - # a client_secret but omit token_endpoint_auth_method. The MCP SDK - # defaults that missing field to "none", which causes token exchange - # to omit client_secret and fail with "Required parameter: client_secret". - # If a secret is present, use client_secret_post unless the provider - # explicitly saved a different method. - if getattr(info, "client_secret", None) and data.get("token_endpoint_auth_method") in (None, "none", ""): - data["token_endpoint_auth_method"] = "client_secret_post" - info = OAuthClientInformationFull.model_validate(data) - _write_json(self._client_info_path(), info.model_dump(mode="json", exclude_none=True)) - return info - except (ValueError, TypeError, KeyError) as exc: - logger.warning("Corrupt client info at %s -- ignoring: %s", self._client_info_path(), exc) - return None + coerced: list[bool] = [] + info = self._load_model( + self._client_info_path(), "OAuthClientInformationFull", "client info", + lambda data: coerced.append(self._coerce_secret_auth_method(data)), + ) + 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)) + return info async def set_client_info(self, client_info: "OAuthClientInformationFull") -> None: data = client_info.model_dump(mode="json", exclude_none=True) - # Supabase MCP dynamic client registration returns a client_secret but - # omits token_endpoint_auth_method. The MCP SDK defaults that to - # "none", which makes token exchange omit client_secret and loops the - # browser authorization page. Persist the effective method immediately - # so this flow and subsequent retries use client_secret_post. - if data.get("client_secret") and data.get("token_endpoint_auth_method") in (None, "none", ""): - data["token_endpoint_auth_method"] = "client_secret_post" + 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 MCP SDK keeps discovered ``OAuthMetadata`` (token endpoint URL, - # etc.) in memory only. Persisting it here lets a restarted process - # refresh tokens without re-running metadata discovery. Without this, - # cold-start refresh requests fall back to the SDK's guessed - # ``{server_url}/token`` which returns 404 on most real providers and - # forces a full browser re-authorization. + # 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(exclude_none=True, mode="json")) logger.debug("OAuth metadata saved for %s", self._server_name) def load_oauth_metadata(self) -> "OAuthMetadata | None": - data = _read_json(self._meta_path()) - if data is None: - return None - if OAuthMetadata is None and not _ensure_sdk_loaded(): - return None - try: - return OAuthMetadata.model_validate(data) - except (ValueError, TypeError, KeyError) as exc: - logger.warning("Corrupt OAuth metadata at %s -- ignoring: %s", self._meta_path(), exc) - return None + return self._load_model(self._meta_path(), "OAuthMetadata", "OAuth metadata") # -- CIMD refusal ------------------------------------------------------ def mark_cimd_rejected(self) -> None: - """Record that this server refused our Client ID Metadata Document. + """Durably record that this server refused our Client ID Metadata Document. - Without a durable marker the in-memory fallback in - ``mcp_oauth_manager`` only holds for the current process, so every - restart re-presents a client_id the server has already fetched and - refused. Cleared by ``remove()``, i.e. by ``hermes mcp login`` / - ``hermes mcp remove``, so a fixed document gets another chance. + 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: @@ -631,23 +530,14 @@ class HermesTokenStorage: def remove(self) -> None: """Delete all stored OAuth state for this server.""" - for p in ( - self._tokens_path(), - self._client_info_path(), - self._meta_path(), - self._cimd_rejected_path(), - ): + for p in (*self._state_paths(), self._cimd_rejected_path()): p.unlink(missing_ok=True) def snapshot(self) -> dict[str, bytes]: - """Capture on-disk OAuth state so a failed re-auth can restore it. - - Maps filename -> bytes for whichever of the three state files exist. - Feed back to ``restore()`` to undo an intervening ``remove()`` when a - re-authentication attempt fails, so a still-valid token isn't destroyed. - """ + """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._tokens_path(), self._client_info_path(), self._meta_path()): + for p in self._state_paths(): try: snap[p.name] = p.read_bytes() except OSError: @@ -656,10 +546,7 @@ class HermesTokenStorage: def restore(self, snapshot: dict[str, bytes], *, only_if_absent: bool = False) -> None: """Revert to a snapshot without overwriting a concurrent successful write.""" - if only_if_absent and any( - path.exists() - for path in (self._tokens_path(), self._client_info_path(), self._meta_path()) - ): + if only_if_absent and any(path.exists() for path in self._state_paths()): logger.info( "Skipping OAuth rollback for %s because newer state exists", self._server_name, @@ -673,32 +560,17 @@ class HermesTokenStorage: for fname, data in snapshot.items(): path = token_dir / fname try: - fd = os.open( - str(path), - os.O_WRONLY | os.O_CREAT | os.O_TRUNC, - stat.S_IRUSR | stat.S_IWUSR, - ) + fd = os.open(str(path), os.O_WRONLY | os.O_CREAT | os.O_TRUNC, stat.S_IRUSR | stat.S_IWUSR) with os.fdopen(fd, "wb") as fh: fh.write(data) except OSError as exc: logger.warning("Failed to restore OAuth state %s: %s", fname, exc) def poison_client_registration(self) -> bool: - """Discard a dead dynamically-registered client so it gets re-created. - - Called when the IdP rejects our cached ``client_id`` with - ``invalid_client`` on the token endpoint — proof the server-side - registration is gone (IdP redeploy / DB wipe / rebrand). Deleting - ``client.json`` makes the MCP SDK's ``async_auth_flow`` take the - ``if not client_info`` branch and re-run RFC 7591 dynamic client - registration on the next flow. The stale ``meta.json`` is dropped - too so discovery re-runs against a freshly fetched document. - - Tokens are intentionally left in place — the subsequent - re-authorization overwrites them, and keeping them avoids losing a - still-valid refresh token if the re-registration never completes. - - A single ``.bak`` copy of the client file is kept for recovery. + """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() @@ -719,23 +591,20 @@ class HermesTokenStorage: return True def has_cached_tokens(self) -> bool: - """Return True if we have tokens on disk (may be expired).""" + """True if we have tokens on disk (may be expired).""" return self._tokens_path().exists() # --------------------------------------------------------------------------- -# Callback handler factory -- each invocation gets its own result dict +# 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 changed ``callback_handler``'s contract from a - ``tuple[str, str | None]`` to an ``AuthorizationCodeResult`` model, and the - SDK now reads ``result.state`` / ``result.iss`` off it — a tuple raises - ``AttributeError`` mid-flow. Fall back to the tuple when the model is - absent so the handler still satisfies an older SDK. + mcp 2.0's ``callback_handler`` contract returns an ``AuthorizationCodeResult`` + (the SDK reads ``.state`` / ``.iss`` off it); older SDKs take a tuple. """ try: from mcp.shared.auth import AuthorizationCodeResult @@ -744,42 +613,46 @@ def _authorization_code_result(code: str, state: "str | None", iss: "str | None" return AuthorizationCodeResult(code=code, state=state, iss=iss) +def _parse_redirect_query(query: str) -> dict[str, Any]: + """Extract code/state/error/iss from a redirect query string. + + ``iss`` is the RFC 9207 authorization-response issuer: mcp 2.0 rejects a + response that omits it when the server advertised + ``authorization_response_iss_parameter_supported``, so it must be kept. + """ + params = parse_qs(query) + return {k: params.get(k, [None])[0] for k in ("code", "state", "error", "iss")} + + +def _result_taken(result: dict) -> bool: + return result.get("auth_code") is not None or result.get("error") is not None + + +def _fill_result(result: dict, parsed: dict[str, Any]) -> None: + result["auth_code"] = parsed["code"] + result["state"] = parsed["state"] + result["error"] = parsed["error"] + result["iss"] = parsed["iss"] + + def _make_callback_handler() -> tuple[type, dict]: """Create a per-flow callback HTTP handler class with its own result dict. - Returns ``(HandlerClass, result_dict)`` where *result_dict* is a mutable - dict that the handler writes ``auth_code`` and ``state`` into when the - OAuth redirect arrives. Each call returns a fresh pair so concurrent + Each call returns a fresh ``(HandlerClass, result_dict)`` so concurrent flows don't stomp on each other. """ - result: dict[str, Any] = { - "auth_code": None, "state": None, "error": None, "iss": None, - } + result: dict[str, Any] = {"auth_code": None, "state": None, "error": None, "iss": None} class _Handler(BaseHTTPRequestHandler): def do_GET(self) -> None: # noqa: N802 - params = parse_qs(urlparse(self.path).query) - code = params.get("code", [None])[0] - state = params.get("state", [None])[0] - error = params.get("error", [None])[0] - # RFC 9207 authorization-response issuer. mcp 2.0 validates it - # against the discovered metadata and *rejects* a response that - # omits it when the authorization server advertised - # `authorization_response_iss_parameter_supported`, so dropping it - # here would break login against those providers. - iss = params.get("iss", [None])[0] - - result["auth_code"] = code - result["state"] = state - result["error"] = error - result["iss"] = iss - + parsed = _parse_redirect_query(urlparse(self.path).query) + _fill_result(result, parsed) body = ( "
You can close this tab and return to Hermes.
" - ) if code else ( + ) if parsed["code"] else ( "Error: {error or 'unknown'}
" + f"Error: {parsed['error'] or 'unknown'}