Fix issue when client info is written but flow is incomplete

This commit is contained in:
Jeremiah Lowin 2025-05-30 21:38:17 -04:00
commit b8e1e8114b

View file

@ -1,7 +1,6 @@
from __future__ import annotations
import asyncio
import datetime
import json
import webbrowser
from pathlib import Path
@ -20,10 +19,9 @@ from mcp.shared.auth import (
OAuthMetadata as _MCPServerOAuthMetadata,
)
from mcp.shared.auth import (
OAuthToken as _MCPOAuthToken,
OAuthToken as OAuthToken,
)
from pydantic import AnyHttpUrl, ValidationError, model_validator
from typing_extensions import Self
from pydantic import AnyHttpUrl, ValidationError
from fastmcp.client.oauth_callback import (
create_oauth_callback_server,
@ -37,19 +35,8 @@ __all__ = ["OAuth"]
logger = get_logger(__name__)
class OAuthToken(_MCPOAuthToken):
"""
OAuth token that stores expiration as a datetime object
"""
expires_at: datetime.datetime | None = None
@model_validator(mode="after")
def set_expires_at(self) -> Self:
if self.expires_in is not None and self.expires_at is None:
now = datetime.datetime.now(datetime.timezone.utc)
self.expires_at = now + datetime.timedelta(seconds=self.expires_in)
return self
def default_cache_dir() -> Path:
return fastmcp_global_settings.home / "oauth-mcp-client-cache"
# Flexible OAuth models for real-world compatibility
@ -140,9 +127,7 @@ class FileTokenStorage(TokenStorage):
def __init__(self, server_url: str, cache_dir: Path | None = None):
"""Initialize storage for a specific server URL."""
self.server_url = server_url
self.cache_dir = (
cache_dir or fastmcp_global_settings.home / "oauth-mcp-client-cache"
)
self.cache_dir = cache_dir or default_cache_dir()
self.cache_dir.mkdir(exist_ok=True, parents=True)
@staticmethod
@ -172,10 +157,10 @@ class FileTokenStorage(TokenStorage):
try:
tokens = OAuthToken.model_validate_json(path.read_text())
now = datetime.datetime.now(datetime.timezone.utc)
if tokens.expires_at is not None and tokens.expires_at <= now:
logger.debug(f"Token expired for {self.get_base_url(self.server_url)}")
return None
# now = datetime.datetime.now(datetime.timezone.utc)
# if tokens.expires_at is not None and tokens.expires_at <= now:
# logger.debug(f"Token expired for {self.get_base_url(self.server_url)}")
# return None
return tokens
except (FileNotFoundError, json.JSONDecodeError, ValidationError) as e:
logger.debug(
@ -183,10 +168,8 @@ class FileTokenStorage(TokenStorage):
)
return None
async def set_tokens(self, tokens: _MCPOAuthToken) -> None:
async def set_tokens(self, tokens: OAuthToken) -> None:
"""Save tokens to file storage."""
# Convert to custom model with expiration datetime
tokens = OAuthToken.model_validate(tokens.model_dump())
path = self._get_file_path("tokens")
path.write_text(tokens.model_dump_json(indent=2))
logger.debug(f"Saved tokens for {self.get_base_url(self.server_url)}")
@ -195,7 +178,24 @@ class FileTokenStorage(TokenStorage):
"""Load client information from file storage."""
path = self._get_file_path("client_info")
try:
return OAuthClientInformationFull.model_validate_json(path.read_text())
client_info = OAuthClientInformationFull.model_validate_json(
path.read_text()
)
# Check if we have corresponding valid tokens
# If no tokens exist, the OAuth flow was incomplete and we should
# force a fresh client registration
tokens = await self.get_tokens()
if tokens is None:
logger.debug(
f"No tokens found for client info at {self.get_base_url(self.server_url)}. "
"OAuth flow may have been incomplete. Clearing client info to force fresh registration."
)
# Clear the incomplete client info
client_info_path = self._get_file_path("client_info")
client_info_path.unlink(missing_ok=True)
return None
return client_info
except (FileNotFoundError, json.JSONDecodeError, ValidationError) as e:
logger.debug(
f"Could not load client info for {self.get_base_url(self.server_url)}: {e}"
@ -208,7 +208,7 @@ class FileTokenStorage(TokenStorage):
path.write_text(client_info.model_dump_json(indent=2))
logger.debug(f"Saved client info for {self.get_base_url(self.server_url)}")
def clear_cache(self) -> None:
def clear(self) -> None:
"""Clear all cached data for this server."""
file_types: list[Literal["client_info", "tokens"]] = ["client_info", "tokens"]
for file_type in file_types:
@ -219,7 +219,7 @@ class FileTokenStorage(TokenStorage):
@classmethod
def clear_all(cls, cache_dir: Path | None = None) -> None:
"""Clear all cached data for all servers."""
cache_dir = cache_dir or fastmcp_global_settings.home / "oauth-mcp-client-cache"
cache_dir = cache_dir or default_cache_dir()
if not cache_dir.exists():
return