Close upstream OAuth clients (#4248)

This commit is contained in:
Jeremiah Lowin 2026-05-27 19:45:45 -04:00 committed by GitHub
commit 53b20168c8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 93 additions and 32 deletions

View file

@ -23,6 +23,8 @@ import secrets
import time
from base64 import urlsafe_b64encode
from collections import OrderedDict
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from typing import Any, Literal
from urllib.parse import urlencode, urlparse, urlunparse
@ -697,6 +699,14 @@ class OAuthProxy(OAuthProvider, ConsentMixin):
timeout=HTTP_TIMEOUT_SECONDS,
)
@asynccontextmanager
async def _upstream_oauth_client(self) -> AsyncIterator[AsyncOAuth2Client]:
oauth_client = self._create_upstream_oauth_client()
try:
yield oauth_client
finally:
await oauth_client.aclose()
def _get_refresh_lock(self, token_id: str) -> anyio.Lock:
"""Get or create a per-token refresh lock, evicting LRU entries when at capacity."""
lock = self._refresh_locks.get(token_id)
@ -1389,9 +1399,6 @@ class OAuthProxy(OAuthProvider, ConsentMixin):
logger.error("No upstream refresh token available")
raise TokenError("invalid_grant", "Refresh not supported for this token")
# Refresh upstream token using authlib
oauth_client = self._create_upstream_oauth_client()
# Allow child classes to transform scopes before sending to upstream
# This enables provider-specific scope formatting (e.g., Azure prefixing)
# while keeping original scopes in storage
@ -1399,12 +1406,13 @@ class OAuthProxy(OAuthProvider, ConsentMixin):
try:
logger.debug("Refreshing upstream token (jti=%s)", refresh_jti[:8])
token_response: dict[str, Any] = await oauth_client.refresh_token(
url=self._upstream_token_endpoint,
refresh_token=upstream_token_set.refresh_token,
scope=" ".join(upstream_scopes) if upstream_scopes else None,
**self._extra_token_params,
)
async with self._upstream_oauth_client() as oauth_client:
token_response: dict[str, Any] = await oauth_client.refresh_token(
url=self._upstream_token_endpoint,
refresh_token=upstream_token_set.refresh_token,
scope=" ".join(upstream_scopes) if upstream_scopes else None,
**self._extra_token_params,
)
logger.debug("Successfully refreshed upstream token")
except Exception as e:
logger.error("Upstream token refresh failed: %s", e)
@ -1634,14 +1642,14 @@ class OAuthProxy(OAuthProvider, ConsentMixin):
"""
scopes = upstream_token_set.scope.split() if upstream_token_set.scope else []
upstream_scopes = self._prepare_scopes_for_upstream_refresh(scopes)
oauth_client = self._create_upstream_oauth_client()
token_response: dict[str, Any] = await oauth_client.refresh_token(
url=self._upstream_token_endpoint,
refresh_token=upstream_token_set.refresh_token,
scope=" ".join(upstream_scopes) if upstream_scopes else None,
**self._extra_token_params,
)
async with self._upstream_oauth_client() as oauth_client:
token_response: dict[str, Any] = await oauth_client.refresh_token(
url=self._upstream_token_endpoint,
refresh_token=upstream_token_set.refresh_token,
scope=" ".join(upstream_scopes) if upstream_scopes else None,
**self._extra_token_params,
)
logger.debug(
"Transparent upstream refresh succeeded (token_id=%s)",
upstream_token_set.upstream_token_id[:8],
@ -1894,16 +1902,16 @@ class OAuthProxy(OAuthProvider, ConsentMixin):
# Use the factory method when available (supports alternative auth like
# client assertions for managed identity), falling back to basic auth
# or client_id-only for public clients per RFC 7009
oauth_client = self._create_upstream_oauth_client()
if oauth_client.client_secret is not None:
# Client secret is available, use HTTP Basic auth
request_kwargs["auth"] = (
self._upstream_client_id,
oauth_client.client_secret,
)
else:
# No secret; public client must still identify itself per RFC 7009
revocation_data["client_id"] = self._upstream_client_id
async with self._upstream_oauth_client() as oauth_client:
if oauth_client.client_secret is not None:
# Client secret is available, use HTTP Basic auth
request_kwargs["auth"] = (
self._upstream_client_id,
oauth_client.client_secret,
)
else:
# No secret; public client must still identify itself per RFC 7009
revocation_data["client_id"] = self._upstream_client_id
await http_client.post(
self._upstream_revocation_endpoint,
@ -2133,9 +2141,6 @@ class OAuthProxy(OAuthProvider, ConsentMixin):
transaction = transaction_model.model_dump()
# Exchange IdP code for tokens (server-side)
oauth_client = self._create_upstream_oauth_client()
try:
idp_redirect_uri = (
f"{str(self.base_url).rstrip('/')}{self._redirect_path}"
@ -2176,9 +2181,11 @@ class OAuthProxy(OAuthProvider, ConsentMixin):
list(self._extra_token_params.keys()),
)
idp_tokens: dict[str, Any] = await oauth_client.fetch_token(
**token_params
)
# Exchange IdP code for tokens (server-side)
async with self._upstream_oauth_client() as oauth_client:
idp_tokens: dict[str, Any] = await oauth_client.fetch_token(
**token_params
)
logger.debug(
f"Successfully exchanged IdP code for tokens (transaction: {txn_id}, PKCE: {bool(proxy_code_verifier)})"

View file

@ -24,6 +24,7 @@ from fastmcp.server.auth.oauth_proxy.models import (
DEFAULT_REFRESH_TOKEN_EXPIRY_SECONDS,
ClientCode,
JTIMapping,
OAuthTransaction,
RefreshTokenMetadata,
UpstreamTokenSet,
_hash_token,
@ -195,6 +196,57 @@ class TestOAuthProxyTokenEndpointAuth:
token_endpoint_auth_method="client_secret_post",
timeout=30.0,
)
mock_client.aclose.assert_awaited_once()
async def test_callback_closes_upstream_oauth_client(self, jwt_verifier):
proxy = OAuthProxy(
upstream_authorization_endpoint="https://oauth.example.com/authorize",
upstream_token_endpoint="https://oauth.example.com/token",
upstream_client_id="client-id",
upstream_client_secret="client-secret",
token_verifier=jwt_verifier,
base_url="https://proxy.example.com",
require_authorization_consent=False,
jwt_signing_key="test-secret",
client_storage=MemoryStore(),
)
await proxy._transaction_store.put(
key="txn-id",
value=OAuthTransaction(
txn_id="txn-id",
client_id="test-client",
client_redirect_uri="http://localhost:12345/callback",
client_state="client-state",
code_challenge="",
code_challenge_method="S256",
scopes=["read"],
created_at=time.time(),
),
)
mock_request = Mock()
mock_request.query_params = {"code": "idp-code", "state": "txn-id"}
mock_request.cookies = {}
mock_client = AsyncMock()
mock_client.fetch_token = AsyncMock(
return_value={
"access_token": "upstream-access-token",
"refresh_token": "upstream-refresh-token",
"expires_in": 3600,
"token_type": "Bearer",
}
)
with patch.object(
proxy, "_create_upstream_oauth_client", return_value=mock_client
):
response = await proxy._handle_idp_callback(mock_request)
assert response.status_code == 302
mock_client.fetch_token.assert_awaited_once()
mock_client.aclose.assert_awaited_once()
class TestTokenHandlerErrorTransformation:
@ -614,6 +666,7 @@ class TestFallbackRefreshTokenExpiry:
oauth_client_mock = Mock()
oauth_client_mock.refresh_token = AsyncMock(side_effect=fake_refresh)
oauth_client_mock.aclose = AsyncMock()
with patch.object(
proxy,
"_create_upstream_oauth_client",
@ -1042,6 +1095,7 @@ class TestTransparentUpstreamRefresh:
assert result is not None
assert result.token == "refreshed-upstream-access"
mock_oauth_client.refresh_token.assert_called_once()
mock_oauth_client.aclose.assert_awaited_once()
async def test_transparent_refresh_updates_stored_token(self, proxy):
"""After transparent refresh, the stored upstream token is updated."""