diff --git a/tests/tools/test_mcp_oauth.py b/tests/tools/test_mcp_oauth.py index faa4ebc3a0..780a6752b8 100644 --- a/tests/tools/test_mcp_oauth.py +++ b/tests/tools/test_mcp_oauth.py @@ -3,7 +3,9 @@ import json import stat import sys +import time from io import BytesIO +from pathlib import Path from unittest.mock import patch, MagicMock from urllib.parse import quote @@ -21,6 +23,7 @@ from tools.mcp_oauth import ( _make_callback_handler, _make_redirect_handler, _paste_callback_reader, + _token_store_lock, ) @@ -1185,3 +1188,116 @@ def test_humanize_non_registration_403_passthrough(): ) is None ) + + +# --------------------------------------------------------------------------- +# Token-store advisory lock +# +# Two Hermes backends share one HERMES_HOME (desktop `serve` + `gateway run`). +# Providers that issue single-use refresh tokens punish an unserialized +# read-modify-write: both processes POST the same token and the loser's +# refresh is rejected. cron/jobs.py guards jobs.json the same way. +# --------------------------------------------------------------------------- + + +_LOCK_HOLDER_SRC = ''' +import os, sys, time +from pathlib import Path +sys.path.insert(0, sys.argv[1]) +os.environ["HERMES_HOME"] = sys.argv[2] +from tools.mcp_oauth import _token_store_lock +with _token_store_lock(Path(sys.argv[3])): + print("HELD", flush=True) + time.sleep(float(sys.argv[4])) +''' + + +def test_token_store_lock_is_reentrant_within_a_thread(tmp_path): + """Nested acquisition must not deadlock the same thread.""" + target = tmp_path / "srv.json" + target.write_text("{}", encoding="utf-8") + + with _token_store_lock(target): + with _token_store_lock(target): + pass # would hang if the in-process guard were not re-entrant + + +def test_token_store_lock_survives_unwritable_lock_path(tmp_path, monkeypatch): + """A lock that cannot be created must not block token access.""" + target = tmp_path / "srv.json" + target.write_text("{}", encoding="utf-8") + + def _boom(*a, **k): + raise OSError("read-only filesystem") + + monkeypatch.setattr("builtins.open", _boom) + + # Degrades to in-process locking rather than raising. + with _token_store_lock(target): + pass + + +def test_token_store_lock_releases_on_exception(tmp_path): + """An exception inside the section must still release the lock.""" + target = tmp_path / "srv.json" + target.write_text("{}", encoding="utf-8") + + with pytest.raises(ValueError): + with _token_store_lock(target): + raise ValueError("boom") + + # Reacquire immediately; a leaked lock would stall until the timeout. + start = time.monotonic() + with _token_store_lock(target): + pass + assert time.monotonic() - start < 1.0 + + +@pytest.mark.skipif( + not hasattr(sys, "executable") or not sys.executable, + reason="needs a real interpreter to spawn a peer process", +) +def test_token_store_lock_excludes_a_separate_process(tmp_path): + """The point of an ADVISORY FILE lock: exclude another process. + + An in-process test cannot show this — threading.RLock alone would pass. + Spawn a peer that holds the lock and assert we actually wait for it. + """ + import subprocess + + repo_root = Path(__file__).resolve().parents[2] + target = tmp_path / "srv.json" + target.write_text("{}", encoding="utf-8") + + holder = tmp_path / "holder.py" + holder.write_text(_LOCK_HOLDER_SRC, encoding="utf-8") + + hold_seconds = 2.0 + proc = subprocess.Popen( + [ + sys.executable, str(holder), str(repo_root), str(tmp_path), + str(target), str(hold_seconds), + ], + stdout=subprocess.PIPE, + text=True, + ) + try: + ready = proc.stdout.readline().strip() + if ready != "HELD": + pytest.skip(f"peer process could not take the lock (got {ready!r})") + + start = time.monotonic() + with _token_store_lock(target): + waited = time.monotonic() - start + finally: + try: + proc.wait(timeout=15) + except subprocess.TimeoutExpired: # pragma: no cover + proc.kill() + + # Allow slack for process startup, but a near-zero wait means the advisory + # lock did nothing and only the in-process RLock was involved. + assert waited > 0.5, ( + f"acquired in {waited:.2f}s while a peer held the lock — " + "cross-process exclusion is not working" + ) diff --git a/tests/tools/test_mcp_oauth_manager.py b/tests/tools/test_mcp_oauth_manager.py index 9bc4496eb3..08c15c57e6 100644 --- a/tests/tools/test_mcp_oauth_manager.py +++ b/tests/tools/test_mcp_oauth_manager.py @@ -539,3 +539,207 @@ async def test_refresh_response_with_new_refresh_token_rotates(tmp_path, monkeyp on_disk = json.loads((tmp_path / "mcp-tokens" / "srv.json").read_text(encoding="utf-8")) assert provider.context.current_tokens.refresh_token == "rt-new" == on_disk["refresh_token"] + + +# --------------------------------------------------------------------------- +# Cross-process refresh-token rotation (single-use refresh tokens) +# +# Two Hermes backends routinely share one HERMES_HOME (desktop `serve` + +# `gateway run`). With a provider that rotates refresh tokens, the loser of the +# race POSTs a token the winner already consumed and gets 400 — while a valid +# replacement sits on disk. Clearing state there forces an interactive browser +# reauth that a cron/background context cannot satisfy. +# --------------------------------------------------------------------------- + + +def _token(access, refresh, expires_in=3600): + from mcp.shared.auth import OAuthToken + + return OAuthToken( + access_token=access, + token_type="Bearer", + expires_in=expires_in, + refresh_token=refresh, + ) + + +@pytest.mark.asyncio +async def test_refresh_400_recovers_token_rotated_by_peer(tmp_path, monkeypatch): + """A peer rotated the refresh token: recover from disk instead of clearing.""" + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + provider = _provider_with_token_endpoint( + tmp_path, {}, "https://idp.example.com/oauth/token", monkeypatch + ) + + # We hold R1 in memory and are about to fail with it. + provider.context.current_tokens = _token("A1", "R1") + # The peer process already persisted its replacement. + await provider.context.storage.set_tokens(_token("A2", "R2")) + + resp = _fake_response( + 400, "https://idp.example.com/oauth/token", b'{"error":"invalid_grant"}' + ) + result = await provider._handle_refresh_response(resp) + + assert result is True, "a rotated-token race must be recoverable" + assert provider.context.current_tokens.access_token == "A2" + assert provider.context.current_tokens.refresh_token == "R2" + + +@pytest.mark.asyncio +async def test_refresh_400_rejects_disk_token_without_refresh_token( + tmp_path, monkeypatch +): + """A disk token with no refresh token is a dead end, not a recovery. + + Its access token may still be inside its TTL, so the naive "is it + different and currently valid?" test says yes — but adopting it only + defers the reauth to expiry, with no way to refresh in between. Recovery + must require a refresh token to recover *onto*. + """ + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + provider = _provider_with_token_endpoint( + tmp_path, {}, "https://idp.example.com/oauth/token", monkeypatch + ) + + provider.context.current_tokens = _token("A1", "R1") + # Different access token, still valid, but nothing to refresh with later. + await provider.context.storage.set_tokens(_token("A2", None)) + + resp = _fake_response( + 400, "https://idp.example.com/oauth/token", b'{"error":"invalid_grant"}' + ) + result = await provider._handle_refresh_response(resp) + + assert result is False, "a token with no refresh token must not be adopted" + assert provider.context.current_tokens is None + + +@pytest.mark.asyncio +async def test_refresh_400_does_not_strand_a_rejected_token_in_the_context( + tmp_path, monkeypatch +): + """A rejected candidate must not be left installed on the context. + + is_token_valid() reads the context, so the candidate has to be published + to be tested. This asserts on the state the recovery helper itself leaves + behind, because the caller's clear_tokens() would otherwise mask the + difference: without the restore, current_tokens still points at the + rejected candidate when the helper returns. + """ + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + provider = _provider_with_token_endpoint( + tmp_path, {}, "https://idp.example.com/oauth/token", monkeypatch + ) + + stale = _token("A1", "R1") + provider.context.current_tokens = stale + await provider.context.storage.set_tokens(_token("A2", "R2")) + + seen = [] + provider.context.is_token_valid = lambda: ( + seen.append(provider.context.current_tokens) or False + ) + + recovered = await provider._hermes_reload_tokens_after_refresh_failure() + + assert recovered is False + assert seen and seen[0].access_token == "A2", "candidate must be testable" + assert provider.context.current_tokens is stale, ( + "a rejected candidate must not be left on the context" + ) + + + +@pytest.mark.asyncio +async def test_refresh_400_still_clears_when_disk_is_same_token(tmp_path, monkeypatch): + + """No peer wrote anything: the credential really is dead — clear it.""" + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + provider = _provider_with_token_endpoint( + tmp_path, {}, "https://idp.example.com/oauth/token", monkeypatch + ) + + provider.context.current_tokens = _token("A1", "R1") + await provider.context.storage.set_tokens(_token("A1", "R1")) + + resp = _fake_response( + 400, "https://idp.example.com/oauth/token", b'{"error":"invalid_grant"}' + ) + result = await provider._handle_refresh_response(resp) + + assert result is False + assert provider.context.current_tokens is None + + +@pytest.mark.asyncio +async def test_refresh_400_does_not_recover_expired_disk_token(tmp_path, monkeypatch): + """A *different* but already-expired disk token is not a recovery.""" + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + provider = _provider_with_token_endpoint( + tmp_path, {}, "https://idp.example.com/oauth/token", monkeypatch + ) + + provider.context.current_tokens = _token("A1", "R1") + await provider.context.storage.set_tokens(_token("A2", "R2", expires_in=-60)) + + resp = _fake_response( + 400, "https://idp.example.com/oauth/token", b'{"error":"invalid_grant"}' + ) + result = await provider._handle_refresh_response(resp) + + assert result is False + assert provider.context.current_tokens is None + + +@pytest.mark.asyncio +async def test_refresh_400_does_not_recover_tokenless_disk_entry( + tmp_path, monkeypatch +): + """A disk entry without an access token is not a recovery.""" + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + provider = _provider_with_token_endpoint( + tmp_path, {}, "https://idp.example.com/oauth/token", monkeypatch + ) + + provider.context.current_tokens = _token("A1", "R1") + # Rotated refresh token, but the access token is empty — recovering here + # would ship an Authorization header with no credential. + await provider.context.storage.set_tokens(_token("", "R2")) + + resp = _fake_response( + 400, "https://idp.example.com/oauth/token", b'{"error":"invalid_grant"}' + ) + result = await provider._handle_refresh_response(resp) + + assert result is False + assert provider.context.current_tokens is None + + +@pytest.mark.asyncio +async def test_refresh_400_recovery_never_logs_token_material( + tmp_path, monkeypatch, caplog +): + """The recovery path must not leak secrets into logs.""" + import logging + + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + provider = _provider_with_token_endpoint( + tmp_path, {}, "https://idp.example.com/oauth/token", monkeypatch + ) + + provider.context.current_tokens = _token("access-secret", "refresh-secret") + await provider.context.storage.set_tokens( + _token("rotated-access-secret", "rotated-refresh-secret") + ) + + resp = _fake_response( + 400, "https://idp.example.com/oauth/token", b'{"error":"invalid_grant"}' + ) + with caplog.at_level(logging.DEBUG): + result = await provider._handle_refresh_response(resp) + + assert result is True + assert "refresh-secret" not in caplog.text + assert "rotated-refresh-secret" not in caplog.text + assert "rotated-access-secret" not in caplog.text diff --git a/tools/mcp_oauth.py b/tools/mcp_oauth.py index cfe5cda6f4..2f5fdae179 100644 --- a/tools/mcp_oauth.py +++ b/tools/mcp_oauth.py @@ -21,9 +21,23 @@ import socket import stat import sys import threading +from contextlib import contextmanager as _contextmanager import time import webbrowser from functools import partialmethod + +# Cross-process advisory file locking for the token store's critical sections. +# Mirrors cron/jobs.py: fcntl is Unix-only, msvcrt is the Windows fallback, and +# either may be absent - in which case locking degrades to in-process only +# (the historical behaviour) rather than failing. +try: + import fcntl +except ImportError: # pragma: no cover - non-Unix + fcntl = None +try: + import msvcrt +except ImportError: # pragma: no cover - non-Windows + msvcrt = None from http.server import BaseHTTPRequestHandler, HTTPServer from pathlib import Path from typing import TYPE_CHECKING, Any @@ -39,6 +53,92 @@ if TYPE_CHECKING: # annotations only; the SDK is imported lazily at runtime logger = logging.getLogger(__name__) +# Bounded acquisition for the token-store lock. Deliberately short: every +# critical section it guards is a local file read/write, never a network call. +_TOKEN_LOCK_TIMEOUT_SECONDS = 10.0 + +# In-process mutual exclusion, keyed by lock path, so threads inside one +# process don't fight over the same file before the advisory lock is reached. +_token_locks: dict[str, threading.RLock] = {} +_token_locks_guard = threading.Lock() + + +@_contextmanager +def _token_store_lock(path: "Path"): + """Serialize a read-modify-write on one server's token file. + + Two Hermes backends routinely share one HERMES_HOME (the desktop app spawns + ``serve`` while the scheduled task runs ``gateway run``); cron/jobs.py + already guards jobs.json the same way. Without this, a provider that issues + single-use refresh tokens can have both processes POST the same token, and + the loser's refresh is rejected. + + Acquisition is bounded and non-blocking (the lesson of cron's #60703: a + plain blocking ``flock`` with no timeout lets one wedged process freeze + every other one forever). On timeout we log and proceed with in-process + locking only — a briefly-contended refresh is strictly better than a + permanently stuck client. + """ + lock_path = path.with_suffix(path.suffix + ".lock") + key = str(lock_path) + + with _token_locks_guard: + local_lock = _token_locks.setdefault(key, threading.RLock()) + + with local_lock: + lock_fd = None + acquired = False + try: + try: + secure_parent_dir(lock_path) + lock_fd = open(lock_path, "a+", encoding="utf-8") + lock_fd.seek(0) + deadline = time.monotonic() + _TOKEN_LOCK_TIMEOUT_SECONDS + while True: + try: + if fcntl is not None: + fcntl.flock(lock_fd, fcntl.LOCK_EX | fcntl.LOCK_NB) + acquired = True + elif msvcrt is not None: + getattr(msvcrt, "locking")( + lock_fd.fileno(), getattr(msvcrt, "LK_NBLCK"), 1 + ) + acquired = True + break + except (OSError, IOError): + if time.monotonic() >= deadline: + logger.warning( + "Token store lock timed out after %.0fs (%s); " + "proceeding with in-process locking only", + _TOKEN_LOCK_TIMEOUT_SECONDS, + lock_path.name, + ) + break + time.sleep(0.05) + except OSError as exc: + # An unwritable lock path must never block token access. + logger.debug("Token store lock unavailable (%s): %s", lock_path.name, exc) + + yield + finally: + if lock_fd is not None: + try: + if acquired: + if fcntl is not None: + fcntl.flock(lock_fd, fcntl.LOCK_UN) + elif msvcrt is not None: + getattr(msvcrt, "locking")( + lock_fd.fileno(), getattr(msvcrt, "LK_UNLCK"), 1 + ) + except (OSError, IOError): + pass + finally: + lock_fd.close() + +# --------------------------------------------------------------------------- +# Lazy imports -- MCP SDK with OAuth support is optional +# --------------------------------------------------------------------------- + # SDK availability is detected WITHOUT importing mcp (~170 ms); classes bind lazily via _sdk_class(). _OAUTH_AVAILABLE = _importlib_util.find_spec("mcp") is not None if not _OAUTH_AVAILABLE: @@ -338,7 +438,10 @@ class HermesTokenStorage: async def get_tokens(self) -> "OAuthToken | None": self.loaded_issuer = None - return self._load_model(self._tokens_path(), "OAuthToken", "tokens", self._fixup_loaded_tokens) + # Held across the read so a peer process mid-rotation cannot expose a + # half-written token file (see _token_store_lock). + with _token_store_lock(self._tokens_path()): + return self._load_model(self._tokens_path(), "OAuthToken", "tokens", self._fixup_loaded_tokens) async def set_tokens(self, tokens: "OAuthToken") -> None: payload = _model_json(tokens) @@ -349,7 +452,8 @@ class HermesTokenStorage: if self._bound_issuer: # which authorization server granted these tokens (never sent on the wire) payload["hermes_issuer"] = self._bound_issuer self.loaded_issuer = self._bound_issuer - _write_json(self._tokens_path(), payload) + with _token_store_lock(self._tokens_path()): + _write_json(self._tokens_path(), payload) logger.debug("OAuth tokens saved for %s", self._server_name) def bind_issuer(self, issuer: str | None) -> None: diff --git a/tools/mcp_oauth_provider.py b/tools/mcp_oauth_provider.py index c51a9cb484..ac2807c975 100644 --- a/tools/mcp_oauth_provider.py +++ b/tools/mcp_oauth_provider.py @@ -110,6 +110,15 @@ class HermesProviderMixin: """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) + # A peer process (gateway vs desktop sharing one HERMES_HOME) may have + # rotated the refresh token microseconds ago and already persisted the + # replacement. Providers issuing single-use refresh tokens reject our + # now-stale copy with a 400. Re-read disk before destroying the session. + if await self._hermes_reload_tokens_after_refresh_failure(): + self._hermes_logger.info( + "Recovered a peer-rotated refresh token instead of clearing the session" + ) + return True self.context.clear_tokens() return False from httpx import HTTPError @@ -134,6 +143,85 @@ class HermesProviderMixin: await self._store_tokens(token_response) return True + async def _hermes_reload_tokens_after_refresh_failure(self) -> bool: + """Re-read tokens from disk after a rejected refresh. + + Returns True only when disk holds a token that is BOTH different + from the one we just failed with AND still valid. That is the + signature of a peer process having rotated the refresh token + between our read and our POST — a recoverable race, not a dead + credential. + + Returns False for the genuinely-expired case (nobody else wrote a + newer token), so the caller still clears state and surfaces the + reauth prompt. Never raises: a failure to recover must degrade to + the pre-existing clear-and-reauth path. + """ + try: + storage = getattr(self.context, "storage", None) + if storage is None: + return False + + stale = getattr(self.context, "current_tokens", None) + stale_refresh = getattr(stale, "refresh_token", None) + + fresh = await storage.get_tokens() + if fresh is None: + return False + + fresh_refresh = getattr(fresh, "refresh_token", None) + # Same credential we just failed with: disk has nothing newer. + if fresh_refresh is not None and fresh_refresh == stale_refresh: + return False + + # No refresh token on the disk entry. The access token may + # still be usable right now, but adopting it would trade an + # explicit reauth today for a silent one at expiry, with no + # way to refresh in between. Treat it as unrecoverable. + if fresh_refresh is None: + return False + + # Defence-in-depth. Not load-bearing: is_token_valid() below + # already rejects an empty access_token, so mutating this + # guard away leaves the suite green. Kept because recovering + # onto a credential-less token would be a security-relevant + # failure if that SDK behaviour ever changed. + access = getattr(fresh, "access_token", None) + if not access: + return False + + # Storage clamps a past-due token to ``expires_in == 0`` on + # read (see HermesTokenStorage.get_tokens). The SDK's + # is_token_valid() compares ``time.time() <= expiry`` and so + # still reports True for that boundary value, which would make + # us "recover" onto a token the server will immediately reject. + # Require a real, positive TTL. + try: + if int(getattr(fresh, "expires_in", 0) or 0) <= 0: + return False + except (TypeError, ValueError): + return False + + # Publish, then restore on rejection. is_token_valid() reads + # the context rather than taking a token argument, so the + # candidate has to be installed to be tested; keeping the + # previous value lets a losing probe leave the context exactly + # as it found it instead of stranding a rejected token there + # for the caller to clean up. + previous_tokens = self.context.current_tokens + self.context.current_tokens = fresh + self.context.update_token_expiry(fresh) + + if not self.context.is_token_valid(): + self.context.current_tokens = previous_tokens + self.context.update_token_expiry(previous_tokens) + return False + + return True + except Exception as exc: # noqa: BLE001 — recovery is best-effort + logger.debug("Post-refresh disk reload failed: %s", exc) + return False + def _metadata_issuer(context: Any) -> str | None: """Discovered authorization-server issuer from the SDK auth context, without trailing slash."""