diff --git a/src/fastmcp/client/auth.py b/src/fastmcp/client/auth.py index bd63419a2..6165fdfd8 100644 --- a/src/fastmcp/client/auth.py +++ b/src/fastmcp/client/auth.py @@ -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