diff --git a/tests/tools/test_mcp_oauth_manager.py b/tests/tools/test_mcp_oauth_manager.py index 2fd484866c..ba3d6fc261 100644 --- a/tests/tools/test_mcp_oauth_manager.py +++ b/tests/tools/test_mcp_oauth_manager.py @@ -848,3 +848,57 @@ async def test_refresh_fails_closed_while_a_peer_holds_the_fence(tmp_path, monke assert provider.context.current_tokens.refresh_token == "R1" assert (await provider.context.storage.get_tokens()).refresh_token == "R1" assert provider._hermes_fence is None + + +@pytest.mark.asyncio +async def test_refresh_adopts_expired_peer_pair_and_posts_its_refresh_token(tmp_path, monkeypatch): + """A peer rotated to (A2, R2) but A2 already expired: we must POST R2, never R1. + + The adopt path installs the rotated pair even without a live access + token, because the POST we are about to build needs the new grant. + """ + from urllib.parse import parse_qs + + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + endpoint = "https://idp.example.com/oauth/token" + provider = _fenced_provider(tmp_path, monkeypatch, endpoint) + await provider.context.storage.set_tokens(_token("A2", "R2", expires_in=0)) + + presented = [] + + def responder(request): + if request.method != "POST": + return _fake_response(200, str(request.url), b"{}") + presented.append(parse_qs(request.content.decode())["refresh_token"][0]) + body = json.dumps(_token("A3", "R3").model_dump(mode="json", exclude_none=True)).encode() + return _fake_response(200, endpoint, body) + + await _drive_flow(provider, responder) + + assert presented == ["R2"], presented + assert provider.context.current_tokens.refresh_token == "R3" + + +@pytest.mark.asyncio +async def test_refresh_restarts_flow_when_disk_pair_is_from_another_issuer(tmp_path, monkeypatch): + """A disk pair bound to a different issuer loses its refresh token on adoption. + + With nothing left to refresh, _refresh_token must restart the SDK flow + (401 -> full auth) instead of building a POST from the foreign grant or + raising OAuthTokenError, and it must not keep the fence. + """ + from tools.mcp_oauth_provider import _RefreshCompletedByPeer + + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + endpoint = "https://idp.example.com/oauth/token" + provider = _fenced_provider(tmp_path, monkeypatch, endpoint) + storage = provider.context.storage + storage.bind_issuer("https://other-idp.example.com") + await storage.set_tokens(_token("A2", "R2")) + + with pytest.raises(_RefreshCompletedByPeer): + await provider._refresh_token() + + assert not provider.context.current_tokens.refresh_token, "foreign refresh token must be stripped" + assert (await storage.get_tokens()).refresh_token is None, "strip must reach disk" + assert provider._hermes_fence is None diff --git a/tools/mcp_oauth_provider.py b/tools/mcp_oauth_provider.py index 019aec8309..3b1c196e63 100644 --- a/tools/mcp_oauth_provider.py +++ b/tools/mcp_oauth_provider.py @@ -145,7 +145,7 @@ class HermesProviderMixin: # 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. adopted = await self._hermes_adopt_tokens_from_disk() - if adopted and self.context.is_token_valid(): + if adopted and self._hermes_live_ttl() 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. @@ -156,6 +156,19 @@ class HermesProviderMixin: await self._hermes_release_refresh_fence() raise + def _hermes_live_ttl(self) -> bool: + """True when the installed token has a real, positive TTL. + + 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 still reports True for that boundary, + which would make us adopt a token the server rejects immediately. + """ + try: + return int(getattr(self.context.current_tokens, "expires_in", 0) or 0) > 0 + except (TypeError, ValueError): + return False + async def _hermes_acquire_refresh_fence(self) -> None: """Enter the fence, or let RefreshFenceTimeout abort this attempt. @@ -184,25 +197,55 @@ class HermesProviderMixin: 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_rotated_candidate(self): + """The on-disk pair, if a peer rotated it past the one we hold. + + A candidate must carry a refresh token different from ours (same + token: disk has nothing newer) and a non-empty access token. A disk + entry with no refresh token is never adopted: its access token may + still be inside its TTL, but taking it trades an explicit reauth now + for a silent one at expiry with no way to refresh in between. + ``get_tokens`` already returns None for absent or corrupt files. + """ + stored = await self.context.storage.get_tokens() + if stored is None: + return None + current = self.context.current_tokens + stored_refresh = getattr(stored, "refresh_token", None) + if not stored_refresh or not getattr(stored, "access_token", None): + return None + if stored_refresh == getattr(current, "refresh_token", None): + return None + return stored + + def _hermes_install_disk_pair(self, tokens) -> None: + """Publish a disk pair to the context and re-run issuer binding on it. + + If the enforcer strips the refresh token (the pair was minted by a + different issuer) there is nothing left to refresh with: restart the + SDK flow so it lands in 401 -> full authorization instead of raising + OAuthTokenError over an unusable grant. + """ + self.context.current_tokens = tokens + self.context.update_token_expiry(tokens) + enforce_refresh_token_issuer(self.context) + if not getattr(self.context.current_tokens, "refresh_token", None): + raise _RefreshCompletedByPeer + 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. Returns True when the - in-memory pair was replaced. + Called under the fence. If disk already holds a different refresh + token, the peer that held the fence before us won this generation; + its value is the only one the provider will still accept, so it is + installed even when its access token has already expired (the POST + we are about to build needs the new refresh token). 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 + candidate = await self._hermes_rotated_candidate() + if candidate is None: return False - if stored is None: - return False - current = self.context.current_tokens - if current is not None and getattr(stored, "refresh_token", None) == getattr(current, "refresh_token", None): - return False - self.context.current_tokens = stored - self.context.update_token_expiry(stored) + self._hermes_install_disk_pair(candidate) return True async def _initialize(self) -> None: @@ -251,10 +294,11 @@ class HermesProviderMixin: async def _hermes_handle_refresh_response(self, response) -> bool: 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. + # A writer outside the fence (interactive `hermes mcp login`, or a + # pre-fence Hermes sharing this HERMES_HOME) may have rotated the + # grant and persisted the replacement. Providers issuing single-use + # refresh tokens reject our 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" @@ -287,81 +331,30 @@ class HermesProviderMixin: 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 True only when disk holds a pair that is BOTH different from + the one we just failed with AND still live. That is the signature of + a writer outside the fence (an interactive ``hermes mcp login`` or a + pre-fence Hermes) having rotated the grant 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. + Returns False for the genuinely-expired case (nobody wrote a newer + pair), so the caller still clears state and surfaces the reauth + prompt. """ - 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) + candidate = await self._hermes_rotated_candidate() + if candidate is None: 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; a losing probe must leave the context + # exactly as it found it. + previous_tokens = self.context.current_tokens + self._hermes_install_disk_pair(candidate) + if self._hermes_live_ttl() and self.context.is_token_valid(): + return True + self.context.current_tokens = previous_tokens + self.context.update_token_expiry(previous_tokens) + return False def _metadata_issuer(context: Any) -> str | None: