From 2a00b06212adfee1a0c55713c7309c0ab6ff4b34 Mon Sep 17 00:00:00 2001 From: teknium1 <127238744+teknium1@users.noreply.github.com> Date: Sun, 20 Sep 2026 15:19:40 -0700 Subject: [PATCH] fix(mcp): wrap the redirect handler for Google offline access once and fold the tests to two MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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). --- tests/tools/test_mcp_google_offline_access.py | 160 ++++++++---------- tools/mcp_oauth_provider.py | 11 +- 2 files changed, 76 insertions(+), 95 deletions(-) diff --git a/tests/tools/test_mcp_google_offline_access.py b/tests/tools/test_mcp_google_offline_access.py index b9e2ddcf69..9ffb39d755 100644 --- a/tests/tools/test_mcp_google_offline_access.py +++ b/tests/tools/test_mcp_google_offline_access.py @@ -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 diff --git a/tools/mcp_oauth_provider.py b/tools/mcp_oauth_provider.py index 60d88d1e95..2d8e230154 100644 --- a/tools/mcp_oauth_provider.py +++ b/tools/mcp_oauth_provider.py @@ -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):