fix(mcp-oauth): keep refresh_token when a refresh response omits it (#62333)
HermesProviderMixin._handle_refresh_response overrides the SDK's handler (to accept any 2xx and keep token bodies out of logs) but dropped the SDK's RFC 6749 section 6 carry-forward. An authorization server that does not rotate refresh tokens (TinyFish, Google, Zoho, Asana, Futu) answers the refresh grant without a refresh_token; we then stored the response verbatim, erasing the only refresh token we had, so the next expiry had nothing to refresh with and forced a browser re-auth roughly one TTL after every login. Carry the prior refresh_token (and scope, per section 5.1) forward on the OAuthToken before _store_tokens, so both the live provider and the on-disk token file keep it. A rotating AS still wins: only None fields are filled. Tests: two invariants on the real HermesMCPOAuthProvider + HermesTokenStorage (omitted -> preserved in memory and on disk; provided -> rotated). The carry-forward test is red on main.
This commit is contained in:
@@ -27,7 +27,7 @@ def test_manager_isolates_same_named_servers_by_profile_home(tmp_path, monkeypat
|
||||
storage._tokens_path().write_text(
|
||||
'{"access_token":"%s","token_type":"Bearer","expires_in":3600}'
|
||||
% access_token
|
||||
)
|
||||
, encoding="utf-8")
|
||||
finally:
|
||||
reset_hermes_home_override(token)
|
||||
|
||||
@@ -102,7 +102,7 @@ async def test_disk_watch_invalidates_on_mtime_change(tmp_path, monkeypatch):
|
||||
tokens_file.write_text(json.dumps({
|
||||
"access_token": "OLD",
|
||||
"token_type": "Bearer",
|
||||
}))
|
||||
}), encoding="utf-8")
|
||||
|
||||
mgr = MCPOAuthManager()
|
||||
provider = mgr.get_or_build_provider("srv", "https://example.com/mcp", None)
|
||||
@@ -270,8 +270,8 @@ def test_invalid_client_at_token_endpoint_poisons(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
d = tmp_path / "mcp-tokens"
|
||||
d.mkdir(parents=True)
|
||||
(d / "srv.client.json").write_text('{"client_id": "dead"}')
|
||||
(d / "srv.meta.json").write_text("{}")
|
||||
(d / "srv.client.json").write_text('{"client_id": "dead"}', encoding="utf-8")
|
||||
(d / "srv.meta.json").write_text("{}", encoding="utf-8")
|
||||
provider = _provider_with_token_endpoint(
|
||||
tmp_path, {}, "https://idp.example.com/oauth/token", monkeypatch
|
||||
)
|
||||
@@ -292,7 +292,7 @@ def test_invalid_client_metadata_does_not_trip(tmp_path, monkeypatch):
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
d = tmp_path / "mcp-tokens"
|
||||
d.mkdir(parents=True)
|
||||
(d / "srv.client.json").write_text('{"client_id": "live"}')
|
||||
(d / "srv.client.json").write_text('{"client_id": "live"}', encoding="utf-8")
|
||||
provider = _provider_with_token_endpoint(
|
||||
tmp_path, {}, "https://idp.example.com/oauth/token", monkeypatch
|
||||
)
|
||||
@@ -328,7 +328,7 @@ def test_bridge_forwards_requests_and_poisons_on_token_endpoint_400(
|
||||
token_ep = "https://idp.example.com/oauth/token"
|
||||
d = tmp_path / "mcp-tokens"
|
||||
d.mkdir(parents=True)
|
||||
(d / "srv.client.json").write_text('{"client_id": "dead"}')
|
||||
(d / "srv.client.json").write_text('{"client_id": "dead"}', encoding="utf-8")
|
||||
|
||||
forwarded = []
|
||||
|
||||
@@ -488,3 +488,54 @@ async def test_manager_refresh_read_error_clears_tokens(tmp_path, monkeypatch):
|
||||
|
||||
assert result is False
|
||||
assert provider.context.current_tokens is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_response_without_refresh_token_keeps_stored_one(tmp_path, monkeypatch):
|
||||
"""RFC 6749 §6: an AS that does not rotate omits refresh_token; the prior one must survive in
|
||||
the live provider AND on disk, or the server dies at the next expiry (#62333)."""
|
||||
import json
|
||||
from mcp.shared.auth import OAuthToken
|
||||
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
provider = _provider_with_token_endpoint(
|
||||
tmp_path, {}, "https://idp.example.com/oauth/token", monkeypatch
|
||||
)
|
||||
provider.context.current_tokens = OAuthToken(
|
||||
access_token="at-1", token_type="Bearer", expires_in=3600, refresh_token="rt-keep", scope="read"
|
||||
)
|
||||
provider.context.client_info = SimpleNamespace(client_id="cid")
|
||||
|
||||
body = b'{"access_token": "at-2", "token_type": "Bearer", "expires_in": 3600}'
|
||||
assert await provider._handle_refresh_response(
|
||||
_fake_response(200, "https://idp.example.com/oauth/token", body)
|
||||
)
|
||||
|
||||
on_disk = json.loads((tmp_path / "mcp-tokens" / "srv.json").read_text(encoding="utf-8"))
|
||||
assert provider.context.current_tokens.access_token == "at-2"
|
||||
assert provider.context.current_tokens.refresh_token == "rt-keep" == on_disk["refresh_token"]
|
||||
assert provider.context.current_tokens.scope == "read" == on_disk["scope"]
|
||||
assert provider.context.can_refresh_token()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_refresh_response_with_new_refresh_token_rotates(tmp_path, monkeypatch):
|
||||
"""A rotating AS's new refresh_token replaces the stored one (carry-forward fills gaps only)."""
|
||||
import json
|
||||
from mcp.shared.auth import OAuthToken
|
||||
|
||||
monkeypatch.setenv("HERMES_HOME", str(tmp_path))
|
||||
provider = _provider_with_token_endpoint(
|
||||
tmp_path, {}, "https://idp.example.com/oauth/token", monkeypatch
|
||||
)
|
||||
provider.context.current_tokens = OAuthToken(
|
||||
access_token="at-1", token_type="Bearer", expires_in=3600, refresh_token="rt-old"
|
||||
)
|
||||
|
||||
body = b'{"access_token": "at-2", "token_type": "Bearer", "expires_in": 3600, "refresh_token": "rt-new"}'
|
||||
assert await provider._handle_refresh_response(
|
||||
_fake_response(200, "https://idp.example.com/oauth/token", body)
|
||||
)
|
||||
|
||||
on_disk = json.loads((tmp_path / "mcp-tokens" / "srv.json").read_text(encoding="utf-8"))
|
||||
assert provider.context.current_tokens.refresh_token == "rt-new" == on_disk["refresh_token"]
|
||||
|
||||
@@ -107,6 +107,16 @@ class HermesProviderMixin:
|
||||
self._hermes_logger.warning("Invalid refresh response: %s", response.status_code)
|
||||
self.context.clear_tokens()
|
||||
return False
|
||||
# RFC 6749 §6: a refresh response may omit refresh_token (AS does not rotate) and scope
|
||||
# (unchanged). The SDK's own _handle_refresh_response carries both forward; this override
|
||||
# must too, or every non-rotating refresh erases the stored refresh_token and the server
|
||||
# dies at the NEXT expiry with a forced browser re-auth (#62333).
|
||||
prior = self.context.current_tokens
|
||||
if prior is not None:
|
||||
if token_response.refresh_token is None:
|
||||
token_response.refresh_token = prior.refresh_token
|
||||
if token_response.scope is None:
|
||||
token_response.scope = prior.scope
|
||||
await self._store_tokens(token_response)
|
||||
return True
|
||||
|
||||
|
||||
Reference in New Issue
Block a user