fix(mcp-oauth): serialize token-store access across processes

The desktop app spawns 'serve' while the scheduled task runs 'gateway run',
so two backends routinely share one HERMES_HOME. With a provider that issues
single-use refresh tokens, both could POST the same token and the loser's
refresh was rejected, clearing an otherwise healthy session.

Guard the token file's read-modify-write with a bounded advisory file lock
(fcntl on Unix, msvcrt on Windows, in-process only where neither exists),
mirroring how cron/jobs.py guards jobs.json. Acquisition is non-blocking with
a 10s ceiling: a briefly-contended refresh beats a permanently stuck client.

(cherry picked from commit 16f0d811de446a66ed5fd061fd7594fca3d230da)
This commit is contained in:
anhtahaylove
2026-09-09 03:24:12 +07:00
committed by kshitij
parent 45760ac022
commit 47ee79a647
4 changed files with 514 additions and 2 deletions

View File

@@ -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"
)

View File

@@ -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

View File

@@ -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:

View File

@@ -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."""