Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 35 additions & 3 deletions src/mcp/client/auth/oauth2.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@

import base64
import hashlib
import json
import logging
import secrets
import string
Expand Down Expand Up @@ -81,6 +82,15 @@
)


def _refresh_error_code(body: bytes) -> str | None:
"""Extract the RFC 6749 ``error`` code from a failed token response, if any."""
try:
error = json.loads(body).get("error")
except (json.JSONDecodeError, UnicodeDecodeError, AttributeError):
return None
return error if isinstance(error, str) else None


def check_registration_usable(client_info: OAuthClientInformationFull) -> None:
"""Confirm a registration this flow completed is one it can act on.

Expand Down Expand Up @@ -486,8 +496,15 @@ async def _handle_token_response(self, response: httpx2.Response) -> None:
self.context.update_token_expiry(token_response)
await self.context.storage.set_tokens(token_response)

async def _refresh_token(self) -> httpx2.Request:
"""Build token refresh request."""
async def _refresh_token(self, *, include_resource: bool = True) -> httpx2.Request:
"""Build token refresh request.

Args:
include_resource: Whether to attach the RFC 8707 ``resource`` parameter
(when the protocol version calls for it). The retry path passes False
for authorization servers that reject the parameter on
``refresh_token`` grants (e.g. Microsoft Entra ID v2.0, AADSTS9010010).
"""
if not self.context.current_tokens or not self.context.current_tokens.refresh_token:
raise OAuthTokenError("No refresh token available") # pragma: no cover

Expand All @@ -507,7 +524,7 @@ async def _refresh_token(self) -> httpx2.Request:
}

# Only include resource param if conditions are met
if self.context.should_include_resource_param(self.context.protocol_version):
if include_resource and self.context.should_include_resource_param(self.context.protocol_version):
refresh_data["resource"] = self.context.get_resource_url() # RFC 8707

# Prepare authentication based on preferred method
Expand Down Expand Up @@ -588,9 +605,24 @@ async def async_auth_flow(self, request: httpx2.Request) -> AsyncGenerator[httpx

if not self.context.is_token_valid() and self.context.can_refresh_token():
# Try to refresh token
resource_included = self.context.should_include_resource_param(self.context.protocol_version)
refresh_request = await self._refresh_token()
refresh_response = yield refresh_request

if (
refresh_response.status_code == 400
and resource_included
and _refresh_error_code(await refresh_response.aread()) != "invalid_grant"
):
# Some authorization servers (e.g. Microsoft Entra ID v2.0,
# AADSTS9010010) reject the RFC 8707 resource parameter on
# refresh_token grants. Retry once without it before giving
# up and forcing a full interactive re-authentication.
# `invalid_grant` is excluded: it means the refresh token
# itself is no longer valid, so a retry cannot succeed.
refresh_request = await self._refresh_token(include_resource=False)
refresh_response = yield refresh_request

if not await self._handle_refresh_response(refresh_response):
# Refresh failed, need full re-authentication
self._initialized = False
Expand Down
118 changes: 118 additions & 0 deletions tests/client/test_auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -824,6 +824,124 @@ async def test_resource_param_included_with_protected_resource_metadata(self, oa
assert "resource=" in content


class TestRefreshResourceParamFallback:
"""Refresh keeps the RFC 8707 resource param per the MCP spec, but retries once
without it when the authorization server rejects the refresh with a 400 —
some servers (e.g. Microsoft Entra ID v2.0, AADSTS9010010) reject the
parameter on refresh_token grants (#2578)."""

def _prepare_expired_session(self, oauth_provider: OAuthClientProvider) -> httpx2.Request:
oauth_provider._initialized = True
oauth_provider.context.client_info = OAuthClientInformationFull(
client_id="test_client",
client_secret="test_secret",
redirect_uris=[AnyUrl("http://localhost:3030/callback")],
)
oauth_provider.context.current_tokens = OAuthToken(
access_token="expired_access",
token_type="Bearer",
refresh_token="test_refresh_token",
)
oauth_provider.context.token_expiry_time = time.time() - 3600
return httpx2.Request("GET", "https://api.example.com/mcp", headers={"mcp-protocol-version": "2025-06-18"})

@pytest.mark.anyio
async def test_refresh_retries_without_resource_on_400(self, oauth_provider: OAuthClientProvider):
test_request = self._prepare_expired_session(oauth_provider)
auth_flow = oauth_provider.async_auth_flow(test_request)

first_refresh = await auth_flow.__anext__()
first_body = first_refresh.content.decode()
assert "grant_type=refresh_token" in first_body
assert "resource=" in first_body

entra_rejection = httpx2.Response(
400,
content=b'{"error": "invalid_request", "error_description": "AADSTS9010010: ..."}',
request=first_refresh,
)
retry_refresh = await auth_flow.asend(entra_rejection)
retry_body = retry_refresh.content.decode()
assert "grant_type=refresh_token" in retry_body
assert "resource=" not in retry_body

token_response = httpx2.Response(
200,
content=(
b'{"access_token": "new_access_token", "token_type": "Bearer", '
b'"expires_in": 3600, "refresh_token": "new_refresh_token"}'
),
request=retry_refresh,
)
original_request = await auth_flow.asend(token_response)
assert original_request.headers["Authorization"] == "Bearer new_access_token"

with pytest.raises(StopAsyncIteration):
await auth_flow.asend(httpx2.Response(200, request=original_request))

@pytest.mark.anyio
async def test_refresh_falls_back_to_reauth_when_retry_fails(self, oauth_provider: OAuthClientProvider):
test_request = self._prepare_expired_session(oauth_provider)
auth_flow = oauth_provider.async_auth_flow(test_request)

first_refresh = await auth_flow.__anext__()
entra_rejection = httpx2.Response(
400,
content=b'{"error": "invalid_request", "error_description": "AADSTS9010010: ..."}',
request=first_refresh,
)
retry_refresh = await auth_flow.asend(entra_rejection)
assert "resource=" not in retry_refresh.content.decode()

original_request = await auth_flow.asend(httpx2.Response(400, request=retry_refresh))
# Both refresh attempts failed: the original request goes out unauthenticated
# and the provider is flagged for full re-authentication.
assert "Authorization" not in original_request.headers
assert oauth_provider._initialized is False

with pytest.raises(StopAsyncIteration):
await auth_flow.asend(httpx2.Response(200, request=original_request))

@pytest.mark.anyio
async def test_no_retry_on_invalid_grant(self, oauth_provider: OAuthClientProvider):
test_request = self._prepare_expired_session(oauth_provider)
auth_flow = oauth_provider.async_auth_flow(test_request)

first_refresh = await auth_flow.__anext__()
assert "resource=" in first_refresh.content.decode()

# invalid_grant means the refresh token itself is dead: retrying
# without the resource param cannot help, so the flow goes straight
# to full re-authentication.
dead_grant = httpx2.Response(400, content=b'{"error": "invalid_grant"}', request=first_refresh)
original_request = await auth_flow.asend(dead_grant)
assert str(original_request.url) == "https://api.example.com/mcp"
assert "Authorization" not in original_request.headers
assert oauth_provider._initialized is False

with pytest.raises(StopAsyncIteration):
await auth_flow.asend(httpx2.Response(200, request=original_request))

@pytest.mark.anyio
async def test_no_retry_when_resource_was_not_sent(self, oauth_provider: OAuthClientProvider):
test_request = self._prepare_expired_session(oauth_provider)
test_request.headers["mcp-protocol-version"] = "2025-03-26"
auth_flow = oauth_provider.async_auth_flow(test_request)

first_refresh = await auth_flow.__anext__()
assert "resource=" not in first_refresh.content.decode()

# A 400 without the resource param present is a real failure: no retry,
# the next request is the original one, unauthenticated.
original_request = await auth_flow.asend(httpx2.Response(400, request=first_refresh))
assert str(original_request.url) == "https://api.example.com/mcp"
assert "Authorization" not in original_request.headers
assert oauth_provider._initialized is False

with pytest.raises(StopAsyncIteration):
await auth_flow.asend(httpx2.Response(200, request=original_request))


@pytest.mark.parametrize(
("protocol_version", "expected"),
[
Expand Down
Loading