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:
@@ -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()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user