From 20d80bb36f14e643d1477ad3e845cac3bfbd9be9 Mon Sep 17 00:00:00 2001 From: KoNit-K <124019182+KoNit-K@users.noreply.github.com> Date: Tue, 15 Sep 2026 00:25:40 +0800 Subject: [PATCH] fix(mcp): refresh expired cold-loaded OAuth tokens --- .../tools/test_mcp_oauth_cold_load_expiry.py | 70 +++++++++++++++++++ tools/mcp_oauth_manager.py | 9 ++- 2 files changed, 78 insertions(+), 1 deletion(-) diff --git a/tests/tools/test_mcp_oauth_cold_load_expiry.py b/tests/tools/test_mcp_oauth_cold_load_expiry.py index 6e59a59c0c..13017d9502 100644 --- a/tests/tools/test_mcp_oauth_cold_load_expiry.py +++ b/tests/tools/test_mcp_oauth_cold_load_expiry.py @@ -280,6 +280,76 @@ async def test_initialize_seeds_token_expiry_time_from_stored_tokens( assert provider.context.token_expiry_time <= time.time() + 7200 + 5 +@pytest.mark.asyncio +async def test_initialize_marks_zero_ttl_cold_loaded_token_invalid( + tmp_path, monkeypatch +): + """An expired token must not pass the SDK's same-tick validity check. + + ``OAuthContext.update_token_expiry`` maps ``expires_in=0`` to ``time.time()`` + and ``is_token_valid`` accepts equality. Cold-loaded expired tokens therefore + need an expiry that is already in the past before the SDK selects its refresh + or authorization-code path. + """ + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + from mcp.shared.auth import OAuthClientInformationFull, OAuthClientMetadata, OAuthToken + from pydantic import AnyUrl + + from tools.mcp_oauth import HermesTokenStorage, _get_token_dir + from tools.mcp_oauth_manager import _HERMES_PROVIDER_CLS, reset_manager_for_tests + + assert _HERMES_PROVIDER_CLS is not None + reset_manager_for_tests() + + storage = HermesTokenStorage("srv") + await storage.set_tokens( + OAuthToken( + access_token="expired-access", + token_type="Bearer", + expires_in=3600, + refresh_token="refresh-token", + ) + ) + token_path = _get_token_dir() / "srv.json" + persisted = json.loads(token_path.read_text()) + fixed_now = time.time() + persisted["expires_at"] = fixed_now - 60 + token_path.write_text(json.dumps(persisted)) + await storage.set_client_info( + OAuthClientInformationFull( + client_id="test-client", + redirect_uris=[AnyUrl("http://127.0.0.1:12345/callback")], + grant_types=["authorization_code", "refresh_token"], + response_types=["code"], + token_endpoint_auth_method="none", + ) + ) + + provider = _HERMES_PROVIDER_CLS( + server_name="srv", + server_url="https://example.com/mcp", + client_metadata=OAuthClientMetadata( + redirect_uris=[AnyUrl("http://127.0.0.1:12345/callback")], + client_name="Hermes Agent", + ), + storage=storage, + redirect_handler=_noop_redirect, + callback_handler=_noop_callback, + ) + + # Reproduce the SDK's equality boundary deterministically: its zero-TTL + # expiry is exactly ``time.time()``, and validity accepts ``<=``. + monkeypatch.setattr("mcp.client.auth.oauth2.time.time", lambda: fixed_now) + await provider._initialize() + + assert provider.context.current_tokens is not None + assert provider.context.current_tokens.expires_in == 0 + assert not provider.context.is_token_valid(), ( + "An expired cold-loaded token must be invalid before the SDK chooses " + "between refresh and authorization-code flow." + ) + + async def _noop_redirect(_url: str) -> None: return None diff --git a/tools/mcp_oauth_manager.py b/tools/mcp_oauth_manager.py index f88ab41a5e..721756a805 100644 --- a/tools/mcp_oauth_manager.py +++ b/tools/mcp_oauth_manager.py @@ -11,6 +11,7 @@ import asyncio import logging import re import threading +import time from dataclasses import dataclass, field from pathlib import Path from typing import Any, Optional @@ -81,7 +82,13 @@ class HermesMCPOAuthProvider(HermesProviderMixin, *_SDK_BASES): await super()._initialize() # HermesProviderMixin: restores metadata from disk, enforces issuer binding tokens = self.context.current_tokens if tokens is not None and tokens.expires_in is not None: - self.context.update_token_expiry(tokens) + # The SDK maps a zero TTL to ``time.time()`` and accepts equality + # in ``is_token_valid()``. On a cold load that same-tick boundary + # can send an already-expired access token instead of refreshing it. + if tokens.expires_in <= 0: + self.context.token_expiry_time = time.time() - 1 + else: + self.context.update_token_expiry(tokens) if tokens is not None and self.context.oauth_metadata is None: try: await self._prefetch_oauth_metadata()