fix(mcp-oauth): make the refresh fence async and skip the POST after adopting a peer's rotation

The fence is entered from the SDK's coroutine-driven auth flow, so its
acquire loop spun on time.sleep(0.05) and blocked the whole event loop
for up to 60 s while a peer finished its network round trip. The fence
is now an async context manager that polls a non-blocking flock with
asyncio.sleep, and it creates the token directory itself (the parent may
not exist yet on a first refresh; secure_parent_dir only chmods).

The in-process RLock layer is dropped: an advisory lock on a fresh
descriptor already excludes sibling tasks and threads of the same
process, and a thread RLock is reentrant across asyncio tasks on one
thread, so it excluded nothing there anyway.

After acquiring the fence the provider re-reads the store; when a peer
already rotated the pair and the adopted access token is valid, it
raises _RefreshCompletedByPeer instead of building the refresh request.
The auth-flow wrapper restarts the SDK flow so the original request goes
out with the winner's access token. Previously the loser adopted the
new pair but still presented its stale refresh token, burning a
generation on every single-use provider.

Design lifted from #71715.

Co-authored-by: Kevin Yin <182213728+yinkev@users.noreply.github.com>
This commit is contained in:
kshitijk4poor
2026-09-14 23:17:19 +05:30
committed by kshitij
parent a0810c9cc9
commit 1a1345aba4
2 changed files with 113 additions and 102 deletions

View File

