fix(mcp): refresh expired cold-loaded OAuth tokens

This commit is contained in:
KoNit-K
2026-09-15 00:25:40 +08:00
committed by Teknium
parent 417b707f59
commit 20d80bb36f
2 changed files with 78 additions and 1 deletions
@@ -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
+8 -1
View File
@@ -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()