Remove orphaned OAuth proxy code (#1722)

This commit is contained in:
Jeremiah Lowin 2025-09-02 16:30:50 -04:00 committed by GitHub
commit 694f0fe149
No known key found for this signature in database
GPG key ID: B5690EEEBB952194

View file

@ -39,7 +39,7 @@ from mcp.server.auth.settings import (
from mcp.shared.auth import OAuthClientInformationFull, OAuthToken
from pydantic import AnyHttpUrl, AnyUrl, SecretStr
from starlette.requests import Request
from starlette.responses import JSONResponse, RedirectResponse
from starlette.responses import RedirectResponse
from starlette.routing import Route
from fastmcp.server.auth.auth import OAuthProvider, TokenVerifier
@ -730,155 +730,6 @@ class OAuthProxy(OAuthProvider):
logger.debug("Token revoked successfully")
# -------------------------------------------------------------------------
# Custom Route Handling
# -------------------------------------------------------------------------
async def _handle_proxy_token_request(self, request: Request) -> JSONResponse:
"""Custom token endpoint using authlib for upstream requests.
This handler uses authlib's OAuth2Client to forward token requests to the
upstream OAuth server, automatically handling response format differences.
"""
try:
# Parse the incoming request form data
form_data = await request.form()
# Log the incoming request (with sensitive data redacted)
redacted_form = {
k: (
str(v)[:8] + "..."
if k in {"code", "code_verifier", "client_secret", "refresh_token"}
and v
else str(v)
)
for k, v in form_data.items()
}
logger.debug("Proxy token request form data: %s", redacted_form)
# Create authlib OAuth2 client
oauth_client = AsyncOAuth2Client(
client_id=self._upstream_client_id,
client_secret=self._upstream_client_secret.get_secret_value(),
timeout=HTTP_TIMEOUT_SECONDS,
)
grant_type = str(form_data.get("grant_type", ""))
if grant_type == "authorization_code":
# Authorization code grant
try:
token_data: dict[str, Any] = await oauth_client.fetch_token( # type: ignore[misc]
url=self._upstream_token_endpoint,
code=str(form_data.get("code", "")),
redirect_uri=str(form_data.get("redirect_uri", "")),
code_verifier=str(form_data.get("code_verifier"))
if "code_verifier" in form_data
else None,
)
# Store tokens locally for tracking
if "access_token" in token_data:
self._store_tokens_from_response(token_data)
logger.debug(
"Successfully proxied authorization code exchange via authlib"
)
except Exception as e:
logger.error("Authlib authorization code exchange failed: %s", e)
return JSONResponse(
content={
"error": "invalid_grant",
"error_description": f"Authorization code exchange failed: {e}",
},
status_code=400,
)
elif grant_type == "refresh_token":
# Refresh token grant
try:
token_data: dict[str, Any] = await oauth_client.refresh_token( # type: ignore[misc]
url=self._upstream_token_endpoint,
refresh_token=str(form_data.get("refresh_token", "")),
scope=str(form_data.get("scope"))
if "scope" in form_data
else None,
)
logger.debug(
"Successfully proxied refresh token exchange via authlib"
)
except Exception as e:
logger.error("Authlib refresh token exchange failed: %s", e)
return JSONResponse(
content={
"error": "invalid_grant",
"error_description": f"Refresh token exchange failed: {e}",
},
status_code=400,
)
else:
# Unsupported grant type
logger.error("Unsupported grant type: %s", grant_type)
return JSONResponse(
content={
"error": "unsupported_grant_type",
"error_description": f"Grant type '{grant_type}' not supported by proxy",
},
status_code=400,
)
return JSONResponse(content=token_data)
except Exception as e:
logger.error("Error in proxy token handler: %s", e, exc_info=True)
return JSONResponse(
content={
"error": "server_error",
"error_description": "Internal server error",
},
status_code=500,
)
def _store_tokens_from_response(self, token_data: dict[str, Any]) -> None:
"""Store tokens from upstream response for local tracking."""
try:
access_token_value = token_data.get("access_token")
refresh_token_value = token_data.get("refresh_token")
expires_in = int(
token_data.get("expires_in", DEFAULT_ACCESS_TOKEN_EXPIRY_SECONDS)
)
expires_at = int(time.time() + expires_in)
if access_token_value:
access_token = AccessToken(
token=access_token_value,
client_id=self._upstream_client_id,
scopes=[], # Will be determined by token validation
expires_at=expires_at,
)
self._access_tokens[access_token_value] = access_token
if refresh_token_value:
refresh_token = RefreshToken(
token=refresh_token_value,
client_id=self._upstream_client_id,
scopes=[],
expires_at=None,
)
self._refresh_tokens[refresh_token_value] = refresh_token
# Maintain token relationships
self._access_to_refresh[access_token_value] = refresh_token_value
self._refresh_to_access[refresh_token_value] = access_token_value
logger.debug("Stored tokens from upstream response for tracking")
except Exception as e:
logger.warning("Failed to store tokens from upstream response: %s", e)
def get_routes(
self,
mcp_path: str | None = None,