@@ -153,8 +153,8 @@ class RefreshFenceTimeout(RuntimeError):
"""
@_contextmanager
def _refresh_fence(path: "Path", *, timeout: float = _REFRESH_FENCE_TIMEOUT_SECONDS):
@contextlib.asynccontextmanager
async def _refresh_fence(path: "Path", *, timeout: float = _REFRESH_FENCE_TIMEOUT_SECONDS):
"""Own one refresh generation across read -> POST -> persist.
``_token_store_lock`` is deliberately narrow: it makes a single file
@@ -176,75 +176,66 @@ def _refresh_fence(path: "Path", *, timeout: float = _REFRESH_FENCE_TIMEOUT_SECO
A separate ``.refresh.lock`` sibling keeps the two scopes independent:
the fence holder can still call get_tokens()/set_tokens() normally.
Entered from the SDK's coroutine-driven auth flow, so the wait is an
``asyncio.sleep`` poll on a non-blocking lock: a peer's slow network
round trip must not freeze every other task on this event loop. No
in-process lock layer is needed: an advisory lock on a fresh descriptor
already excludes sibling tasks and threads of the same process.
Unlike ``_token_store_lock``, acquisition failure RAISES. Degrading to
"proceed unlocked" here would reintroduce the exact race.
"""
lock_path = path.with_suffix(path.suffix + ".refresh.lock")
key = str(lock_path)
with _token_locks_guard:
local_lock = _token_locks.setdefault(key, threading.RLock())
# Bound the in-process wait too: a sibling thread holding the fence is
# just as capable of stranding us as a sibling process.
if not local_lock.acquire(timeout=timeout):
raise RefreshFenceTimeout(
f"refresh fence busy in this process after {timeout:.0f}s ({lock_path.name})"
)
try:
lock_fd = None
acquired = False
try:
lock_path.parent.mkdir(parents=True, exist_ok=True)
secure_parent_dir(lock_path)
lock_fd = open(lock_path, "a+", encoding="utf-8")
except OSError as exc:
# No lock file means no ownership proof. Fail closed: see the
# class docstring for why proceeding is worse than aborting.
raise RefreshFenceTimeout(
f"refresh fence unavailable ({lock_path.name}): {exc}"
) from exc
acquired = False
try:
lock_fd.seek(0)
deadline = time.monotonic() + timeout
while True:
try:
secure_parent_dir(lock_path)
lock_fd = open(lock_path, "a+", encoding="utf-8")
lock_fd.seek(0)
except OSError as exc:
# No lock file means no ownership proof. Fail closed: see the
# class docstring for why proceeding is worse than aborting.
raise RefreshFenceTimeout(
f"refresh fence unavailable ({lock_path.name}): {exc}"
) from exc
if fcntl is not None:
fcntl.flock(lock_fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
elif msvcrt is not None:
getattr(msvcrt, "locking")(
lock_fd.fileno(), getattr(msvcrt, "LK_NBLCK"), 1
)
else: # pragma: no cover - no advisory locking primitive
raise RefreshFenceTimeout(
"refresh fence unsupported: no flock/msvcrt on this platform"
)
acquired = True
break
except (OSError, IOError):
if time.monotonic() >= deadline:
raise RefreshFenceTimeout(
f"refresh fence held by a peer for {timeout:.0f}s ({lock_path.name})"
) from None
await asyncio.sleep(0.05)
deadline = time.monotonic() + timeout
while True:
try:
if fcntl is not None:
fcntl.flock(lock_fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
elif msvcrt is not None:
getattr(msvcrt, "locking")(
lock_fd.fileno(), getattr(msvcrt, "LK_NBLCK"), 1
)
else: # pragma: no cover - no advisory locking primitive
raise RefreshFenceTimeout(
"refresh fence unsupported: no flock/msvcrt on this platform"
)
acquired = True
break
except (OSError, IOError):
if time.monotonic() >= deadline:
raise RefreshFenceTimeout(
f"refresh fence held by a peer for {timeout:.0f}s ({lock_path.name})"
) from None
time.sleep(0.05)
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()
yield
finally:
local_lock.release()
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()
# ---------------------------------------------------------------------------

View File

@@ -16,6 +16,10 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__)
class _RefreshCompletedByPeer(Exception):
"""Restart the SDK auth flow: a peer rotated the grant we were about to present."""
class HermesProviderMixin:
"""Token-endpoint fixes layered over the SDK's ``OAuthClientProvider`` (must precede it in
the MRO; subclasses set ``_hermes_logger`` to keep their own logger name).
@@ -84,32 +88,41 @@ class HermesProviderMixin:
``_handle_refresh_response`` never runs. Without this wrapper the fence
file would stay locked for the life of the process and every later
refresh would fail closed -- trading a race for a deadlock.
When a peer already rotated the grant while we waited for the fence,
``_refresh_token`` adopts it and raises ``_RefreshCompletedByPeer``;
the SDK flow is then restarted so it re-evaluates validity and sends
the original request with the winner's access token, never POSTing
the burned refresh token.
"""
inner = super().async_auth_flow(request)
try:
sent, thrown = None, None
while True:
try:
if thrown is not None:
exc, thrown = thrown, None
out = await inner.athrow(exc)
else:
out = await inner.asend(sent)
except StopAsyncIteration:
return
# Full bidirectional delegation: the SDK drives this flow with
# asend(response), so `async for` would swallow the response and
# feed the inner generator None. Async generators have no
# `yield from`, hence the manual pump.
try:
sent = yield out
except GeneratorExit:
await inner.aclose()
raise
except BaseException as exc:
sent, thrown = None, exc
finally:
self._hermes_release_refresh_fence()
while True:
inner = super().async_auth_flow(request)
try:
sent, thrown = None, None
while True:
try:
if thrown is not None:
exc, thrown = thrown, None
out = await inner.athrow(exc)
else:
out = await inner.asend(sent)
except StopAsyncIteration:
return
except _RefreshCompletedByPeer:
break
# Full bidirectional delegation: the SDK drives this flow with
# asend(response), so `async for` would swallow the response and
# feed the inner generator None. Async generators have no
# `yield from`, hence the manual pump.
try:
sent = yield out
except GeneratorExit:
await inner.aclose()
raise
except BaseException as exc:
sent, thrown = None, exc
finally:
await self._hermes_release_refresh_fence()
async def _refresh_token(self):
"""Take the refresh fence, then build the request from the token we own.
@@ -127,18 +140,23 @@ class HermesProviderMixin:
refreshes inside one process.
"""
self._coerce_client_secret_post()
self._hermes_acquire_refresh_fence()
await self._hermes_acquire_refresh_fence()
try:
# Re-read under the fence: a peer may have rotated while we waited
# for it, in which case the token we were about to POST is dead.
await self._hermes_adopt_tokens_from_disk()
adopted = await self._hermes_adopt_tokens_from_disk()
if adopted and self.context.is_token_valid():
# The peer's access token is live: presenting our copy of the
# refresh token would only burn a generation on a single-use
# provider. Skip the POST and let the flow restart.
raise _RefreshCompletedByPeer
return self._prepare_token_request(await super()._refresh_token())
except BaseException:
# Never hold the fence when no POST will follow.
self._hermes_release_refresh_fence()
await self._hermes_release_refresh_fence()
raise
def _hermes_acquire_refresh_fence(self) -> None:
async def _hermes_acquire_refresh_fence(self) -> None:
"""Enter the fence, or let RefreshFenceTimeout abort this attempt.
Fails closed on purpose: a refresh we are not certain we own must not
@@ -147,43 +165,45 @@ class HermesProviderMixin:
"""
from tools.mcp_oauth import _refresh_fence
self._hermes_release_refresh_fence()
await self._hermes_release_refresh_fence()
storage = self.context.storage
tokens_path = getattr(storage, "_tokens_path", None)
if tokens_path is None: # pragma: no cover - non-Hermes storage
return
fence = _refresh_fence(tokens_path())
fence.__enter__()
await fence.__aenter__()
self._hermes_fence = fence
def _hermes_release_refresh_fence(self) -> None:
async def _hermes_release_refresh_fence(self) -> None:
"""Release the fence if held. Idempotent and never raises."""
fence, self._hermes_fence = self._hermes_fence, None
if fence is None:
return
try:
fence.__exit__(None, None, None)
await fence.__aexit__(None, None, None)
except Exception: # pragma: no cover - release must never mask the outcome
self._hermes_logger.debug("Refresh fence release failed", exc_info=True)
async def _hermes_adopt_tokens_from_disk(self) -> None:
async def _hermes_adopt_tokens_from_disk(self) -> bool:
"""Adopt a peer's newer tokens before POSTing our own copy.
Called under the fence. If disk already holds a different token, the
peer that held the fence before us won this generation; its value is
the only one the provider will still accept.
the only one the provider will still accept. Returns True when the
in-memory pair was replaced.
"""
try:
stored = await self.context.storage.get_tokens()
except Exception: # pragma: no cover - unreadable store: keep what we have
return
return False
if stored is None:
return
return False
current = self.context.current_tokens
if current is not None and getattr(stored, "refresh_token", None) == getattr(current, "refresh_token", None):
return
return False
self.context.current_tokens = stored
self.context.update_token_expiry(stored)
return True
async def _initialize(self) -> None:
"""Load stored state, restore persisted server metadata when the SDK has none (so the issuer
@@ -226,7 +246,7 @@ class HermesProviderMixin:
try:
return await self._hermes_handle_refresh_response(response)
finally:
self._hermes_release_refresh_fence()
await self._hermes_release_refresh_fence()
async def _hermes_handle_refresh_response(self, response) -> bool:
if not (200 <= response.status_code < 300):