fix(mcp): wrap the redirect handler for Google offline access once and fold the tests to two

Review follow-up: _request_google_offline_access re-wrapped
context.redirect_handler on every _perform_authorization, so N
re-authorizations ran N nested wrappers. Mark the wrapper and return early
when it is already installed; it reads the issuer at call time so the
single wrapper is right for any issuer. The six test functions fold into
two — browser flow (offline access + consent once, wrap-once, other-issuer
and lookalike controls) and device flow (issuer slash-normalization +
access_type on the device request).
This commit is contained in:
teknium1
2026-09-20 15:19:40 -07:00
committed by Teknium
parent 5be70b5e25
commit 2a00b06212
2 changed files with 76 additions and 95 deletions

View File

@@ -89,120 +89,96 @@ async def _run_browser_flow(tmp_path, monkeypatch, *, issuer, authorization_serv
async with httpx.AsyncClient(auth=provider, transport=httpx.MockTransport(_standin(
httpx, issuer=issuer, authorization_servers=authorization_servers))) as client:
response = await client.get(RESOURCE)
seen["provider"] = provider
return response, seen
@pytest.mark.asyncio
async def test_google_authorization_url_asks_for_offline_access(tmp_path, monkeypatch):
response, seen = await _run_browser_flow(tmp_path, monkeypatch, issuer=GOOGLE, authorization_servers=[GOOGLE])
assert response.status_code == 200
query = dict(parse_qsl(urlsplit(seen["authorize_url"]).query))
assert query["access_type"] == "offline"
assert query["prompt"] == "consent"
assert (tmp_path / "mcp-tokens" / "srv.json").exists()
@pytest.mark.asyncio
async def test_other_issuer_authorization_url_keeps_sdk_parameters(tmp_path, monkeypatch):
response, seen = await _run_browser_flow(tmp_path, monkeypatch, issuer=OTHER_AS, authorization_servers=[OTHER_AS])
assert response.status_code == 200
query = parse_qsl(urlsplit(seen["authorize_url"]).query)
assert not any(key in ("access_type", "prompt") for key, _ in query)
@pytest.mark.asyncio
async def test_google_params_never_duplicate_an_existing_prompt(tmp_path, monkeypatch):
_, seen = await _run_browser_flow(
response, seen = await _run_browser_flow(
tmp_path, monkeypatch, issuer=GOOGLE, authorization_servers=[GOOGLE], scope="email offline_access")
assert response.status_code == 200
pairs = parse_qsl(urlsplit(seen["authorize_url"]).query)
assert [value for key, value in pairs if key == "prompt"] == ["consent"]
# Exactly once each, even with an offline_access scope already requested.
assert [value for key, value in pairs if key == "access_type"] == ["offline"]
assert [value for key, value in pairs if key == "prompt"] == ["consent"]
assert (tmp_path / "mcp-tokens" / "srv.json").exists()
# Wraps once: a re-authorization must not nest another wrapper around the redirect handler.
provider = seen["provider"]
wrapped = provider.context.redirect_handler
provider._request_google_offline_access()
assert provider.context.redirect_handler is wrapped
# Control: any other issuer keeps the SDK-built parameters untouched, and only the Google
# authorization server itself (not a path or suffix lookalike) gets the parameters.
response, seen = await _run_browser_flow(
tmp_path / "other", monkeypatch, issuer=OTHER_AS, authorization_servers=[OTHER_AS])
assert response.status_code == 200
assert not any(key in ("access_type", "prompt") for key, _ in parse_qsl(urlsplit(seen["authorize_url"]).query))
from tools.mcp_oauth_provider import google_offline_access_params
def ctx(issuer):
return SimpleNamespace(oauth_metadata=SimpleNamespace(issuer=issuer) if issuer is not None else None)
for lookalike in (None, OTHER_AS, "https://evil.example/accounts.google.com", "https://accounts.google.com.evil.example"):
assert google_offline_access_params(ctx(lookalike)) == {}, lookalike
assert google_offline_access_params(ctx(GOOGLE)) == {"access_type": "offline"}
@pytest.mark.parametrize("document_issuer, advertised, accepted", [
pytest.param(GOOGLE, f"{GOOGLE}/", True, id="slash-less-doc-issuer-vs-normalized-advertised"),
pytest.param(f"{GOOGLE}/", f"{GOOGLE}/", True, id="both-slash-forms-still-match"),
pytest.param("https://evil.example", f"{GOOGLE}/", False, id="different-issuer-still-rejected"),
])
@pytest.mark.asyncio
async def test_device_metadata_issuer_normalization(document_issuer, advertised, accepted):
async def test_device_flow_normalizes_issuer_and_asks_google_for_offline_access(capsys):
from mcp.client.auth.exceptions import OAuthFlowError
from tools.mcp_oauth_device import DeviceOAuthMetadata, _device_metadata
from tools.mcp_tool import sdk_httpx
def handler(request):
doc = _asm_doc(document_issuer, device_authorization_endpoint=f"{GOOGLE}/device/code",
grant_types_supported=["authorization_code", "refresh_token",
"urn:ietf:params:oauth:grant-type:device_code"])
return sdk_httpx().Response(200, json=doc, request=request)
httpx = sdk_httpx()
async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client:
if accepted:
metadata = await _device_metadata(client, RESOURCE, advertised)
assert isinstance(metadata, DeviceOAuthMetadata)
else:
with pytest.raises(OAuthFlowError):
await _device_metadata(client, RESOURCE, advertised)
@pytest.mark.parametrize("issuer, expects_offline", [
pytest.param(GOOGLE, True, id="google-gets-offline-access"),
pytest.param(OTHER_AS, False, id="other-issuer-untouched"),
])
@pytest.mark.asyncio
async def test_device_authorization_request_carries_access_type(issuer, expects_offline, capsys):
from mcp.shared.auth import OAuthClientInformationFull, OAuthClientMetadata
from pydantic import AnyUrl
from tools.mcp_oauth import HermesTokenStorage
from tools.mcp_oauth_device import DeviceOAuthMetadata, _authorize
from tools.mcp_oauth_device import DeviceOAuthMetadata, _authorize, _device_metadata
from tools.mcp_oauth_manager import _HERMES_PROVIDER_CLS
from tools.mcp_tool import sdk_httpx
seen = {}
def handler(request):
if urlsplit(str(request.url)).path == "/device/code":
seen["device_form"] = dict(request.url.params) if hasattr(request.url, "params") else {}
# httpx2 form posts: read the body
body = request.content.decode() if request.content else ""
seen["device_form"] = dict(pair.split("=", 1) for pair in body.split("&")) if body else {}
return sdk_httpx().Response(200, json={
"device_code": "DC-1", "user_code": "UC-1", "verification_uri": f"{issuer}/activate",
"interval": 0.05, "expires_in": 60}, request=request)
return sdk_httpx().Response(200, json={"access_token": "AT-1", "token_type": "Bearer",
"expires_in": 3600}, request=request)
httpx = sdk_httpx()
provider = _HERMES_PROVIDER_CLS(
server_name="srv", server_url=RESOURCE, storage=HermesTokenStorage("srv"),
client_metadata=OAuthClientMetadata(redirect_uris=[AnyUrl("http://127.0.0.1:1/cb")], client_name="Hermes Agent"))
provider.context.oauth_metadata = DeviceOAuthMetadata.model_validate(
_asm_doc(issuer, device_authorization_endpoint=f"{issuer}/device/code"))
provider.context.client_info = OAuthClientInformationFull.model_validate(
{"client_id": "cid-1", "redirect_uris": ["http://127.0.0.1:1/cb"], "token_endpoint_auth_method": "none"})
async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client:
tokens = await _authorize(client, provider, {"timeout": 30})
capsys.readouterr()
assert tokens.access_token == "AT-1"
assert ("access_type" in seen["device_form"]) is expects_offline
# Issuer check: the advertised host-only server arrives as ``str(AnyHttpUrl)`` with a trailing
# "/"; both slash forms match Google's slash-less document issuer, a different issuer still fails.
for document_issuer, advertised, accepted in (
(GOOGLE, f"{GOOGLE}/", True), (f"{GOOGLE}/", f"{GOOGLE}/", True), ("https://evil.example", f"{GOOGLE}/", False)):
def handler(request, document_issuer=document_issuer):
doc = _asm_doc(document_issuer, device_authorization_endpoint=f"{GOOGLE}/device/code",
grant_types_supported=["authorization_code", "refresh_token",
"urn:ietf:params:oauth:grant-type:device_code"])
return httpx.Response(200, json=doc, request=request)
def test_google_offline_access_params_matches_only_google():
from tools.mcp_oauth_provider import google_offline_access_params
async with httpx.AsyncClient(transport=httpx.MockTransport(handler)) as client:
if accepted:
assert isinstance(await _device_metadata(client, RESOURCE, advertised), DeviceOAuthMetadata)
else:
with pytest.raises(OAuthFlowError):
await _device_metadata(client, RESOURCE, advertised)
def ctx(issuer):
metadata = SimpleNamespace(issuer=issuer) if issuer is not None else None
return SimpleNamespace(oauth_metadata=metadata)
# Device authorization request: Google's carries access_type=offline, any other issuer's does not.
for issuer, expects_offline in ((GOOGLE, True), (OTHER_AS, False)):
seen = {}
assert google_offline_access_params(ctx(None)) == {}
assert google_offline_access_params(ctx(OTHER_AS)) == {}
# Path- or suffix-lookalikes must not match; only the authorization server itself does.
assert google_offline_access_params(ctx("https://evil.example/accounts.google.com")) == {}
assert google_offline_access_params(ctx("https://accounts.google.com.evil.example")) == {}
assert google_offline_access_params(ctx(GOOGLE)) == {"access_type": "offline"}
def device_handler(request, issuer=issuer):
if urlsplit(str(request.url)).path == "/device/code":
body = request.content.decode() if request.content else ""
seen["device_form"] = dict(pair.split("=", 1) for pair in body.split("&")) if body else {}
return httpx.Response(200, json={
"device_code": "DC-1", "user_code": "UC-1", "verification_uri": f"{issuer}/activate",
"interval": 0.05, "expires_in": 60}, request=request)
return httpx.Response(200, json={"access_token": "AT-1", "token_type": "Bearer",
"expires_in": 3600}, request=request)
provider = _HERMES_PROVIDER_CLS(
server_name="srv", server_url=RESOURCE, storage=HermesTokenStorage("srv"),
client_metadata=OAuthClientMetadata(redirect_uris=[AnyUrl("http://127.0.0.1:1/cb")], client_name="Hermes Agent"))
provider.context.oauth_metadata = DeviceOAuthMetadata.model_validate(
_asm_doc(issuer, device_authorization_endpoint=f"{issuer}/device/code"))
provider.context.client_info = OAuthClientInformationFull.model_validate(
{"client_id": "cid-1", "redirect_uris": ["http://127.0.0.1:1/cb"], "token_endpoint_auth_method": "none"})
async with httpx.AsyncClient(transport=httpx.MockTransport(device_handler)) as client:
tokens = await _authorize(client, provider, {"timeout": 30})
capsys.readouterr()
assert tokens.access_token == "AT-1"
assert ("access_type" in seen["device_form"]) is expects_offline, issuer

View File

@@ -132,20 +132,25 @@ class HermesProviderMixin:
``access_type=offline`` is what makes Google issue one at all, and ``prompt=consent`` is what
makes it re-issue one on repeat logins (the first consent already spent the grant); MCP
discovery advertises neither. The SDK builds the URL itself, so the two parameters are
appended here — never overwriting values already present in the query."""
params = google_offline_access_params(self.context)
appended here — never overwriting values already present in the query. Wraps once: every
authorization runs through here, and the wrapper reads the issuer at call time."""
inner = self.context.redirect_handler
if not params or inner is None:
if inner is None or getattr(inner, "_hermes_offline_access", False):
return
from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit
async def _with_offline_access(authorization_url: str) -> None:
params = google_offline_access_params(self.context)
if not params:
await inner(authorization_url)
return
parts = urlsplit(authorization_url)
query = dict(parse_qsl(parts.query, keep_blank_values=True))
query.update(params)
query.setdefault("prompt", "consent")
await inner(urlunsplit((parts.scheme, parts.netloc, parts.path, urlencode(query), parts.fragment)))
_with_offline_access._hermes_offline_access = True # type: ignore[attr-defined]
self.context.redirect_handler = _with_offline_access
async def _hermes_accept_origin_issued_metadata(self, response):