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 = ( "

Authorization Successful

" "

You can close this tab and return to Hermes.

" - ) if code else ( + ) if parsed["code"] else ( "

Authorization Failed

" - f"

Error: {error or 'unknown'}

" + f"

Error: {parsed['error'] or 'unknown'}

" ) self.send_response(200) self.send_header("Content-Type", "text/html; charset=utf-8") @@ -792,29 +665,99 @@ def _make_callback_handler() -> tuple[type, dict]: return _Handler, result +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. + """ + try: + line = sys.stdin.readline() + except (KeyboardInterrupt, OSError, ValueError): + return + 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( + " OAuth skipped. Run `hermes mcp login ` later to " + "authenticate, or set ``enabled: false`` on that server in " + "config.yaml to disable persistently.", + file=sys.stderr, + ) + return + + # Full URL or "?code=...": take everything after the first "?". + query = line.split("?", 1)[1] if "?" in line else line + if query.startswith("?"): + query = query[1:] + try: + parsed = _parse_redirect_query(query) + except (ValueError, TypeError): + print(" Could not parse pasted input as an OAuth redirect — ignoring.", file=sys.stderr) + return + if not parsed["code"] and not parsed["error"]: + print(" Pasted input did not contain ``code=`` or ``error=`` — ignoring.", file=sys.stderr) + return + if _result_taken(result): # one more race-check before writing + return + _fill_result(result, parsed) + if parsed["code"]: + print(" Got authorization code from paste — completing flow.", file=sys.stderr) + + # --------------------------------------------------------------------------- # Async redirect + callback handlers for OAuthClientProvider # --------------------------------------------------------------------------- +def _print_ssh_hint(port: int, redirect_uri: str | None) -> None: + """Remote-session guidance printed under the authorization URL.""" + if redirect_uri: + # A proxy callback (e.g. Tailscale Funnel) forwards the redirect to the + # listener on this machine, so no tunnel/paste is needed. + print( + f" Remote session detected. After you authorize, the provider redirects to\n" + f" {redirect_uri}\n" + f" which forwards to the callback listener on this machine — no SSH tunnel needed.\n", + file=sys.stderr, + ) + elif port: + # Loopback default: the redirect reaches the *remote* machine's listener, + # not the browser's machine. Paste the redirect URL back, or SSH-forward. + print( + f" Remote session detected. After you authorize, the provider redirects to\n" + f" http://127.0.0.1:{port}/callback\n" + f" which only the listener on THIS machine can receive. Two options:\n" + f"\n" + f" 1. Easiest — when your browser shows a connection error after\n" + f" authorizing, copy the full URL from the address bar and paste\n" + f" it at the prompt below. The pasted ``code=...&state=...`` is\n" + f" enough to complete the flow.\n" + f"\n" + f" 2. Or forward the port first in a separate terminal:\n" + f" ssh -N -L {port}:127.0.0.1:{port} @\n" + f" then open the URL above and let it redirect normally.\n" + f"\n" + f" See: https://hermes-agent.nousresearch.com/docs/guides/oauth-over-ssh\n", + file=sys.stderr, + ) + + def _make_redirect_handler(port: int, redirect_uri: str | None = None): - """Return a redirect handler closure that closes over the given port. + """Return a redirect handler closing over this flow's port. - Using a closure instead of reading the module-level ``_oauth_port`` avoids - cross-server state pollution when multiple MCP servers run OAuth - concurrently (fixes #44588). - - ``redirect_uri`` is the configured proxy callback (e.g. a Tailscale Funnel - URL), or ``None`` for the loopback default. It tailors the remote-session - hint: a proxied callback reaches this machine on its own, so the loopback - SSH-tunnel guidance would be misleading. + A closure (not the module-level ``_oauth_port``) keeps concurrent + server flows isolated. ``redirect_uri`` is a configured proxy callback + (or None for loopback) and only tailors the remote-session hint. """ async def _redirect_handler(authorization_url: str) -> None: - """Show the authorization URL to the user. - - Opens the browser automatically when possible; always prints the URL - as a fallback for headless/SSH/gateway environments. - """ + """Open the browser when possible; always print the URL as a fallback.""" from tools.mcp_dashboard_oauth import get_dashboard_oauth_flow dashboard_flow = get_dashboard_oauth_flow() @@ -822,129 +765,80 @@ def _make_redirect_handler(port: int, redirect_uri: str | None = None): await dashboard_flow.publish_authorization_url(authorization_url) return - # Fail fast at the authorization boundary in non-interactive contexts - # (systemd gateway, cron, background MCP discovery). A cached-but-unusable - # token (expired/revoked, refresh rejected) makes the SDK fall through to - # the authorization-code flow even though build_oauth_auth's token-file - # guard passed. Without this check we would print a URL and launch a - # browser flow no operator can complete, then block in _wait_for_callback - # for the full timeout. Raise before launching so gateway adapters start - # promptly and the caller can skip this server with an actionable warning. - # This intentionally re-checks interactivity here rather than trusting the - # token-file existence guard alone. See #57836. + # 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)." ) - msg = ( + print( f"\n MCP OAuth: authorization required.\n" f" Open this URL in your browser:\n\n" - f" {authorization_url}\n" + f" {authorization_url}\n", + file=sys.stderr, ) - print(msg, file=sys.stderr) + if os.getenv("SSH_CLIENT") or os.getenv("SSH_TTY"): + _print_ssh_hint(port, redirect_uri) - on_ssh = bool(os.getenv("SSH_CLIENT") or os.getenv("SSH_TTY")) - if on_ssh and redirect_uri: - # A configured proxy callback (e.g. Tailscale Funnel) forwards the - # redirect to the listener on this machine, so no tunnel/paste is needed. - print( - f" Remote session detected. After you authorize, the provider redirects to\n" - f" {redirect_uri}\n" - f" which forwards to the callback listener on this machine — no SSH tunnel needed.\n", - file=sys.stderr, - ) - elif on_ssh and port: - # Loopback default: the provider redirects to - # http://127.0.0.1:/callback, which reaches the callback server on - # the *remote* machine — not the user's local machine where the browser - # opened. Two ways out: paste the redirect URL back (default fallback, - # offered by _wait_for_callback on interactive TTYs), or set up an SSH - # port forward so the redirect tunnels through. - print( - f" Remote session detected. After you authorize, the provider redirects to\n" - f" http://127.0.0.1:{port}/callback\n" - f" which only the listener on THIS machine can receive. Two options:\n" - f"\n" - f" 1. Easiest — when your browser shows a connection error after\n" - f" authorizing, copy the full URL from the address bar and paste\n" - f" it at the prompt below. The pasted ``code=...&state=...`` is\n" - f" enough to complete the flow.\n" - f"\n" - f" 2. Or forward the port first in a separate terminal:\n" - f" ssh -N -L {port}:127.0.0.1:{port} @\n" - f" then open the URL above and let it redirect normally.\n" - f"\n" - f" See: https://hermes-agent.nousresearch.com/docs/guides/oauth-over-ssh\n", - file=sys.stderr, - ) - - if _can_open_browser(): - try: - opened = webbrowser.open(authorization_url) - if opened: - print(" (Browser opened automatically.)\n", file=sys.stderr) - else: - print(" (Could not open browser — please open the URL manually.)\n", file=sys.stderr) - except Exception: - print(" (Could not open browser — please open the URL manually.)\n", file=sys.stderr) - else: + if not _can_open_browser(): print(" (Headless environment detected — open the URL manually.)\n", file=sys.stderr) + return + try: + opened = webbrowser.open(authorization_url) + except Exception: + opened = False + if opened: + print(" (Browser opened automatically.)\n", file=sys.stderr) + else: + print(" (Could not open browser — please open the URL manually.)\n", file=sys.stderr) return _redirect_handler -async def _wait_for_callback() -> tuple[str, str | None]: - """Wait for the OAuth callback on the legacy module-level port. +def _start_callback_server(port: int, handler_cls: type) -> HTTPServer: + """Bind the callback listener on *port*, adopting a parked reserved socket. - Kept for backwards compatibility with callers that never went through - :func:`build_oauth_auth`'s per-flow wiring. New code paths receive a - per-flow waiter from :func:`_make_callback_waiter` so concurrent OAuth - flows cannot cross ports (#34260). - - Raises: - RuntimeError: If ``_oauth_port`` has not been set, which would indicate - that ``build_oauth_auth`` was skipped — the asserting form below - was a silent bug when running Python with ``-O``/``-OO``. + 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. """ - if _oauth_port is None: - raise RuntimeError( - "OAuth callback port not set — build_oauth_auth must be called " - "before _wait_for_oauth_callback" - ) - return await _make_callback_waiter(_oauth_port)() + try: + server = HTTPServer(("127.0.0.1", port), handler_cls, bind_and_activate=False) + reserved = _reserved_sockets.pop(port, None) + if reserved is not None: + server.socket.close() + server.socket = reserved + server.server_address = reserved.getsockname() + else: + server.allow_reuse_address = True + 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". + 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." + ) from exc + return server def _make_callback_waiter( port: int, cimd_url: str | None = None, timeout: float = 300.0 ): - """Return a callback waiter bound to a single OAuth flow's port. + """Return a callback waiter bound to one flow's port (isolating concurrent flows). - ``timeout`` bounds how long the waiter polls for the redirect. It used to - be passed to ``OAuthClientProvider(timeout=...)`` as well, but mcp 2.0 - dropped that constructor argument — the wait happens here, so this is now - the only place the configured ``oauth.timeout`` takes effect. - - Closing over the port (instead of reading the module-level - ``_oauth_port``) keeps concurrent OAuth flows isolated: flow A's waiter - listens on flow A's port even when flow B's ``_configure_callback_port`` - overwrites the legacy global afterwards (#34260, the callback-side - sibling of the #44588 redirect-handler fix). - - ``cimd_url`` is the Client ID Metadata Document this flow presents, when - it presents one. It only tailors the timeout message: a server that - fetches the document and refuses it aborts at the *authorization* - endpoint (draft section 5.1), so no redirect ever reaches us and a bare - "timed out" hides the real cause. - - The waiter polls for the redirect without blocking the event loop. On an - interactive TTY it races the HTTP listener against a stdin paste fallback - so users without an SSH tunnel can paste the redirect URL (or just the - ``code=...&state=...`` query string) from a browser on another machine. - - Raises (when awaited): - OAuthNonInteractiveError: If the callback times out (no user present - to complete the browser auth), or in non-interactive contexts. + ``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(): @@ -952,22 +846,14 @@ def _make_callback_waiter( dashboard_flow = get_dashboard_oauth_flow() if dashboard_flow is not None: - # The dashboard flow still speaks the legacy tuple; normalize it - # here so both callback sources hand the SDK one shape. + # 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) - # Reject before binding the callback listener in non-interactive - # contexts. Reaching here means the SDK entered the authorization-code - # flow (a valid or refreshable token would never call the callback - # handler), so a cached token file is present but unusable. Binding the - # listener here would block for the full 300s timeout and — on the next - # connection retry — collide with the still-bound/TIME_WAIT port, - # surfacing as ``OSError: [Errno 98] Address already in use``. Failing - # fast keeps gateway startup independent of an unusable optional MCP - # server. This guard holds "regardless of whether a token file exists" - # — the point the build_oauth_auth token-file guard cannot cover. - # See #57836. + # 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 " @@ -975,50 +861,10 @@ def _make_callback_waiter( ) handler_cls, result = _make_callback_handler() + server = _start_callback_server(port, handler_cls) + threading.Thread(target=server.handle_request, daemon=True).start() - # Start a temporary server on this flow's port, adopting the socket - # reserved at port-selection time when one exists. Holding the bound - # socket from _reserve_callback_port() until here closes the TOCTOU - # window where another process could steal the port between selection - # and bind (#22161). allow_reuse_address is set BEFORE binding (setting - # it after the constructor has already bound is a no-op) so a lingering - # TIME_WAIT socket from a previous flow cannot block the next one - # (#44590). - try: - server = HTTPServer( - ("127.0.0.1", port), handler_cls, bind_and_activate=False - ) - reserved = _reserved_sockets.pop(port, None) - if reserved is not None: - # Adopt the reserved (already bound) socket and start listening. - server.socket.close() - server.socket = reserved - server.server_address = reserved.getsockname() - server.server_activate() - else: - server.allow_reuse_address = True - server.server_bind() - server.server_activate() - except OSError as exc: - # The loopback callback port is genuinely in use: a concurrent OAuth - # flow, a leftover listener, or a fixed `oauth.redirect_port` that - # collided. build_oauth_auth does not start its own callback server, - # so there is nothing to poll here; surface a clear, actionable error - # instead of a misleading "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." - ) from exc - - server_thread = threading.Thread(target=server.handle_request, daemon=True) - server_thread.start() - - # Optional paste-fallback thread: only on interactive TTYs. Reads one - # line from stdin and writes the parsed code/state into the shared - # result dict. The HTTP listener and this thread race for the result; - # whichever fills it first wins. - paste_thread: threading.Thread | None = None + # 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=...`` " @@ -1027,17 +873,12 @@ def _make_callback_waiter( file=sys.stderr, flush=True, ) - paste_thread = threading.Thread( - target=_paste_callback_reader, args=(result,), daemon=True - ) - paste_thread.start() + threading.Thread(target=_paste_callback_reader, args=(result,), daemon=True).start() poll_interval = 0.5 elapsed = 0.0 try: - while elapsed < timeout: - if result["auth_code"] is not None or result["error"] is not None: - break + while elapsed < timeout and not _result_taken(result): await asyncio.sleep(poll_interval) elapsed += poll_interval finally: @@ -1062,102 +903,13 @@ def _make_callback_waiter( "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") - ) + return _authorization_code_result(result["auth_code"], result["state"], result.get("iss")) return _wait -def _paste_callback_reader(result: dict) -> None: - """Read one line from stdin, parse it as an OAuth redirect, write to result. - - Accepts any of: - - Full redirect URL: ``http://127.0.0.1:37949/callback?code=...&state=...`` - - The provider's own callback URL: ``https://mcp.example.com/callback?code=...&state=...`` - - Just the query string: ``?code=...&state=...`` or ``code=...&state=...`` - - A skip token (``skip``, ``cancel``, ``s``, ``n``, ``no``, ``q``, ``quit``) - — exits the OAuth flow cleanly without auth. Caller raises - :class:`OAuthNonInteractiveError` so MCP connection setup treats this - as a non-fatal "user opted out" and continues without that server. - - Failures to parse, EOF, or interrupts are swallowed — this is best-effort - fallback alongside the HTTP listener, which remains the primary path. - """ - try: - line = sys.stdin.readline() - except (KeyboardInterrupt, OSError, ValueError): - return - if not line: - return # EOF - line = line.strip() - if not line: - return - - # Skip if HTTP listener already won. - if result.get("auth_code") is not None or result.get("error") is not None: - return - - # Skip token: user explicitly opted out of authorization. Mark the - # result with a sentinel error string that _wait_for_callback maps - # to OAuthNonInteractiveError (already handled by mcp_tool.py as a - # non-fatal "skip this server and continue startup" path). - if line.lower() in _SKIP_TOKENS: - if result.get("auth_code") is not None or result.get("error") is not None: - return - result["error"] = _USER_SKIPPED_SENTINEL - print( - " OAuth skipped. Run `hermes mcp login ` later to " - "authenticate, or set ``enabled: false`` on that server in " - "config.yaml to disable persistently.", - file=sys.stderr, - ) - return - - # Strip a leading "?" if user pasted just a query string. - query = line - if "?" in line: - # Either a full URL or "?code=...". Take everything after the first "?". - query = line.split("?", 1)[1] - if query.startswith("?"): - query = query[1:] - - try: - params = parse_qs(query) - except (ValueError, TypeError): - print( - " Could not parse pasted input as an OAuth redirect — ignoring.", - file=sys.stderr, - ) - return - - code = params.get("code", [None])[0] - state = params.get("state", [None])[0] - error = params.get("error", [None])[0] - iss = params.get("iss", [None])[0] # RFC 9207 — see _make_callback_handler - - if not code and not error: - print( - " Pasted input did not contain ``code=`` or ``error=`` — ignoring.", - file=sys.stderr, - ) - return - - # One more race-check before writing. - if result.get("auth_code") is not None or result.get("error") is not None: - return - - result["auth_code"] = code - result["state"] = state - result["error"] = error - result["iss"] = iss - if code: - print(" Got authorization code from paste — completing flow.", file=sys.stderr) - - # --------------------------------------------------------------------------- -# OAuth provider compatibility shims +# OAuth provider class (legacy build_oauth_auth path) # --------------------------------------------------------------------------- @@ -1170,94 +922,12 @@ def _get_hermes_oauth_provider_class() -> type | None: return HermesOAuthClientProvider if not _ensure_sdk_loaded(): return None + from tools.mcp_oauth_provider import HermesProviderMixin - class _HermesOAuthClientProvider(OAuthClientProvider): - """OAuth provider with pragmatic fixes for real-world MCP providers. + class _HermesOAuthClientProvider(HermesProviderMixin, OAuthClientProvider): + """SDK provider plus Hermes' token-endpoint fixes (see ``HermesProviderMixin``).""" - Supabase MCP dynamic registration returns ``client_secret`` but omits - ``token_endpoint_auth_method``. The upstream MCP SDK treats the missing - method as ``none`` and therefore omits ``client_secret`` from the token - request, causing Supabase to reject the exchange and the browser to show - the authorization page again. Coerce the in-memory client info right before - token/refresh requests as well as persisting the fixed shape in storage. - - ``token_user_agent`` (from ``oauth.user_agent``) is stamped onto the - token-endpoint requests the SDK builds — some authorization servers - and WAFs reject httpx's default User-Agent there (#75576). - """ - - def __init__(self, *args: Any, token_user_agent: "str | None" = None, **kwargs: Any): - super().__init__(*args, **kwargs) - self._hermes_token_user_agent = token_user_agent - - def _stamp_token_user_agent(self, request): - ua = getattr(self, "_hermes_token_user_agent", None) - if ua: - request.headers["User-Agent"] = ua - return request - - def _coerce_client_secret_post(self) -> None: - info = getattr(self.context, "client_info", None) - if not info or not getattr(info, "client_secret", None): - return - method = getattr(info, "token_endpoint_auth_method", None) - if method not in (None, "none", ""): - return - data = info.model_dump(mode="json", exclude_none=True) - data["token_endpoint_auth_method"] = "client_secret_post" - self.context.client_info = OAuthClientInformationFull.model_validate(data) - - async def _exchange_token_authorization_code(self, *args: Any, **kwargs: Any): - self._coerce_client_secret_post() - request = await super()._exchange_token_authorization_code(*args, **kwargs) - return self._stamp_token_user_agent(request) - - async def _refresh_token(self): - self._coerce_client_secret_post() - request = await super()._refresh_token() - return self._stamp_token_user_agent(request) - - async def _handle_token_response(self, response): - """Accept any 2xx token response and avoid leaking token bodies in errors.""" - if 200 <= response.status_code < 300: - from mcp.client.auth.utils import handle_token_response_scopes - from mcp.client.auth.oauth2 import OAuthTokenError - from httpx import HTTPError - - try: - token_response = await handle_token_response_scopes(response) - except (HTTPError, OAuthTokenError): - raise OAuthTokenError("Invalid token response") from None - self.context.current_tokens = token_response - self.context.update_token_expiry(token_response) - await self.context.storage.set_tokens(token_response) - return - - from mcp.client.auth.oauth2 import OAuthTokenError - - raise OAuthTokenError(f"Token exchange failed ({response.status_code})") - - async def _handle_refresh_response(self, response) -> bool: - """Accept any 2xx refresh response and avoid logging token bodies.""" - if not (200 <= response.status_code < 300): - logger.warning("Token refresh failed: %s", response.status_code) - self.context.clear_tokens() - return False - - from pydantic import ValidationError - from httpx import HTTPError - - try: - content = await response.aread() - token_response = OAuthToken.model_validate_json(content) - self.context.current_tokens = token_response - self.context.update_token_expiry(token_response) - await self.context.storage.set_tokens(token_response) - return True - except (HTTPError, ValidationError): - logger.warning("Invalid refresh response: %s", response.status_code) - self.context.clear_tokens() - return False + _hermes_logger = logger _HermesOAuthClientProvider.__name__ = "HermesOAuthClientProvider" _HermesOAuthClientProvider.__qualname__ = "HermesOAuthClientProvider" @@ -1276,45 +946,32 @@ def remove_oauth_tokens( hermes_home: str | Path | None = None, ) -> None: """Delete stored OAuth tokens and client info for a server.""" - storage = HermesTokenStorage(server_name, hermes_home=hermes_home) - storage.remove() + HermesTokenStorage(server_name, hermes_home=hermes_home).remove() logger.info("OAuth tokens removed for '%s'", server_name) -# --------------------------------------------------------------------------- -# Extracted helpers (Task 3 of MCP OAuth consolidation) -# -# These compose into ``build_oauth_auth`` below, and are also used by -# ``tools.mcp_oauth_manager.MCPOAuthManager._build_provider`` so the two -# construction paths share one implementation. -# --------------------------------------------------------------------------- - - # --------------------------------------------------------------------------- # CIMD -- OAuth Client ID Metadata Documents # -# Under CIMD the client_id IS an HTTPS URL that the authorization server -# fetches to learn our app name, logo and permitted redirect URIs, replacing -# the per-install RFC 7591 registration that the MCP spec deprecated in -# 2026-07-28. The SDK does the protocol work; Hermes only decides whether a -# given 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 -# (draft-ietf-oauth-client-id-metadata-document section 5), and +# 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. The redirect URI in the -# authorization request must be an exact string match against a listed one -# (section 4.2), so a CIMD flow cannot use the ephemeral port Hermes picks -# otherwise. These sit below Linux's 32768 ephemeral floor, so the kernel never -# hands one to an unrelated process. Keep in sync with the document — the -# cross-artifact test in tests/tools/test_mcp_cimd.py enforces that. +# Loopback callback ports declared in that document. The redirect URI must be +# an exact string match against a listed one, 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 @@ -1325,15 +982,10 @@ _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 own validator so we never hand - ``OAuthClientProvider`` a URL its constructor would reject outright. An - ImportError means the SDK predates CIMD, leaving DCR as the only option. - - The SDK checks only the https-scheme and non-root-path halves of - draft-ietf-oauth-client-id-metadata-document section 3. The rest is - enforced here because a URL that violates it fails at the authorization - server, mid-browser-flow, where the user sees an opaque invalid-client - page instead of a config error. + 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 @@ -1343,8 +995,7 @@ def _is_valid_cimd_url(url: str) -> bool: return False try: parsed = urlparse(url) - # Accessing username/password parses the netloc, which can raise. - has_userinfo = bool(parsed.username or parsed.password) + has_userinfo = bool(parsed.username or parsed.password) # netloc parse can raise except ValueError: return False if has_userinfo or parsed.fragment: @@ -1352,11 +1003,10 @@ def _is_valid_cimd_url(url: str) -> bool: return not any(seg in {".", ".."} for seg in parsed.path.split("/")) -# Pinned ports this process has committed to, in the order they were taken. -# A provider is built once per configured OAuth server and keeps its port for -# the process lifetime, so assignments are never released. Includes a port -# restored from a cached client registration, so a sibling server is never -# handed a port another one is already registered on (#34260). +# 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]" = [] @@ -1366,64 +1016,32 @@ def _note_assigned_cimd_port(port: int) -> None: _assigned_cimd_ports.append(port) -def _reserve_cimd_port(port: int) -> bool: - """Bind *port* and park the socket, or return False if it's taken.""" - sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - try: - sock.bind(("127.0.0.1", port)) - except OSError: - sock.close() - return False - _park_reserved_socket(port, sock) - return True - - def _pick_cimd_port() -> int | None: """Reserve a pinned CIMD callback port, or None when none is usable. - Holding the bound socket until ``_wait_for_callback`` adopts it does the - same job here as ``_reserve_callback_port`` does for ephemeral ports - (#22161): a fixed port is just as stealable in the minutes between - selection and the browser redirect arriving. It also makes contention - cooperative — a second profile mid-login, or a sibling server in this - process, finds the bind refused and moves down the range instead of - racing us to the same listener. - - Once every pinned port belongs to this process the range wraps rather - than falling back to DCR: a reused port only bites if both of its - servers authorize at the same moment, and ``_wait_for_callback`` reports - that collision clearly, whereas the DCR fallback would silently use a - mechanism the server may not support at all. + 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 in _assigned_cimd_ports: continue - if _reserve_cimd_port(port): + if _bind_reserved(port) is not None: _assigned_cimd_ports.append(port) return port return _assigned_cimd_ports[0] if _assigned_cimd_ports else None -def _has_cached_client_info(storage: "HermesTokenStorage | None") -> bool: - """True when a client registration is already on disk for this server.""" - if storage is None: - return False - try: - return _read_json(storage._client_info_path()) is not None - except (AttributeError, TypeError, ValueError): - return False - - def _server_declined_cimd(storage: "HermesTokenStorage | None") -> bool: """True when cached metadata shows this server doesn't advertise CIMD. - Pinning a callback port is only needed for a flow that actually ends up - using CIMD, but the SDK decides that during its 401 branch — long after - Hermes has to fix the redirect URI. Cached authorization-server metadata - from an earlier connection closes the gap for every server the user has - already reached: one that never advertised - ``client_id_metadata_document_supported`` keeps the reserved ephemeral - port it has always used, and only a genuinely unknown server pays the + 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. """ if storage is None: @@ -1443,61 +1061,42 @@ def _maybe_use_cimd( ) -> "tuple[str, int] | None": """Return ``(client_id URL, pinned callback port)``, or None to use DCR. - Every early return below is a case where the redirect URI Hermes would - send is not one the published document declares, where the client - identity is already settled, or where the server is known not to want a - document — DCR remains correct in all of them. Passing a metadata URL - anyway would make the SDK present a client_id whose registered redirect - URIs don't match the request, and the authorization server would reject - the flow. + Every early return is a case 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. """ if cfg.get("cimd") is False: return None - url = cfg.get("client_metadata_url") or _CIMD_CLIENT_METADATA_URL if not _is_valid_cimd_url(url): return None - - # A client pinned in config.yaml is the user's explicit choice, and a - # secret means they want a confidential client — the document forbids - # shared secrets (draft section 4.1). + # A pinned client is the user's explicit choice; a secret means a + # confidential client, which the document forbids. if cfg.get("client_id") or cfg.get("client_secret"): return None - - # The document, not the config, supplies the name and auth method the - # server sees, so a caller that set either is asking for an identity CIMD - # cannot present. Figma's DCR name allowlist (applied by - # apply_oauth_provider_defaults) is the in-tree example. - if cfg.get("client_name"): + # The document supplies name and auth method; a caller setting either asks + # for an identity CIMD cannot present (Figma's DCR name allowlist, e.g.). + if cfg.get("client_name") or (cfg.get("token_endpoint_auth_method") or "none") != "none": return None - if (cfg.get("token_endpoint_auth_method") or "none") != "none": - return None - - # Dashboard/desktop flows redirect to the server's own externally - # reachable URL (``/api/mcp/oauth/callback/``), which is - # deployment-specific and 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. from tools.mcp_dashboard_oauth import get_dashboard_oauth_flow if get_dashboard_oauth_flow() is not None: return None - if cfg.get("redirect_uri") or cfg.get("redirect_port"): return None - if (cfg.get("redirect_host") or "127.0.0.1") not in _CIMD_REDIRECT_HOSTS: return None - - # An existing registration is bound to the redirect URI it registered - # with; swapping in a CIMD client_id now would invalidate stored tokens. - if _has_cached_client_info(storage): + # An existing registration is bound to its redirect URI; swapping in a + # CIMD client_id now would invalidate stored tokens. + if _cached_client_info(storage) is not None: return None - if storage is not None and storage.cimd_rejected(): return None - if _server_declined_cimd(storage): return None - port = _pick_cimd_port() if port is None: return None @@ -1507,31 +1106,24 @@ def _maybe_use_cimd( def cimd_provider_kwargs(cfg: dict) -> dict[str, Any]: """``client_metadata_url=`` for ``OAuthClientProvider``, when CIMD applies. - Returned as kwargs rather than a plain value so the argument is omitted - entirely on a DCR flow. An SDK old enough to lack CIMD support — the case - ``_is_valid_cimd_url`` already refuses to produce a URL for — rejects the - keyword outright, and that must not take every other OAuth flow with it. + 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: - """The configured ``oauth.user_agent`` for token-endpoint requests, or None. + """Configured ``oauth.user_agent`` for token-endpoint requests, or None. - Some authorization servers and network protection layers (WAFs) reject - the default python-httpx User-Agent on the token endpoint. The value is - opt-in and per-server; anything that is not a non-empty string is - treated as unset so a null/empty YAML value never sends a blank header. - Applied ONLY to authorization-code exchange and refresh-token requests — - never to MCP traffic or discovery, and no other headers are configurable - (arbitrary token headers risk secrets landing in config.yaml). + 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") - if isinstance(ua, str): - ua = ua.strip() - if ua: - return ua + if isinstance(ua, str) and ua.strip(): + return ua.strip() return None @@ -1539,26 +1131,12 @@ def _configure_callback_port( cfg: dict, storage: "HermesTokenStorage | None" = None, ) -> int: - """Pick or validate the OAuth callback port. + """Resolve the callback port into ``cfg['_resolved_port']`` (0 = non-loopback URI). - Stores the resolved port into ``cfg['_resolved_port']`` so sibling - helpers (and the manager) can read it from the same dict. Returns the - resolved port. - - Port choice precedence: - 1. explicit ``oauth.redirect_port`` config - 2. cached client registration redirect URI port - 3. a pinned CIMD port, when the flow is CIMD-eligible - 4. newly allocated free port - - A CIMD-eligible flow also records the client_id URL in - ``cfg['_cimd_url']`` for the provider constructors to forward. - - NOTE: also sets the legacy module-level ``_oauth_port`` so existing - calls to ``_wait_for_callback`` keep working. The legacy global is - the root cause of issue #5344 (port collision on concurrent OAuth - flows); replacing it with a ContextVar is out of scope for this - consolidation PR. + 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 from tools.mcp_dashboard_oauth import get_dashboard_oauth_flow @@ -1576,44 +1154,25 @@ def _configure_callback_port( cimd = _maybe_use_cimd(cfg, storage) if cimd is not None: cfg["_cimd_url"], port = cimd - cfg["_resolved_port"] = port - _oauth_port = port - return port - requested = int(cfg.get("redirect_port", 0)) - # Precedence: explicit config port → cached client-registration port → - # fresh ephemeral port. The cached port keeps re-auth consistent with the - # redirect URI pinned at dynamic client registration (providers reject a - # mismatched URI). Only a truly fresh ephemeral pick goes through - # _reserve_callback_port(), which keeps the socket bound until - # _wait_for_callback adopts it — closing the select→bind TOCTOU race - # (#22161). Explicit and cached ports are fixed, known values and bind - # via the reuse_address path instead. - port = requested or _cached_redirect_port(storage) or _reserve_callback_port() - # A cached port can be one of the pinned CIMD ports, left behind by an - # earlier CIMD login for this server. Claim it so a sibling server's - # _pick_cimd_port doesn't hand the same port out a second time. - _note_assigned_cimd_port(port) + 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) cfg["_resolved_port"] = port - _oauth_port = port # legacy consumer: _wait_for_callback reads this + _oauth_port = port return port def _resolve_redirect_uri(cfg: dict, port: int) -> str: - """Resolve the OAuth callback URL: configured ``redirect_uri`` or loopback. + """Configured ``redirect_uri`` (proxy, e.g. Tailscale Funnel) or + ``http://:/callback``. - A configured ``redirect_uri`` lets the callback go through a proxy (e.g. a - Tailscale Funnel exposing a public HTTPS URL that forwards to localhost); - otherwise we default to ``http://:/callback``. An empty - value is treated as unset. Both the client metadata and any pre-registered - client info must derive the redirect_uri here so they stay identical — a - mismatch makes the authorization server reject the callback. - - ``redirect_host`` (default ``127.0.0.1``) tweaks only the hostname of the - loopback callback. Some providers' WAFs (e.g. Reclaim.ai's AWS API Gateway) - reject any authorize request whose query string contains a literal - ``127.0.0.1``, returning ``{"message":"Forbidden"}``; ``redirect_host: - localhost`` works around that. The callback listener still binds - ``127.0.0.1`` either way. + 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``. """ configured = cfg.get("redirect_uri") if configured: @@ -1622,16 +1181,9 @@ def _resolve_redirect_uri(cfg: dict, port: int) -> str: return f"http://{host}:{port}/callback" -# Figma's remote MCP (https://mcp.figma.com/mcp) implement RFC 7591 DCR as a -# *name allowlist*, not open registration. POST /v1/oauth/mcp/register returns -# 403 Forbidden for any client_name outside a short fixed set. Empirically (as -# of 2026-07, verified by live call against api.figma.com): -# "Claude Code" → 200 -# "Codex" → 200 -# "Hermes Agent" / "Hermes" / "Cursor" / "VS Code" / … → 403 -# pi-figma-remote-auth and similar tools work around this the same way — register -# under an allowlisted name so the browser flow can start. User can still pin a -# different name via oauth.client_name if Figma ever admits one. +# 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" @@ -1649,9 +1201,7 @@ def _is_figma_remote_mcp( ): return True # Name-only match only when the URL isn't some other host called figma-*. - if "figma" in name and (not url or "figma" in base_url_hostname(url)): - return True - return False + return "figma" in name and (not url or "figma" in base_url_hostname(url)) def apply_oauth_provider_defaults( @@ -1662,9 +1212,8 @@ def apply_oauth_provider_defaults( ) -> dict: """Mutate *cfg* with provider-specific OAuth workarounds. Returns *cfg*. - Call this before :func:`_build_client_metadata` / - :func:`_maybe_preregister_client`. Only fills keys the user left unset — - an explicit ``oauth.client_name`` / ``oauth.scope`` always wins. + 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"): @@ -1677,22 +1226,16 @@ def apply_oauth_provider_defaults( ) if not cfg.get("scope"): cfg["scope"] = _FIGMA_DEFAULT_SCOPE - # Figma's register response advertises token_endpoint_auth_method=none - # *and* returns a client_secret — then the token endpoint rejects the - # exchange with "Client secret is required". Request confidential- - # client registration so the SDK includes client_secret on the token - # POST (auth method client_secret_post). + # 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. if not cfg.get("token_endpoint_auth_method"): cfg["token_endpoint_auth_method"] = "client_secret_post" return cfg def _build_client_metadata(cfg: dict) -> "OAuthClientMetadata": - """Build OAuthClientMetadata from the oauth config dict. - - Requires ``cfg['_resolved_port']`` to have been populated by - :func:`_configure_callback_port` first. - """ + """Build OAuthClientMetadata; requires ``_configure_callback_port`` first.""" port = cfg.get("_resolved_port") if port is None: raise ValueError( @@ -1700,38 +1243,28 @@ def _build_client_metadata(cfg: dict) -> "OAuthClientMetadata": ) if OAuthClientMetadata is None: _ensure_sdk_loaded() - client_name = cfg.get("client_name", "Hermes Agent") - scope = cfg.get("scope") - redirect_uri = _resolve_redirect_uri(cfg, port) - - # Default public client; confidential only when a secret is already known - # or the provider (e.g. Figma) needs confidential-style token posts. - auth_method = cfg.get("token_endpoint_auth_method") - if not auth_method: - auth_method = "client_secret_post" if cfg.get("client_secret") else "none" - + # 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": client_name, - "redirect_uris": [AnyUrl(redirect_uri)], + "client_name": cfg.get("client_name", "Hermes Agent"), + "redirect_uris": [AnyUrl(_resolve_redirect_uri(cfg, port))], "grant_types": ["authorization_code", "refresh_token"], "response_types": ["code"], "token_endpoint_auth_method": auth_method, - # SEP-837 (2026-07-28 spec): clients MUST declare an application_type - # during registration so OIDC-strict authorization servers stop - # rejecting loopback redirect_uris. Hermes is a CLI/desktop app - # redirecting to 127.0.0.1/localhost — that is exactly "native". - # Overridable for the rare hosted-dashboard deployment 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 scope: - metadata_kwargs["scope"] = scope - + if cfg.get("scope"): + metadata_kwargs["scope"] = cfg["scope"] try: return OAuthClientMetadata.model_validate(metadata_kwargs) except Exception: - # mcp 1.x metadata models predate SEP-837 and reject the unknown - # field — retry without it rather than failing the whole flow. + # 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) @@ -1743,21 +1276,13 @@ def _invalidate_tokens_on_client_change( ) -> None: """Drop cached tokens when the configured OAuth client identity changes. - Tokens are minted for a specific ``client_id``: after the user edits - ``oauth.client_id`` / ``oauth.client_secret`` in config.yaml (or switches - from dynamic registration to a pre-registered client), the old tokens are - unusable — the token endpoint rejects their refresh with - ``invalid_client``. Pre-registered clients are deliberately exempt from - the ``invalid_client`` auto-poison path (config-supplied identity can't - be healed by re-registration), so without this check the stale tokens - wedge every request until the user manually wipes - ``~/.hermes/mcp-tokens/.*``. - - Compares the on-disk ``client.json`` identity against the incoming - config identity BEFORE the new client info overwrites it. Matching - identity is a no-op so live sessions and valid tokens are preserved. - Port of cline/cline#12983's "invalidate tokens when OAuth client - changes" invariant. + 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()) if not isinstance(existing, dict): @@ -1766,9 +1291,7 @@ def _invalidate_tokens_on_client_change( if not old_client_id: return old_client_secret = existing.get("client_secret") or None - if old_client_id == new_client_id and old_client_secret == ( - new_client_secret or None - ): + if old_client_id == new_client_id and old_client_secret == (new_client_secret or None): return removed = False for path in (storage._tokens_path(), storage._meta_path()): @@ -1802,26 +1325,17 @@ def _maybe_preregister_client( return if OAuthClientInformationFull is None: _ensure_sdk_loaded() - _invalidate_tokens_on_client_change( - storage, client_id, cfg.get("client_secret") - ) - port = cfg["_resolved_port"] - redirect_uri = _resolve_redirect_uri(cfg, port) - + _invalidate_tokens_on_client_change(storage, client_id, cfg.get("client_secret")) info_dict: dict[str, Any] = { "client_id": client_id, - "redirect_uris": [redirect_uri], + "redirect_uris": [_resolve_redirect_uri(cfg, cfg["_resolved_port"])], "grant_types": client_metadata.grant_types, "response_types": client_metadata.response_types, "token_endpoint_auth_method": client_metadata.token_endpoint_auth_method, } - if cfg.get("client_secret"): - info_dict["client_secret"] = cfg["client_secret"] - if cfg.get("client_name"): - info_dict["client_name"] = cfg["client_name"] - if cfg.get("scope"): - info_dict["scope"] = cfg["scope"] - + for key in ("client_secret", "client_name", "scope"): + if cfg.get(key): + info_dict[key] = cfg[key] client_info = OAuthClientInformationFull.model_validate(info_dict) _write_json(storage._client_info_path(), client_info.model_dump(mode="json", exclude_none=True)) logger.debug("Pre-registered client_id=%s for '%s'", client_id, storage._server_name) @@ -1833,14 +1347,11 @@ def humanize_oauth_registration_error( *, server_url: str | None = None, ) -> str | None: - """Turn a Dynamic Client Registration refusal into a useful next step. + """Turn a Dynamic Client Registration 403/Forbidden into a useful next step. - Returns a humanized message when the error is a registration 403/Forbidden, - else ``None`` so the caller keeps the original exception text. - - Figma's remote MCP gates DCR on exact ``client_name``. Hermes auto-sets - ``Claude Code`` (known-good); this message fires when the user overrode - that with something Figma still rejects, or an older Hermes is running. + 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. """ msg = str(exc) lowered = msg.lower() @@ -1884,18 +1395,10 @@ def build_oauth_auth( ) -> "OAuthClientProvider | None": """Build an ``httpx.Auth``-compatible OAuth handler for an MCP server. - Public API preserved for backwards compatibility. New code should use + 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. - - Args: - server_name: Server key in mcp_servers config (used for storage). - server_url: MCP server endpoint URL. - oauth_config: Optional dict from the ``oauth:`` block in config.yaml. - - Returns: - An ``OAuthClientProvider`` instance, or None if the MCP SDK lacks - OAuth support. + 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() @@ -1907,12 +1410,9 @@ def build_oauth_auth( ) return None - cfg = dict(oauth_config or {}) # copy — we mutate _resolved_port - apply_oauth_provider_defaults( - cfg, server_name=server_name, server_url=server_url - ) - storage = HermesTokenStorage(server_name) + 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 " @@ -1922,19 +1422,7 @@ def build_oauth_auth( "initial authorization, then cached tokens will be reused." ) - _configure_callback_port(cfg, storage) - client_metadata = _build_client_metadata(cfg) - _maybe_preregister_client(storage, cfg, client_metadata) - - # Use closure factories to avoid global state pollution (#44588, #34260). - resolved_port = cfg.get("_resolved_port", _oauth_port) - redirect_handler = _make_redirect_handler( - resolved_port, redirect_uri=cfg.get("redirect_uri") or None - ) - callback_handler = _make_callback_waiter( - resolved_port, cfg.get("_cimd_url"), timeout=float(cfg.get("timeout", 300)) - ) - + kwargs = build_provider_kwargs(cfg, storage, ssh_proxy_hint=True) provider_class = _get_hermes_oauth_provider_class() if provider_class is None: logger.warning( @@ -1942,16 +1430,4 @@ def build_oauth_auth( server_name, ) return None - - return provider_class( - server_url=server_url, - client_metadata=client_metadata, - storage=storage, - redirect_handler=redirect_handler, - # mcp 2.0 removed the provider's own `timeout` argument; the configured - # `oauth.timeout` is applied inside the callback waiter above, which is - # where the browser round-trip is actually awaited. - callback_handler=callback_handler, - token_user_agent=token_request_user_agent(cfg), - **cimd_provider_kwargs(cfg), - ) + return provider_class(server_url=server_url, **kwargs) diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index ee681a6335..80ec51a755 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -1,35 +1,20 @@ #!/usr/bin/env python3 """Central manager for per-server MCP OAuth state. -One instance shared across the process. Holds per-server OAuth provider -instances and coordinates: +One instance per process. Holds per-server provider instances and coordinates: -- **Cross-process token reload** via mtime-based disk watch. When an external - process (e.g. a user cron job) refreshes tokens on disk, the next auth flow - picks them up without requiring a process restart. -- **401 deduplication** via in-flight futures. When N concurrent tool calls - all hit 401 with the same access_token, only one recovery attempt fires; - the rest await the same result. -- **Reconnect signalling** for long-lived MCP sessions. The manager itself - does not drive reconnection — the `MCPServerTask` in `mcp_tool.py` does — - but the manager is the single source of truth that decides when reconnect - is warranted. +- **Cross-process token reload** via mtime-based disk watch, so tokens + refreshed by another process (cron, another CLI) are picked up without a + restart (Claude Code's ``invalidateOAuthCacheIfDiskChanged`` bug class). +- **401 deduplication** via in-flight futures: N concurrent tool calls hitting + 401 with the same access_token trigger one recovery attempt. +- **Reconnect signalling** — ``MCPServerTask`` in ``mcp_tool.py`` drives the + reconnect; the manager decides when it is warranted. -Replaces what used to be scattered across eight call sites in `mcp_oauth.py`, -`mcp_tool.py`, and `hermes_cli/mcp_config.py`. This module is the ONLY place -that instantiates the MCP SDK's `OAuthClientProvider` — all other code paths -go through `get_manager()`. - -Design reference: - -- Claude Code's ``invalidateOAuthCacheIfDiskChanged`` - (``claude-code/src/utils/auth.ts:1320``, CC-1096 / GH#24317). Identical - external-refresh staleness bug class. -- Codex's ``refresh_oauth_if_needed`` / ``persist_if_needed`` - (``codex-rs/rmcp-client/src/rmcp_client.rs:805``). We lean on the MCP SDK's - lazy refresh rather than calling refresh before every op, because one - ``stat()`` per tool call is cheaper than an ``await`` + potential refresh - round-trip, and the SDK's in-memory expiry path is already correct. +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 @@ -72,23 +57,10 @@ def _same_endpoint(a: str, b: str) -> bool: @dataclass class _ProviderEntry: - """Per-server OAuth state tracked by the manager. - - Fields: - server_url: The MCP server URL used to build the provider. Tracked - so we can discard a cached provider if the URL changes. - oauth_config: Optional dict from ``mcp_servers..oauth``. - provider: The ``httpx.Auth``-compatible provider wrapping the MCP - SDK. None until first use. - last_mtime_ns: Last-seen ``st_mtime_ns`` of the on-disk tokens file. - Zero if never read. Used by :meth:`MCPOAuthManager.invalidate_if_disk_changed` - to detect external refreshes. - lock: Serialises concurrent access to this entry's state. Bound to - whichever asyncio loop first awaits it (the MCP event loop). - pending_401: In-flight 401-handler futures keyed by the failed - access_token, for deduplicating thundering-herd 401s. Mirrors - Claude Code's ``pending401Handlers`` map. - """ + """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] @@ -104,191 +76,79 @@ class _ProviderEntry: def _make_hermes_provider_class() -> Optional[type]: - """Lazy-import the SDK base class and return our subclass. - - Wrapped in a function so this module imports cleanly even when the - MCP SDK's OAuth module is unavailable (e.g. older mcp versions). - """ + """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 + from tools.mcp_oauth_provider import HermesProviderMixin - class HermesMCPOAuthProvider(OAuthClientProvider): + class HermesMCPOAuthProvider(HermesProviderMixin, OAuthClientProvider): """OAuthClientProvider with pre-flow disk-mtime reload. - Before every ``async_auth_flow`` invocation, asks the manager to - check whether the tokens file on disk has been modified externally. - If so, the manager resets ``_initialized`` so the next flow - re-reads from storage. - - This makes external-process refreshes (cron, another CLI instance) - visible to the running MCP session without requiring a restart. - - Reference: Claude Code's ``invalidateOAuthCacheIfDiskChanged`` - (``src/utils/auth.ts:1320``, CC-1096 / GH#24317). + 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``. """ + _hermes_logger = logger + def __init__( self, *args: Any, server_name: str = "", preregistered: bool = False, - token_user_agent: "str | None" = None, **kwargs: Any, ): super().__init__(*args, **kwargs) - # mcp 2.0.0 uses a task-owned anyio.Lock and holds it across the - # yielded resource request. A session-long GET therefore blocks - # every concurrent POST, and HTTPX may later close the auth-flow - # generator from a different task than the lock owner. A binary - # semaphore preserves mutual exclusion without task ownership; - # async_auth_flow below 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 = "" - # When the client_id comes from config.yaml (pre-registered), an - # invalid_client rejection means the *config* is wrong — deleting - # client.json would just be re-seeded from config and re-running - # registration can't help. Only auto-heal dynamically-registered - # clients. See _maybe_flag_poisoned_client. + # 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 - # oauth.user_agent — stamped onto token-endpoint requests only; - # some authorization servers/WAFs reject httpx's default (#75576). - self._hermes_token_user_agent = token_user_agent - def _stamp_token_user_agent(self, request): - ua = getattr(self, "_hermes_token_user_agent", None) - if ua: - request.headers["User-Agent"] = ua - return request + def _hermes_storage(self): + """The context storage when it is a ``HermesTokenStorage``, else None.""" + from tools.mcp_oauth import HermesTokenStorage - def _coerce_client_secret_post(self) -> None: - """Use client_secret_post when dynamic registration returned a secret. - - Some MCP OAuth providers, notably Supabase, return a - ``client_secret`` from dynamic client registration but omit - ``token_endpoint_auth_method``. The MCP SDK treats the missing - value as public-client auth (``none``), so token exchange omits the - secret and Supabase rejects it with ``Required parameter: - client_secret``. Coerce the in-memory client info before token and - refresh requests. - """ - info = getattr(self.context, "client_info", None) - if not info or not getattr(info, "client_secret", None): - return - method = getattr(info, "token_endpoint_auth_method", None) - if method not in (None, "none", ""): - return - from mcp.shared.auth import OAuthClientInformationFull - - data = info.model_dump(mode="json", exclude_none=True) - data["token_endpoint_auth_method"] = "client_secret_post" - self.context.client_info = OAuthClientInformationFull.model_validate(data) - - async def _exchange_token_authorization_code(self, *args: Any, **kwargs: Any): - self._coerce_client_secret_post() - request = await super()._exchange_token_authorization_code(*args, **kwargs) - return self._stamp_token_user_agent(request) - - async def _refresh_token(self): - self._coerce_client_secret_post() - request = await super()._refresh_token() - return self._stamp_token_user_agent(request) - - async def _handle_token_response(self, response): - """Accept any 2xx token response and avoid leaking token bodies in errors.""" - if 200 <= response.status_code < 300: - from mcp.client.auth.utils import handle_token_response_scopes - from mcp.client.auth.oauth2 import OAuthTokenError - from httpx import HTTPError - - try: - token_response = await handle_token_response_scopes(response) - except (HTTPError, OAuthTokenError): - raise OAuthTokenError("Invalid token response") from None - self.context.current_tokens = token_response - self.context.update_token_expiry(token_response) - await self.context.storage.set_tokens(token_response) - return - - from mcp.client.auth.oauth2 import OAuthTokenError - - raise OAuthTokenError(f"Token exchange failed ({response.status_code})") - - async def _handle_refresh_response(self, response) -> bool: - """Accept any 2xx refresh response and avoid logging token bodies.""" - if not (200 <= response.status_code < 300): - logger.warning("Token refresh failed: %s", response.status_code) - self.context.clear_tokens() - return False - - from mcp.shared.auth import OAuthToken - from httpx import HTTPError - from pydantic import ValidationError - - try: - content = await response.aread() - token_response = OAuthToken.model_validate_json(content) - self.context.current_tokens = token_response - self.context.update_token_expiry(token_response) - await self.context.storage.set_tokens(token_response) - return True - except (HTTPError, ValidationError): - logger.warning("Invalid refresh response: %s", response.status_code) - self.context.clear_tokens() - return False + storage = self.context.storage + return storage if isinstance(storage, HermesTokenStorage) else None async def _initialize(self) -> None: - """Load stored tokens + client info AND seed token_expiry_time. + """Load stored state, seed ``token_expiry_time``, restore/prefetch metadata. - Also eagerly fetches OAuth authorization-server metadata (PRM + - ASM) when we have stored tokens but no cached metadata, so the - SDK's ``_refresh_token`` can build the correct token_endpoint - URL on the preemptive-refresh path. Without this, the SDK - falls back to ``{mcp_server_url}/token`` (wrong for providers - whose AS is a different origin — BetterStack's MCP lives at - ``https://mcp.betterstack.com`` but its token endpoint is at - ``https://betterstack.com/oauth/token``), the refresh 404s, and - we drop through to full browser reauth. + 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 before the first + request; ``HermesTokenStorage`` persists absolute ``expires_at`` so + the TTL reflects wall-clock age. - The SDK's base ``_initialize`` populates ``current_tokens`` but - does NOT call ``update_token_expiry``, so ``token_expiry_time`` - stays ``None`` and ``is_token_valid()`` returns True for any - loaded token regardless of actual age. After a process restart - this ships stale Bearer tokens to the server; some providers - return HTTP 401 (caught by the 401 handler), others return 200 - with an app-level auth error (invisible to the transport layer, - e.g. BetterStack returning "No teams found. Please check your - authentication."). - - Seeding ``token_expiry_time`` from the reloaded token fixes that: - ``is_token_valid()`` correctly reports False for expired tokens, - ``async_auth_flow`` takes the ``can_refresh_token()`` branch, - and the SDK quietly refreshes before the first real request. - - Paired with :class:`HermesTokenStorage` persisting an absolute - ``expires_at`` timestamp (``mcp_oauth.py:set_tokens``) so the - remaining TTL we compute here reflects real 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), the refresh 404s and we fall through to browser reauth. """ await super()._initialize() tokens = self.context.current_tokens if tokens is not None and tokens.expires_in is not None: self.context.update_token_expiry(tokens) - # Cold-load: restore OAuth server metadata from disk before any - # refresh attempt. Without this, a restarted process with cached - # tokens but no in-memory metadata would fall back to the SDK's - # guessed ``{server_url}/token`` path (returns 404 on most real - # providers) and require a full browser re-authorization. - storage = self.context.storage - from tools.mcp_oauth import HermesTokenStorage - if ( - isinstance(storage, HermesTokenStorage) - and self.context.oauth_metadata is None - ): + storage = self._hermes_storage() + if storage is not None and self.context.oauth_metadata is None: meta = storage.load_oauth_metadata() if meta is not None: self.context.oauth_metadata = meta @@ -299,20 +159,11 @@ def _make_hermes_provider_class() -> Optional[type]: meta.token_endpoint, ) - # Pre-flight OAuth AS discovery so ``_refresh_token`` has a - # correct ``token_endpoint`` before the first refresh attempt. - # Only runs when we have tokens on cold-load but no cached - # metadata — i.e. the exact scenario where the SDK's built-in - # 401-branch discovery hasn't had a chance to run yet. - if ( - tokens is not None - and self.context.oauth_metadata is None - ): + 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: if discovery fails, the SDK's normal 401- - # branch discovery will run on the next request. + # Non-fatal: the SDK's 401-branch discovery runs next request. logger.debug( "MCP OAuth '%s': pre-flight metadata discovery " "failed (non-fatal): %s", @@ -320,18 +171,15 @@ def _make_hermes_provider_class() -> Optional[type]: ) async def _prefetch_oauth_metadata(self) -> None: - """Fetch PRM + ASM from the well-known endpoints, cache on context. + """Fetch PRM + ASM from the well-known endpoints and cache on context. - Mirrors the SDK's 401-branch discovery (oauth2.py ~line 511-551) - but runs synchronously before the first request instead of - inside the httpx auth_flow generator. Uses the SDK's own URL - builders and response handlers so we track whatever the SDK - version we're pinned to expects. + 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. """ - # The SDK's httpx flavour, not Hermes' — mcp 2.0 builds on httpx2, - # and `create_oauth_metadata_request` below returns one of *its* - # Request objects, which only its own AsyncClient can send. See - # tools.mcp_tool.sdk_httpx. + # 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). from tools.mcp_tool import sdk_httpx httpx = sdk_httpx() if httpx is None: # pragma: no cover — SDK import would have failed @@ -345,53 +193,46 @@ def _make_hermes_provider_class() -> Optional[type]: ) server_url = self.context.server_url + + async def _send(client, url: str, label: str): + 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, + ) + return None + async with httpx.AsyncClient(timeout=10.0) as client: # Step 1: PRM discovery to learn the authorization_server URL. - for url in build_protected_resource_metadata_discovery_urls( - None, server_url - ): - req = create_oauth_metadata_request(url) - try: - resp = await client.send(req) - except httpx.HTTPError as exc: - logger.debug( - "MCP OAuth '%s': PRM discovery to %s failed: %s", - self._hermes_server_name, url, exc, - ) + for url in build_protected_resource_metadata_discovery_urls(None, server_url): + resp = await _send(client, url, "PRM") + if resp is None: continue prm = await handle_protected_resource_response(resp) if prm: self.context.protected_resource_metadata = prm if prm.authorization_servers: - self.context.auth_server_url = str( - prm.authorization_servers[0] - ) + self.context.auth_server_url = str(prm.authorization_servers[0]) break - # Step 2: ASM discovery against the auth_server_url (or - # 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 ): - req = create_oauth_metadata_request(url) - try: - resp = await client.send(req) - except httpx.HTTPError as exc: - logger.debug( - "MCP OAuth '%s': ASM discovery to %s failed: %s", - self._hermes_server_name, url, exc, - ) + resp = await _send(client, url, "ASM") + if resp is None: continue ok, asm = await handle_auth_metadata_response(resp) if not ok: break if asm: self.context.oauth_metadata = asm - # Persist immediately so a subsequent cold-load can - # skip discovery entirely. - storage = self.context.storage - from tools.mcp_oauth import HermesTokenStorage - if isinstance(storage, HermesTokenStorage): + # Persist now so a later cold-load skips discovery. + storage = self._hermes_storage() + if storage is not None: storage.save_oauth_metadata(asm) logger.debug( "MCP OAuth '%s': pre-flight ASM discovered " @@ -401,61 +242,37 @@ def _make_hermes_provider_class() -> Optional[type]: break def _persist_oauth_metadata_if_changed(self) -> None: - """Persist discovered OAuth metadata for future process restarts. - - Called after the SDK's normal 401-branch auth flow completes so - metadata discovered via the lazy path (not pre-flight) is also - saved. No-op when nothing to persist or metadata hasn't changed. - """ + """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 - if meta is None: - return - storage = self.context.storage - from tools.mcp_oauth import HermesTokenStorage - if not isinstance(storage, HermesTokenStorage): + storage = self._hermes_storage() + if meta is None or storage is None: return existing = storage.load_oauth_metadata() - if ( - existing is None - or str(existing.token_endpoint) != str(meta.token_endpoint) - ): + if existing is None or str(existing.token_endpoint) != str(meta.token_endpoint): storage.save_oauth_metadata(meta) async def _maybe_flag_poisoned_client(self, response: Any) -> None: """Detect a dead client registration and force re-registration. - When the IdP rejects our ``client_id`` with ``invalid_client`` on - the token endpoint (token exchange or refresh), the cached client - registration is provably dead server-side. We delete ``client.json`` - (+ stale metadata) so the SDK's next ``async_auth_flow`` takes the - ``if not client_info`` branch and re-runs RFC 7591 dynamic client - registration. This addresses the recurring manual-reset ritual in - GH#36767 for the auto-detectable subset (token-endpoint rejection); - the browser-side "Redirect URI Mismatch" case has no HTTP signal - and is handled by ``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 all hold: - * status is 400/401, - * the request hit the discovered ``token_endpoint`` (the only - request carrying our ``client_id``), and - * the body carries the ``invalid_client`` error code - (word-boundary match, so RFC 7591's ``invalid_client_metadata`` - registration error does not trip it). - Pre-registered (config-supplied) clients are never poisoned. - Fully best-effort: any failure here is swallowed so a detection - miss never breaks the live auth flow. - - Covers both the authorization-code token exchange and the - preemptive refresh — but only when ``token_endpoint`` was - discovered (``_initialize`` prefetches it on cold-load). If that - discovery was skipped, the guard returns early and the user falls - back 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`` + as a whole word (so RFC 7591's ``invalid_client_metadata`` does not + trip it). 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: if self._hermes_preregistered: return - status = getattr(response, "status_code", None) - if status not in (400, 401): + if getattr(response, "status_code", None) not in (400, 401): return meta = getattr(self.context, "oauth_metadata", None) token_endpoint = ( @@ -470,22 +287,15 @@ def _make_hermes_provider_class() -> Optional[type]: if not _same_endpoint(req_url, token_endpoint): return body = await response.aread() - # Word-boundary match: matches `"error":"invalid_client"` but - # not the RFC 7591 registration error `invalid_client_metadata` - # (the trailing `_metadata` removes the right-hand boundary). if not re.search(rb"\binvalid_client\b", body.lower()): return - storage = self.context.storage - from tools.mcp_oauth import HermesTokenStorage - - # When the rejected client_id was our Client ID Metadata - # Document URL, re-presenting it next flow would loop: the - # server has already fetched that document and refused it. - # Dropping the URL sends the retry down the DCR branch - # instead, and the marker on disk keeps the next process from - # walking back into the same refusal. `hermes mcp login` - # clears the marker, so a fixed document gets another chance. + 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). 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: @@ -496,10 +306,10 @@ def _make_hermes_provider_class() -> Optional[type]: self._hermes_server_name, cimd_url, ) self.context.client_metadata_url = None - if isinstance(storage, HermesTokenStorage): + if storage is not None: storage.mark_cimd_rejected() - if isinstance(storage, HermesTokenStorage): + if storage is not None: storage.poison_client_registration() # Drop the in-memory client so the SDK re-registers next flow. self.context.client_info = None @@ -511,9 +321,7 @@ def _make_hermes_provider_class() -> Optional[type]: ) async def async_auth_flow(self, request): # type: ignore[override] - # Pre-flow hook: ask the manager to refresh from disk if needed. - # Any failure here is non-fatal — we just log and proceed with - # whatever state the SDK already has. + # 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, @@ -525,20 +333,11 @@ def _make_hermes_provider_class() -> Optional[type]: self._hermes_server_name, exc, ) - # Manually bridge the bidirectional generator protocol. httpx's - # auth_flow driver (httpx._client._send_handling_auth) calls - # ``auth_flow.asend(response)`` to feed HTTP responses back into - # the generator. A naive wrapper using ``async for item in inner: - # yield item`` DISCARDS those .asend(response) values and resumes - # the inner generator with None, so the SDK's - # ``response = yield request`` branch in - # mcp/client/auth/oauth2.py sees response=None and crashes at - # ``if response.status_code == 401`` with AttributeError. - # - # The bridge below forwards each .asend() value into the inner - # generator via inner.asend(incoming), preserving the bidirectional - # contract. Regression from PR #11383 caught by - # 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`` (tests/tools/test_mcp_oauth_bidirectional.py). inner = super().async_auth_flow(request) resource_lock_released = False sent_access_token = None @@ -546,10 +345,9 @@ def _make_hermes_provider_class() -> Optional[type]: try: outgoing = await inner.__anext__() while True: - # The SDK holds context.lock for its entire generator, - # including while HTTPX waits on the actual MCP request. - # Release it only for that request. OAuth discovery, - # refresh, registration, and token exchange remain + # 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 @@ -562,9 +360,8 @@ def _make_hermes_provider_class() -> Optional[type]: if resource_lock_released: await self.context.lock.acquire() resource_lock_released = False - # A different request may have completed refresh or full - # authorization while this resource request was in - # flight. Retry with that token instead of starting a + # 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 ( @@ -577,22 +374,18 @@ def _make_hermes_provider_class() -> Optional[type]: await inner.aclose() retry_after_concurrent_auth = True break - # Sniff the response for a dead-client-registration signal - # before handing it back to the SDK (best-effort, GH#36767). + # Sniff for a dead-client-registration signal (best-effort). await self._maybe_flag_poisoned_client(incoming) outgoing = await inner.asend(incoming) except StopAsyncIteration: - # Persist any metadata the SDK discovered lazily during the - # 401 branch so a subsequent cold-load skips discovery. + # Persist metadata discovered lazily in the 401 branch. self._persist_oauth_metadata_if_changed() return finally: if resource_lock_released: # Balance the SDK's surrounding ``async with`` even when - # HTTPX cancels or closes the flow while the resource - # request is still in flight. Shield only this local - # bookkeeping; general inner-generator teardown remains - # the separate concern tracked by the cleanup PR. + # HTTPX cancels/closes the flow mid-request. Shield only + # this local bookkeeping. import anyio with anyio.CancelScope(shield=True): @@ -626,9 +419,8 @@ class MCPOAuthManager: def __init__(self) -> None: self._entries: dict[tuple[str, str], _ProviderEntry] = {} self._entries_lock = threading.Lock() - # Holds strong references to in-flight 401 handler tasks so the - # event loop's weak-reference bookkeeping cannot GC them mid-run - # and leave `await pending` waiters 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 ------------------------------------- @@ -686,15 +478,8 @@ class MCPOAuthManager: server_name: str, entry: _ProviderEntry, ) -> Optional[Any]: - """Build the underlying OAuth provider. - - Constructs :class:`HermesMCPOAuthProvider` directly using the helpers - extracted from ``tools.mcp_oauth``. The subclass injects a pre-flow - disk-watch hook so external token refreshes (cron, other CLI - instances) are visible to running MCP sessions. - - Returns None if the MCP 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, @@ -702,30 +487,13 @@ class MCPOAuthManager: return None # Local imports avoid circular deps at module import time. - from tools.mcp_oauth import ( - HermesTokenStorage, - OAuthNonInteractiveError, - _OAUTH_AVAILABLE, - _build_client_metadata, - _configure_callback_port, - _is_interactive, - _maybe_preregister_client, - _make_callback_waiter, - _make_redirect_handler, - cimd_provider_kwargs, - token_request_user_agent, - ) + 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 = dict(entry.oauth_config or {}) - from tools.mcp_oauth import apply_oauth_provider_defaults - - apply_oauth_provider_defaults( - cfg, server_name=server_name, server_url=entry.server_url - ) - storage = HermesTokenStorage(server_name) + cfg, storage = prepare_oauth_config(server_name, entry.server_url, entry.oauth_config) from tools.mcp_dashboard_oauth import get_dashboard_oauth_flow @@ -742,29 +510,11 @@ class MCPOAuthManager: "authorization." ) - _configure_callback_port(cfg, storage) - client_metadata = _build_client_metadata(cfg) - _maybe_preregister_client(storage, cfg, client_metadata) - - resolved_port = cfg.get("_resolved_port", 0) - redirect_handler = _make_redirect_handler(resolved_port) - # mcp 2.0 removed OAuthClientProvider's `timeout` argument, so the - # configured `oauth.timeout` now bounds the callback waiter's own poll - # loop instead — that is where the browser round-trip is awaited. - callback_handler = _make_callback_waiter( - resolved_port, cfg.get("_cimd_url"), timeout=float(cfg.get("timeout", 300)) - ) - return _HERMES_PROVIDER_CLS( server_name=server_name, preregistered=bool(cfg.get("client_id")), server_url=entry.server_url, - client_metadata=client_metadata, - storage=storage, - redirect_handler=redirect_handler, - callback_handler=callback_handler, - token_user_agent=token_request_user_agent(cfg), - **cimd_provider_kwargs(cfg), + **build_provider_kwargs(cfg, storage, ssh_proxy_hint=False), ) def remove( @@ -778,9 +528,7 @@ class MCPOAuthManager: Called by ``hermes mcp remove `` and (indirectly) by ``hermes mcp login `` during forced re-auth. """ - with self._entries_lock: - entry = self._entries.pop(self._key(server_name, hermes_home), None) - + 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) logger.info( @@ -807,10 +555,10 @@ class MCPOAuthManager: server_name: str, *, hermes_home: str | Path | None = None, - ) -> None: + ) -> _ProviderEntry | None: """Drop only the in-process provider, preserving persisted OAuth state.""" with self._entries_lock: - self._entries.pop(self._key(server_name, hermes_home), None) + return self._entries.pop(self._key(server_name, hermes_home), None) # -- Disk watch ---------------------------------------------------------- @@ -820,13 +568,10 @@ class MCPOAuthManager: *, hermes_home: str | Path | None = None, ) -> bool: - """If the tokens file on disk has a newer mtime than last-seen, force - the MCP SDK provider to reload its in-memory state. + """Force the SDK provider to reload when the tokens file mtime changed. - Returns True if the cache was invalidated (mtime differed). This is - the core fix for the external-refresh workflow: a cron job writes - fresh tokens to disk, and on the next tool call the running MCP - session picks them up without a restart. + Returns 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 @@ -844,9 +589,8 @@ class MCPOAuthManager: if mtime_ns != entry.last_mtime_ns: old = entry.last_mtime_ns entry.last_mtime_ns = mtime_ns - # Force the SDK's OAuthClientProvider to reload from storage - # on its next auth flow. `_initialized` is private API but - # stable across the MCP SDK versions we pin (>=1.26.0). + # `_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( @@ -901,19 +645,16 @@ class MCPOAuthManager: pending.set_result(True) return - # Step 2: No disk change — if the SDK can refresh - # in-place, let the caller retry. The SDK's httpx.Auth - # flow will issue the refresh on the next request. - provider = entry.provider - ctx = getattr(provider, "context", None) - can_refresh = False - if ctx is not None: - can_refresh_fn = getattr(ctx, "can_refresh_token", None) - if callable(can_refresh_fn): - try: - can_refresh = bool(can_refresh_fn()) - except Exception: - can_refresh = False + # 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 + except Exception: + can_refresh = False if not pending.done(): pending.set_result(can_refresh) except Exception as exc: # pragma: no cover — defensive diff --git a/tools/mcp_oauth_provider.py b/tools/mcp_oauth_provider.py new file mode 100644 index 0000000000..144dfc0fb6 --- /dev/null +++ b/tools/mcp_oauth_provider.py @@ -0,0 +1,155 @@ +"""Shared ``OAuthClientProvider`` customizations for Hermes MCP OAuth. + +Two code paths build an SDK provider — ``tools.mcp_oauth.build_oauth_auth`` +(legacy public API) and ``tools.mcp_oauth_manager.MCPOAuthManager`` — and both +need the same real-world fixes and the same config → constructor-kwargs +plumbing. This module holds that shared core once; the origin modules keep +their own subclass (logger name, disk-watch hooks) on top of it. +""" + +from __future__ import annotations + +import logging +from typing import TYPE_CHECKING, Any + +if TYPE_CHECKING: + from tools.mcp_oauth import HermesTokenStorage + +logger = logging.getLogger(__name__) + + +class HermesProviderMixin: + """Token-endpoint fixes layered over the SDK's ``OAuthClientProvider``. + + - Supabase-style dynamic registration returns a ``client_secret`` but omits + ``token_endpoint_auth_method``; the SDK then treats the client as public, + omits the secret, and the token endpoint rejects the exchange (looping the + browser authorization page). Coerce the in-memory client info to + ``client_secret_post`` right before token and refresh requests. + - ``token_user_agent`` (``oauth.user_agent``) is stamped onto token-endpoint + requests only — some authorization servers and WAFs reject httpx's default. + - Any 2xx token/refresh response is accepted, and token bodies never leak + into exception text or log output. + + Must precede the SDK class in the MRO. Subclasses set ``_hermes_logger`` so + warnings keep their origin module's logger name. + """ + + _hermes_logger: logging.Logger = logger + + def __init__(self, *args: Any, token_user_agent: str | None = None, **kwargs: Any): + super().__init__(*args, **kwargs) + self._hermes_token_user_agent = token_user_agent + + def _stamp_token_user_agent(self, request): + ua = getattr(self, "_hermes_token_user_agent", None) # tests build via __new__ + if ua: + request.headers["User-Agent"] = ua + return request + + def _coerce_client_secret_post(self) -> None: + info = self.context.client_info + if not info or not getattr(info, "client_secret", None): + return + if getattr(info, "token_endpoint_auth_method", None) not in (None, "none", ""): + return + from mcp.shared.auth import OAuthClientInformationFull + + data = info.model_dump(mode="json", exclude_none=True) + data["token_endpoint_auth_method"] = "client_secret_post" + self.context.client_info = OAuthClientInformationFull.model_validate(data) + + async def _exchange_token_authorization_code(self, *args: Any, **kwargs: Any): + self._coerce_client_secret_post() + request = await super()._exchange_token_authorization_code(*args, **kwargs) + return self._stamp_token_user_agent(request) + + async def _refresh_token(self): + self._coerce_client_secret_post() + request = await super()._refresh_token() + return self._stamp_token_user_agent(request) + + async def _store_tokens(self, token_response) -> None: + self.context.current_tokens = token_response + self.context.update_token_expiry(token_response) + await self.context.storage.set_tokens(token_response) + + async def _handle_token_response(self, response): + """Accept any 2xx token response; never echo the body into errors.""" + from mcp.client.auth.oauth2 import OAuthTokenError + + if not (200 <= response.status_code < 300): + raise OAuthTokenError(f"Token exchange failed ({response.status_code})") + from httpx import HTTPError + from mcp.client.auth.utils import handle_token_response_scopes + + try: + token_response = await handle_token_response_scopes(response) + except (HTTPError, OAuthTokenError): + raise OAuthTokenError("Invalid token response") from None + await self._store_tokens(token_response) + + async def _handle_refresh_response(self, response) -> bool: + """Accept any 2xx refresh response; never log the body.""" + if not (200 <= response.status_code < 300): + self._hermes_logger.warning("Token refresh failed: %s", response.status_code) + self.context.clear_tokens() + return False + from httpx import HTTPError + from mcp.shared.auth import OAuthToken + from pydantic import ValidationError + + try: + token_response = OAuthToken.model_validate_json(await response.aread()) + except (HTTPError, ValidationError): + self._hermes_logger.warning("Invalid refresh response: %s", response.status_code) + self.context.clear_tokens() + return False + await self._store_tokens(token_response) + return True + + +def prepare_oauth_config( + server_name: str, server_url: str, oauth_config: dict | None +) -> tuple[dict, "HermesTokenStorage"]: + """Copy the ``oauth:`` block, apply provider defaults, open its token storage. + + The copy matters: later steps record ``_resolved_port`` / ``_cimd_url`` in + the dict, which must never leak back into the caller's config. + """ + from tools import mcp_oauth as mo + + cfg = dict(oauth_config or {}) + mo.apply_oauth_provider_defaults(cfg, server_name=server_name, server_url=server_url) + return cfg, mo.HermesTokenStorage(server_name) + + +def build_provider_kwargs( + cfg: dict, storage: "HermesTokenStorage", *, ssh_proxy_hint: bool +) -> dict[str, Any]: + """Resolve the callback port and return the shared provider constructor kwargs. + + Runs the port → client-metadata → pre-registration sequence (order matters: + metadata needs the resolved port, pre-registration needs the metadata). + ``ssh_proxy_hint`` lets the redirect handler tailor its remote-session hint + to a configured proxy ``redirect_uri``. Helpers are looked up on + ``tools.mcp_oauth`` at call time so tests can patch them there. + """ + from tools import mcp_oauth as mo + + port = mo._configure_callback_port(cfg, storage) + client_metadata = mo._build_client_metadata(cfg) + mo._maybe_preregister_client(storage, cfg, client_metadata) + redirect_uri = (cfg.get("redirect_uri") or None) if ssh_proxy_hint else None + return { + "client_metadata": client_metadata, + "storage": storage, + "redirect_handler": mo._make_redirect_handler(port, redirect_uri=redirect_uri), + # mcp 2.0 dropped OAuthClientProvider's own `timeout`; the configured + # `oauth.timeout` bounds the callback waiter's poll loop instead. + "callback_handler": mo._make_callback_waiter( + port, cfg.get("_cimd_url"), timeout=float(cfg.get("timeout", 300)) + ), + "token_user_agent": mo.token_request_user_agent(cfg), + **mo.cimd_provider_kwargs(cfg), + } diff --git a/tools/schema_sanitizer.py b/tools/schema_sanitizer.py index 924fd3d019..c6a0f30665 100644 --- a/tools/schema_sanitizer.py +++ b/tools/schema_sanitizer.py @@ -1,37 +1,21 @@ """Sanitize tool JSON schemas for broad LLM-backend compatibility. -Some local inference backends (notably llama.cpp's ``json-schema-to-grammar`` -converter used to build GBNF tool-call parsers) are strict about what JSON -Schema shapes they accept. Schemas that OpenAI / Anthropic / most cloud -providers silently accept can make llama.cpp fail the entire request with: +Some backends are strict about JSON Schema shapes that OpenAI/Anthropic/most +cloud providers silently accept — llama.cpp's ``json-schema-to-grammar`` fails +the whole request (``Unrecognized schema: "object"``), Anthropic rejects +nullable ``anyOf`` at the top of ``input_schema``, Fireworks rejects ``default`` +beside ``$ref``, OpenAI's Codex backend rejects top-level combinators. Known +hostile constructs: - HTTP 400: Unable to generate parser for this template. - Automatic parser generation failed: JSON schema conversion failed: - Unrecognized schema: "object" +* ``{"type": "object"}`` with no ``properties``. +* A bare string (``"object"``) where a schema dict belongs (malformed MCP output). +* ``"type": ["string", "null"]`` array types. +* ``anyOf``/``oneOf`` unions whose only purpose is to permit ``null``. +* ``default`` (etc.) alongside ``$ref`` — e.g. ``{"$ref": "#/$defs/Foo", "default": null}``. -The failure modes we've seen in the wild: - -* ``{"type": "object"}`` with no ``properties`` — rejected as a node the - grammar generator can't constrain. -* A schema value that is the bare string ``"object"`` instead of a dict - (malformed MCP server output, e.g. ``additionalProperties: "object"``). -* ``"type": ["string", "null"]`` array types — many converters only accept - single-string ``type``. -* ``anyOf`` / ``oneOf`` unions whose only purpose is to permit ``null`` for - optional fields (common Pydantic/MCP shape). Anthropic rejects these at - the top of ``input_schema``; collapse them to the non-null branch. -* Unconstrained ``additionalProperties`` on objects with empty properties. -* ``default`` (and other annotation keywords) alongside ``$ref`` — strict - backends (Fireworks-hosted Kimi, JSON Schema draft-07 validators) reject - sibling keywords at the same level as ``$ref``. Common MCP/Pydantic shape - after nullable-union collapse:: - - {"$ref": "#/$defs/Foo", "default": null} - -This module walks the final tool schema tree (after MCP-level normalization -and any per-tool dynamic rebuilds) and fixes the known-hostile constructs -in-place on a deep copy. It is intentionally conservative: it only modifies -shapes the LLM backend couldn't use anyway. +This module walks the final tool schema tree (after MCP normalization and any +per-tool dynamic rebuilds) and fixes those in place on a deep copy. It is +deliberately conservative: it only modifies shapes the backend couldn't use. """ from __future__ import annotations @@ -39,33 +23,37 @@ from __future__ import annotations import copy import logging import re -from typing import Any +from typing import Any, Callable logger = logging.getLogger(__name__) # Anthropic (and Bedrock/Vertex/Azure fronting it) reject tool input schemas -# whose property keys don't match this pattern. Cloudflare's flat API MCP -# ships 61 such keys (query-filter params like ``issue_class~neq`` and -# ``meta.[]``) — one bad key anywhere in the tools array -# 400s the entire request. +# whose property keys don't match this pattern; one bad key anywhere in the +# tools array 400s the entire request (Cloudflare's MCP ships 61 such keys). _PROP_KEY_RE = re.compile(r"^[a-zA-Z0-9_.-]{1,64}$") _PROP_KEY_BAD_CHARS = re.compile(r"[^a-zA-Z0-9_.-]") +_UNION_KEYS = ("anyOf", "oneOf") +# Outer-node metadata carried onto a union's replacement node. +_UNION_META_KEYS = ("title", "description", "default", "examples") + + +def _empty_object() -> dict: + return {"type": "object", "properties": {}} + def sanitize_property_key(key: str) -> str: """Deterministically map an arbitrary property key to a conforming one.""" - new = _PROP_KEY_BAD_CHARS.sub("_", key)[:64] - return new or "param" + return _PROP_KEY_BAD_CHARS.sub("_", key)[:64] or "param" def _rename_property_keys(props: dict, path: str) -> dict[str, str]: """Return {original_key: conforming_key} for one properties dict. - Identity entries are omitted. Deterministic: keys are processed in - insertion order and collisions deduped with numeric suffixes, so the - model-visible schema AND the dispatch-time reverse map (computed - independently from the registry's original schema) always agree. + Identity entries are omitted. Deterministic (insertion order, numeric + suffixes on collision) so the model-visible schema and the dispatch-time + reverse map computed from the registry's original schema always agree. """ renames: dict[str, str] = {} taken = {k for k in props if _PROP_KEY_RE.match(k)} @@ -90,9 +78,8 @@ def _rename_property_keys(props: dict, path: str) -> dict[str, str]: def unrename_tool_args(params_schema: Any, args: Any) -> Any: """Map sanitized property keys in model-emitted args back to wire names. - ``params_schema`` is the ORIGINAL (unsanitized) parameters schema from the - registry. Recurses into object-typed values and array items so nested - renamed keys are restored too. Unknown keys pass through untouched. + ``params_schema`` is the ORIGINAL (unsanitized) registry schema. Recurses + into object values and array items; unknown keys pass through untouched. """ if not isinstance(params_schema, dict) or not isinstance(args, dict): return args @@ -118,21 +105,11 @@ def unrename_tool_args(params_schema: Any, args: Any) -> Any: def sanitize_tool_schemas(tools: list[dict]) -> list[dict]: - """Return a copy of ``tools`` with each tool's parameter schema sanitized. - - Input is an OpenAI-format tool list: - ``[{"type": "function", "function": {"name": ..., "parameters": {...}}}]`` - - The returned list is a deep copy — callers can safely mutate it without - affecting the original registry entries. - """ + """Return a deep-copied ``tools`` list (OpenAI format) with each tool's + parameter schema sanitized; callers may mutate the result freely.""" if not tools: return tools - - sanitized: list[dict] = [] - for tool in tools: - sanitized.append(_sanitize_single_tool(tool)) - return sanitized + return [_sanitize_single_tool(tool) for tool in tools] def _sanitize_single_tool(tool: dict) -> dict: @@ -143,34 +120,27 @@ def _sanitize_single_tool(tool: dict) -> dict: return out params = fn.get("parameters") - # Missing / non-dict parameters → substitute the minimal valid shape. - if not isinstance(params, dict): - fn["parameters"] = {"type": "object", "properties": {}} + if not isinstance(params, dict): # missing / non-dict → minimal valid shape + fn["parameters"] = _empty_object() return out - fn["parameters"] = _sanitize_node(params, path=fn.get("name", "")) - # After recursion, guarantee the top-level is an object with properties. - top = fn["parameters"] + name = fn.get("name", "") + top = _sanitize_node(params, path=name) + # Guarantee the top level is an object with properties. if not isinstance(top, dict): - fn["parameters"] = {"type": "object", "properties": {}} + top = _empty_object() else: if top.get("type") != "object": top["type"] = "object" - if "properties" not in top or not isinstance(top.get("properties"), dict): + if not isinstance(top.get("properties"), dict): top["properties"] = {} - # Final pass: collapse nullable anyOf/oneOf unions that the recursive - # sanitizer above leaves intact (it only handles the array-form - # ``type: [X, "null"]``). Keep the ``nullable: true`` hint so runtime - # argument coercion (``model_tools._schema_allows_null``) can still - # map a model-emitted ``"null"`` string to Python ``None``. - fn["parameters"] = strip_nullable_unions(fn["parameters"], keep_nullable_hint=True) - # Strip top-level combinators that strict backends (OpenAI's Codex - # endpoint at chatgpt.com/backend-api/codex) reject outright. Nested - # combinators inside properties are preserved. - fn["parameters"] = _strip_top_level_combinators( - fn["parameters"], path=fn.get("name", "") - ) - fn["parameters"] = _strip_ref_siblings(fn["parameters"]) + # Collapse nullable unions the recursive pass leaves intact (it only + # handles the array-form ``type: [X, "null"]``); keep ``nullable: true`` so + # runtime coercion (``model_tools._schema_allows_null``) still maps a + # model-emitted ``"null"`` string to Python ``None``. + top = strip_nullable_unions(top, keep_nullable_hint=True) + top = _strip_top_level_combinators(top, path=name) + fn["parameters"] = _strip_ref_siblings(top) return out @@ -179,26 +149,16 @@ _REF_FORBIDDEN_SIBLINGS = frozenset({"default"}) def _strip_ref_siblings(node: Any) -> Any: - """Drop forbidden sibling keywords from nodes that carry ``$ref``. - - Fireworks (and other draft-07-strict backends) fail tool requests with:: - - JSON Schema not supported: keyword(s) ['default'] not allowed at - the same level as $ref. - - Nullable-union collapse and MCP ingestion can leave ``default`` on a - ``$ref`` node; strip it recursively. - """ + """Recursively drop forbidden sibling keywords from nodes carrying ``$ref`` + (Fireworks: ``keyword(s) ['default'] not allowed at the same level as $ref``).""" if isinstance(node, list): return [_strip_ref_siblings(item) for item in node] if not isinstance(node, dict): return node - out = {key: _strip_ref_siblings(value) for key, value in node.items()} if "$ref" in out: for key in _REF_FORBIDDEN_SIBLINGS: - if key in out: - out.pop(key, None) + out.pop(key, None) return out @@ -206,22 +166,12 @@ _TOP_LEVEL_FORBIDDEN_KEYS = ("allOf", "anyOf", "oneOf", "enum", "not") def _strip_top_level_combinators(params: dict, *, path: str = "") -> dict: - """Drop combinator keywords from the top-level of a function parameters schema. + """Drop combinator keywords from the TOP level of a parameters schema only. - OpenAI's Codex backend (``chatgpt.com/backend-api/codex``) is stricter - than the public Functions API and rejects requests with:: - - Invalid schema for function 'X': schema must have type 'object' and - not have 'oneOf'/'anyOf'/'allOf'/'enum'/'not' at the top level. - - These keywords are typically used for conditional required-fields hints - (``allOf: [{if: ..., then: {required: [...]}}]``). Removing them at the - top level discards the hint but does not change which argument *values* - are valid — the tool handler always re-validates required fields. - - Only the *top* level is stripped; combinators nested inside a property's - schema are preserved (the strict rule only applies to the outermost - parameters object). + OpenAI's Codex backend rejects ``oneOf/anyOf/allOf/enum/not`` at the top + level. They are usually conditional-required hints; dropping them does not + change which argument values are valid (handlers re-validate). Nested + combinators are preserved. """ if not isinstance(params, dict): return params @@ -237,36 +187,34 @@ def _strip_top_level_combinators(params: dict, *, path: str = "") -> dict: return out +def _is_null_branch(item: Any) -> bool: + return isinstance(item, dict) and item.get("type") == "null" + + +def _carry_union_meta(outer: dict, replacement: dict, *, skip_default_on_ref: bool) -> None: + """Copy outer-union metadata onto *replacement* where absent.""" + for meta_key in _UNION_META_KEYS: + if meta_key in outer and meta_key not in replacement: + # ``default`` is illegal alongside ``$ref`` on strict backends. + if skip_default_on_ref and meta_key == "default" and "$ref" in replacement: + continue + replacement[meta_key] = outer[meta_key] + + def strip_nullable_unions( schema: Any, *, keep_nullable_hint: bool = True, ) -> Any: - """Collapse ``anyOf`` / ``oneOf`` nullable unions to the non-null branch. + """Collapse ``anyOf``/``oneOf`` nullable unions to the single non-null branch. - MCP / Pydantic optional fields commonly arrive as:: - - {"anyOf": [{"type": "string"}, {"type": "null"}], "default": null} - - Anthropic's tool input-schema validator rejects the null branch. Tool - optionality is already represented by the parent object's ``required`` - array, so we collapse the union to the single non-null variant. - - Metadata (``title``, ``description``, ``default``, ``examples``) on the - outer union node is carried over to the replacement variant. - - Args: - schema: JSON-Schema fragment (dict, list, or scalar). - keep_nullable_hint: If True, set ``nullable: true`` on the replacement - to preserve the "this field may be None" signal for downstream - consumers that care (e.g. runtime argument coercion that maps the - literal string ``"null"`` to Python ``None``). Anthropic's - validator accepts ``nullable: true`` but strict producers may - prefer False. - - Returns: - The schema with nullable unions collapsed. Non-union nodes are - returned unchanged. + MCP/Pydantic optional fields arrive as + ``{"anyOf": [{"type": "string"}, {"type": "null"}], "default": null}``; + Anthropic rejects the null branch, and optionality is already expressed by + the parent's ``required``. Only collapses when a null branch was dropped + AND exactly one non-null branch survives. Outer metadata is carried over. + ``keep_nullable_hint`` sets ``nullable: true`` on the replacement for + downstream consumers (runtime ``"null"`` → ``None`` coercion). """ if isinstance(schema, list): return [strip_nullable_unions(item, keep_nullable_hint=keep_nullable_hint) for item in schema] @@ -277,27 +225,16 @@ def strip_nullable_unions( k: strip_nullable_unions(v, keep_nullable_hint=keep_nullable_hint) for k, v in schema.items() } - for key in ("anyOf", "oneOf"): + for key in _UNION_KEYS: variants = stripped.get(key) if not isinstance(variants, list): continue - non_null = [ - item for item in variants - if not (isinstance(item, dict) and item.get("type") == "null") - ] - # Only collapse when we actually dropped a null branch AND exactly - # one non-null branch survives (otherwise the union is meaningful - # and we leave it alone). + non_null = [item for item in variants if not _is_null_branch(item)] if len(non_null) == 1 and len(non_null) != len(variants): replacement = dict(non_null[0]) if isinstance(non_null[0], dict) else {} if keep_nullable_hint: replacement.setdefault("nullable", True) - for meta_key in ("title", "description", "default", "examples"): - if meta_key in stripped and meta_key not in replacement: - # ``default`` is illegal alongside ``$ref`` on strict backends. - if meta_key == "default" and "$ref" in replacement: - continue - replacement[meta_key] = stripped[meta_key] + _carry_union_meta(stripped, replacement, skip_default_on_ref=True) return strip_nullable_unions(replacement, keep_nullable_hint=keep_nullable_hint) return stripped @@ -311,59 +248,41 @@ _CONST_PRIMITIVE_TYPES: dict[type, str] = { def _const_branch_type(branch: Any) -> str | None: - """Return the JSON-Schema primitive type of a pure ``const`` branch. + """JSON-Schema primitive type of a pure ``const`` branch, else None. - A branch qualifies when it is a dict carrying ``const`` with a primitive - value, and any declared ``type`` matches the const value's type. Branch - metadata (``title``, ``description``) does not disqualify it, but any - other constraining keyword does. Returns ``None`` for non-qualifying - branches. + Qualifies when the dict carries a primitive ``const`` and any declared + ``type`` matches it; ``title``/``description`` are allowed, any other + constraining keyword disqualifies. """ if not isinstance(branch, dict) or "const" not in branch: return None - extra = set(branch) - {"const", "type", "title", "description"} - if extra: + if set(branch) - {"const", "type", "title", "description"}: return None value = branch["const"] - # bool is a subclass of int in Python; check it first so True/False never - # classify as integers. - for py_type, json_type in _CONST_PRIMITIVE_TYPES.items(): - if type(value) is py_type: - declared = branch.get("type") - if declared is not None and declared != json_type: - return None - return json_type - return None + # ``type(value) is`` (not isinstance): bool is a subclass of int. + json_type = _CONST_PRIMITIVE_TYPES.get(type(value)) + if json_type is None: + return None + declared = branch.get("type") + if declared is not None and declared != json_type: + return None + return json_type def collapse_const_unions(schema: Any) -> Any: - """Collapse ``anyOf`` / ``oneOf`` unions of same-typed consts to ``enum``. + """Collapse ``anyOf``/``oneOf`` unions of same-typed consts to ``enum``. - Ported from block/goose ``tool_schema_normalize.rs`` (Apache-2.0). + Ported from block/goose ``tool_schema_normalize.rs`` (Apache-2.0). MCP + servers generated from Rust/TS union types emit + ``{"anyOf": [{"const": "red"}, {"const": "green"}]}``; strict backends + mishandle these while ``{"type": "string", "enum": [...]}`` is universal. - MCP servers (particularly ones generated from Rust/TypeScript union types) - commonly emit closed value sets as const unions:: - - {"anyOf": [{"const": "red"}, {"const": "green"}, {"const": "blue"}]} - - Strict tool-calling backends reject or mishandle these, while the - equivalent property-level ``enum`` form is universally supported:: - - {"type": "string", "enum": ["red", "green", "blue"]} - - The collapse applies only when EVERY non-null branch is a pure ``const`` - of the same primitive type (bool/int/float/str — ``bool`` never merges - with ``integer``). Mixed unions and non-uniform const types pass through - untouched. A single ``{"type": "null"}`` branch is tolerated: it is - dropped and recorded as ``nullable: true`` (matching the - ``strip_nullable_unions`` convention), since strip_nullable_unions only - collapses unions with exactly one non-null branch and therefore leaves - null+multi-const unions for us. - - Outer-node metadata (``title``, ``description``, ``default``, - ``examples``) is carried onto the replacement. Enum order preserves - branch order, so output is deterministic and byte-stable across - discoveries. Input is never mutated. + Applies only when EVERY non-null branch is a pure ``const`` of one + primitive type (``bool`` never merges with ``integer``). One + ``{"type": "null"}`` branch is tolerated and recorded as ``nullable: true`` + (``strip_nullable_unions`` only handles single-non-null unions, so + null+multi-const unions land here). Enum order preserves branch order; + outer metadata is carried over; input is never mutated. """ if isinstance(schema, list): return [collapse_const_unions(item) for item in schema] @@ -371,13 +290,12 @@ def collapse_const_unions(schema: Any) -> Any: return schema out = {k: collapse_const_unions(v) for k, v in schema.items()} - for key in ("anyOf", "oneOf"): + for key in _UNION_KEYS: variants = out.get(key) if not isinstance(variants, list) or not variants: continue null_branches = [ - item for item in variants - if isinstance(item, dict) and item.get("type") == "null" and "const" not in item + item for item in variants if _is_null_branch(item) and "const" not in item ] const_branches = [item for item in variants if item not in null_branches] if len(null_branches) > 1 or not const_branches: @@ -391,46 +309,67 @@ def collapse_const_unions(schema: Any) -> Any: } if null_branches: replacement["nullable"] = True - for meta_key in ("title", "description", "default", "examples"): - if meta_key in out and meta_key not in replacement: - replacement[meta_key] = out[meta_key] + _carry_union_meta(out, replacement, skip_default_on_ref=False) return replacement return out +_BARE_TYPE_NAMES = frozenset({"object", "string", "number", "integer", "boolean", "array", "null"}) +# Sibling keywords whose values are NOT schemas: recursing would mistake literal +# strings like "path" for bare-string schemas. Passed through unchanged +# (``required`` remapped through property renames). +_NON_SCHEMA_LIST_KEYS = frozenset({"required", "enum", "examples", "dependentRequired"}) + + +def _normalize_type_array(value: list, out: dict) -> None: + """Normalize a ``type: [...]`` array into *out*. + + Several backends reject array types (llama.cpp's grammar generator; Gemini + via OpenAI-compatible transports 400s). Per the AI-SDK behavior: one + non-null type → ``type: X`` (+ ``nullable`` if ``null`` present); several → + ``anyOf`` of single-type schemas so EVERY branch survives; none → ``null`` + or the object fallback. Ported from anomalyco/opencode#31877. + """ + has_null = "null" in value + non_null = [t for t in value if isinstance(t, str) and t != "null"] + if len(non_null) == 1: + out["type"] = non_null[0] + elif len(non_null) >= 2: + out["anyOf"] = [{"type": t} for t in non_null] + else: + out["type"] = "null" if has_null else "object" + return + if has_null: + out.setdefault("nullable", True) + + def _sanitize_node(node: Any, path: str) -> Any: """Recursively sanitize a JSON-Schema fragment. - - Replaces bare-string schema values ("object", "string", ...) with - ``{"type": }`` so downstream consumers see a dict. - - Injects ``properties: {}`` into object-typed nodes missing it. - - Normalizes ``type: [X, "null"]`` arrays to single ``type: X`` (keeping - ``nullable: true`` as a hint), and multi-type arrays like - ``["number", "string"]`` to an ``anyOf`` of single-type schemas so no - branch is dropped (ported from anomalyco/opencode#31877). + - Bare-string schema values become ``{"type": }`` (unknown strings + become a permissive object schema rather than something backends reject). + - Object-typed nodes gain ``properties: {}`` (llama.cpp can't constrain a + free-form object). + - ``type`` arrays are normalized (see ``_normalize_type_array``). - Recurses into ``properties``, ``items``, ``additionalProperties``, - ``anyOf``, ``oneOf``, ``allOf``, and ``$defs`` / ``definitions``. + ``anyOf``/``oneOf``/``allOf`` and ``$defs``/``definitions``; property + keys are renamed to the provider-safe pattern and ``required`` follows. + - ``required`` entries that don't exist in ``properties`` are pruned + (malformed MCP schemas; built-in/plugin tools skip the MCP-level check). """ - # Malformed: the schema position holds a bare string like "object". if isinstance(node, str): - if node in {"object", "string", "number", "integer", "boolean", "array", "null"}: + if node in _BARE_TYPE_NAMES: logger.debug( "schema_sanitizer[%s]: replacing bare-string schema %r " "with {'type': %r}", path, node, node, ) - return {"type": node} if node != "object" else { - "type": "object", - "properties": {}, - } - # Any other stray string is not a schema — drop it by replacing with - # a permissive object schema rather than propagate something the - # backend will reject. + return _empty_object() if node == "object" else {"type": node} logger.debug( "schema_sanitizer[%s]: replacing non-schema string %r " "with empty object schema", path, node, ) - return {"type": "object", "properties": {}} + return _empty_object() if isinstance(node, list): return [_sanitize_node(item, f"{path}[{i}]") for i, item in enumerate(node)] @@ -438,227 +377,70 @@ def _sanitize_node(node: Any, path: str) -> Any: if not isinstance(node, dict): return node - # Compute property-key renames up front so the ``required`` branch below - # can remap regardless of dict iteration order (``required`` may precede - # ``properties`` in the source dict). + # Renames are computed up front so ``required`` can be remapped even when + # it precedes ``properties`` in the source dict. prop_renames: dict[str, str] = {} if isinstance(node.get("properties"), dict): prop_renames = _rename_property_keys(node["properties"], f"{path}.properties") out: dict = {} for key, value in node.items(): - # JSON Schema ``type`` arrays (e.g. ``["number", "string"]``, common - # in MCP tool schemas) are rejected by several tool-call backends: - # * llama.cpp's grammar generator only accepts a singular string type. - # * Gemini (including OpenAI-compatible transports such as GitHub - # Copilot proxying to Gemini) rejects the array form outright — - # plain @ai-sdk/google rewrites it, but the OpenAI-compatible path - # forwards it verbatim and the backend 400s. - # - # Normalize per the SDK's behavior: - # * single non-null type → ``type: X`` (+ ``nullable: true`` if the - # array also contained "null"). No data lost. - # * multiple non-null types → ``anyOf`` of single-type schemas, so - # EVERY branch survives instead of silently dropping all but the - # first. ``null`` is lifted into ``nullable: true``. - # * all-null / empty → ``type: "null"`` (or object fallback). - # Ported from anomalyco/opencode#31877. if key == "type" and isinstance(value, list): - has_null = "null" in value - non_null = [t for t in value if isinstance(t, str) and t != "null"] - if len(non_null) == 1: - out["type"] = non_null[0] - if has_null: - out.setdefault("nullable", True) - continue - if len(non_null) >= 2: - # Preserve all branches as a union instead of dropping them. - out["anyOf"] = [{"type": t} for t in non_null] - if has_null: - out.setdefault("nullable", True) - continue - # No usable non-null type: all-null array → type: "null"; - # otherwise an empty/garbage array → object fallback. - out["type"] = "null" if has_null else "object" - continue - - if key in {"properties", "$defs", "definitions"} and isinstance(value, dict): + _normalize_type_array(value, out) + elif key in {"properties", "$defs", "definitions"} and isinstance(value, dict): renames = prop_renames if key == "properties" else {} - new_props = {} - for sub_k, sub_v in value.items(): - out_k = renames.get(sub_k, sub_k) - new_props[out_k] = _sanitize_node(sub_v, f"{path}.{key}.{out_k}") - out[key] = new_props + out[key] = { + renames.get(sub_k, sub_k): _sanitize_node(sub_v, f"{path}.{key}.{renames.get(sub_k, sub_k)}") + for sub_k, sub_v in value.items() + } elif key in {"items", "additionalProperties"}: - if isinstance(value, bool): - # Keep bool ``additionalProperties`` as-is — it's a valid form - # and widely accepted. ``items: true/false`` is non-standard - # but we preserve rather than drop. - out[key] = value - else: - out[key] = _sanitize_node(value, f"{path}.{key}") + # Bool ``additionalProperties`` is valid and widely accepted; + # ``items: true/false`` is non-standard but preserved rather than dropped. + out[key] = value if isinstance(value, bool) else _sanitize_node(value, f"{path}.{key}") elif key in {"anyOf", "oneOf", "allOf"} and isinstance(value, list): - out[key] = [ - _sanitize_node(item, f"{path}.{key}[{i}]") - for i, item in enumerate(value) - ] - elif key in {"required", "enum", "examples", "dependentRequired"}: - # Schema "sibling" keywords whose values are NOT schemas: - # - ``required``: list of property-name strings - # - ``enum``: list of literal values (any JSON type) - # - ``examples``: list of example values (any JSON type) - # - ``dependentRequired``: mapping of property names to lists of - # required property-name strings (JSON Schema 2020-12) - # Recursing into these with _sanitize_node() would mis-interpret - # literal strings like "path" as bare-string schemas and replace - # them with {"type": "object"} dicts. Pass through unchanged - # (remapping ``required`` entries through the property renames). + out[key] = [_sanitize_node(item, f"{path}.{key}[{i}]") for i, item in enumerate(value)] + elif key in _NON_SCHEMA_LIST_KEYS: if key == "required" and prop_renames and isinstance(value, list): - out[key] = [prop_renames.get(r, r) if isinstance(r, str) else r - for r in value] + out[key] = [prop_renames.get(r, r) if isinstance(r, str) else r for r in value] else: out[key] = copy.deepcopy(value) if isinstance(value, (list, dict)) else value else: out[key] = _sanitize_node(value, f"{path}.{key}") if isinstance(value, (dict, list)) else value - # Object nodes without properties: inject empty properties dict. - # llama.cpp's grammar generator can't constrain a free-form object. - if out.get("type") == "object" and not isinstance(out.get("properties"), dict): - out["properties"] = {} - - # Prune ``required`` entries that don't exist in properties (defense - # against malformed MCP schemas; also caught upstream for MCP tools, but - # built-in tools or plugin tools may not have been through that path). - if out.get("type") == "object" and isinstance(out.get("required"), list): - props = out.get("properties") or {} - valid = [r for r in out["required"] if isinstance(r, str) and r in props] - if not valid: - out.pop("required", None) - elif len(valid) != len(out["required"]): - out["required"] = valid - + if out.get("type") == "object": + if not isinstance(out.get("properties"), dict): + out["properties"] = {} + if isinstance(out.get("required"), list): + props = out.get("properties") or {} + valid = [r for r in out["required"] if isinstance(r, str) and r in props] + if not valid: + out.pop("required", None) + elif len(valid) != len(out["required"]): + out["required"] = valid return out # ============================================================================= -# Reactive strip — only invoked when llama.cpp rejects a schema +# Reactive strips — only invoked after a backend rejects a schema # ============================================================================= _STRIP_ON_RECOVERY_KEYS = frozenset({"pattern", "format"}) -def strip_pattern_and_format(tools: list[dict]) -> tuple[list[dict], int]: - """Strip ``pattern`` and ``format`` JSON Schema keywords from tool schemas. - - This is a *reactive* sanitizer invoked only when llama.cpp's - ``json-schema-to-grammar`` converter has rejected a tool schema with an - HTTP 400 grammar-parse error. llama.cpp's regex engine supports only a - small subset of ECMAScript regex (literals, ``.``, ``[...]``, ``|``, - ``*``, ``+``, ``?``, ``{n,m}``) — it rejects escape classes like ``\\d``, - ``\\w``, ``\\s`` and most ``format`` values. Cloud providers (OpenAI, - Anthropic, OpenRouter, Gemini) accept these keywords fine and rely on - them as prompting hints, so we keep them in the default schema and only - strip on demand. - - The strip operates on a sibling of ``type`` (so schema keywords are - removed) — a property literally *named* ``pattern`` (e.g. the first arg - of the built-in ``search_files`` tool) is not affected because property - names live in the ``properties`` dict, not as siblings of ``type``. - - Args: - tools: OpenAI-format tool list, mutated in place for efficiency. - Callers that need to preserve the original should deep-copy first. - - Returns: - ``(tools, stripped_count)`` — the same list reference plus a count of - how many ``pattern``/``format`` keywords were removed across all tools. - """ +def _reactive_strip(tools: list[dict], strip_node: Callable[[dict], int], log_msg: str) -> tuple[list[dict], int]: + """Walk every tool's parameters in place, applying *strip_node* to each dict + node (it returns how many keywords it removed). Handles OpenAI format + (``{"function": {"parameters": ...}}``) and Responses format + (``{"name": ..., "parameters": ...}`` — codex_responses mode, xAI, etc.). + Returns ``(tools, stripped_count)`` — the same list reference.""" if not tools: return tools, 0 - stripped = 0 def _walk(node: Any) -> None: nonlocal stripped if isinstance(node, dict): - # Only strip as a sibling of ``type`` — i.e. when this node is - # itself a schema. This avoids stripping literal property keys - # named "pattern" (search_files.pattern, etc.) because those live - # inside a ``properties`` dict, not as siblings of ``type``. - is_schema_node = "type" in node or "anyOf" in node or "oneOf" in node or "allOf" in node - for key in list(node.keys()): - if is_schema_node and key in _STRIP_ON_RECOVERY_KEYS: - node.pop(key, None) - stripped += 1 - continue - _walk(node[key]) - elif isinstance(node, list): - for item in node: - _walk(item) - - for tool in tools: - if not isinstance(tool, dict): - continue - - # OpenAI-format: {"function": {"parameters": {...}}} - fn = tool.get("function") - if isinstance(fn, dict): - params = fn.get("parameters") - if isinstance(params, dict): - _walk(params) - continue - - # Responses-format: {"name": "...", "parameters": {...}} - # (used by codex_responses API mode — xAI, OpenAI Codex, etc.) - params = tool.get("parameters") - if isinstance(params, dict): - _walk(params) - continue - - if stripped: - logger.info( - "schema_sanitizer: stripped %d pattern/format keyword(s) from " - "tool schemas (llama.cpp grammar-parse recovery)", - stripped, - ) - return tools, stripped - - -def strip_slash_enum(tools: list[dict]) -> tuple[list[dict], int]: - """Strip ``enum`` keywords whose string values contain a forward slash. - - xAI's ``/v1/responses`` and ``/v1/chat/completions`` endpoints compile - tool schemas to a grammar that rejects ``enum`` values containing ``/`` - (the request fails with HTTP 400 "Invalid arguments passed to the - model" before any token is emitted). Most commonly hit by MCP-derived - tools whose enum lists HuggingFace model IDs (``Qwen/Qwen3.5-0.8B``, - ``openai/gpt-oss-20b``) or owner/name environment IDs. The constraint - is purely a prompting hint; dropping it lets the model still see the - field description and pick a value, without xAI tripping on the slash. - - Args: - tools: OpenAI-format or Responses-format tool list, mutated in - place. Callers that need to preserve the original should - deep-copy first. - - Returns: - ``(tools, stripped_count)`` — same list reference plus a count of - how many ``enum`` keywords were removed. - """ - if not tools: - return tools, 0 - - stripped = 0 - - def _walk(node: Any) -> None: - nonlocal stripped - if isinstance(node, dict): - enum_val = node.get("enum") - if isinstance(enum_val, list) and any( - isinstance(v, str) and "/" in v for v in enum_val - ): - node.pop("enum", None) - stripped += 1 + stripped += strip_node(node) for v in node.values(): _walk(v) elif isinstance(node, list): @@ -669,19 +451,61 @@ def strip_slash_enum(tools: list[dict]) -> tuple[list[dict], int]: if not isinstance(tool, dict): continue fn = tool.get("function") - if isinstance(fn, dict): - params = fn.get("parameters") - if isinstance(params, dict): - _walk(params) - continue - params = tool.get("parameters") - if isinstance(params, dict): - _walk(params) + if isinstance(fn, dict) and isinstance(fn.get("parameters"), dict): + _walk(fn["parameters"]) + continue + if isinstance(tool.get("parameters"), dict): + _walk(tool["parameters"]) if stripped: - logger.info( - "schema_sanitizer: stripped %d enum keyword(s) containing '/' " - "from tool schemas (xAI Responses grammar-compile recovery)", - stripped, - ) + logger.info(log_msg, stripped) return tools, stripped + + +def strip_pattern_and_format(tools: list[dict]) -> tuple[list[dict], int]: + """Strip ``pattern``/``format`` keywords from tool schemas, in place. + + Reactive: invoked only after llama.cpp's grammar converter rejected a + schema with HTTP 400. Its regex engine supports a small ECMAScript subset + (no ``\\d``/``\\w``/``\\s``) and most ``format`` values; cloud providers rely + on these as prompting hints, so they stay in the default schema. + + Only strips as a sibling of ``type``/combinators (i.e. on schema nodes), so + a property literally *named* ``pattern`` (``search_files``) is untouched — + property names live inside ``properties``, not beside ``type``. + """ + def _strip(node: dict) -> int: + if not ("type" in node or "anyOf" in node or "oneOf" in node or "allOf" in node): + return 0 + hits = [k for k in node if k in _STRIP_ON_RECOVERY_KEYS] + for k in hits: + node.pop(k, None) + return len(hits) + + return _reactive_strip( + tools, _strip, + "schema_sanitizer: stripped %d pattern/format keyword(s) from " + "tool schemas (llama.cpp grammar-parse recovery)", + ) + + +def strip_slash_enum(tools: list[dict]) -> tuple[list[dict], int]: + """Strip ``enum`` keywords whose string values contain ``/``, in place. + + xAI's ``/v1/responses`` and ``/v1/chat/completions`` compile schemas to a + grammar that rejects ``/`` in enum values (HTTP 400 before any token) — + typically MCP enums of HuggingFace model IDs or owner/name env IDs. The + constraint is a prompting hint only; the model still sees the description. + """ + def _strip(node: dict) -> int: + enum_val = node.get("enum") + if isinstance(enum_val, list) and any(isinstance(v, str) and "/" in v for v in enum_val): + node.pop("enum", None) + return 1 + return 0 + + return _reactive_strip( + tools, _strip, + "schema_sanitizer: stripped %d enum keyword(s) containing '/' " + "from tool schemas (xAI Responses grammar-compile recovery)", + )