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:
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user