From 62eea7a20bad8a77e9e7fe307309a823efbb14a7 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Fri, 30 May 2025 20:08:42 -0400 Subject: [PATCH] Clean up --- src/fastmcp/client/auth.py | 31 +++++++------- src/fastmcp/client/oauth_callback.py | 40 ++++++++++++++----- src/fastmcp/server/auth/in_memory_provider.py | 3 +- 3 files changed, 47 insertions(+), 27 deletions(-) diff --git a/src/fastmcp/client/auth.py b/src/fastmcp/client/auth.py index 7b589026e..bd63419a2 100644 --- a/src/fastmcp/client/auth.py +++ b/src/fastmcp/client/auth.py @@ -186,7 +186,7 @@ class FileTokenStorage(TokenStorage): async def set_tokens(self, tokens: _MCPOAuthToken) -> None: """Save tokens to file storage.""" # Convert to custom model with expiration datetime - tokens = OAuthToken.model_validate(tokens) + 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)}") @@ -266,7 +266,7 @@ async def discover_oauth_metadata( async def check_if_auth_required( - mcp_endpoint_url: str, httpx_kwargs: dict[str, Any] | None = None + mcp_url: str, httpx_kwargs: dict[str, Any] | None = None ) -> bool: """ Check if the MCP endpoint requires authentication by making a test request. @@ -277,7 +277,7 @@ async def check_if_auth_required( async with httpx.AsyncClient(**(httpx_kwargs or {})) as client: try: # Try a simple request to the endpoint - response = await client.get(mcp_endpoint_url, timeout=5.0) + response = await client.get(mcp_url, timeout=5.0) # If we get 401/403, auth is likely required if response.status_code in (401, 403): @@ -296,7 +296,7 @@ async def check_if_auth_required( def OAuth( - mcp_endpoint_url: str, + mcp_url: str, scopes: str | list[str] | None = None, client_name: str = "FastMCP Client", token_storage_cache_dir: Path | None = None, @@ -309,17 +309,18 @@ def OAuth( httpx.AsyncClient (or appropriate FastMCP client/transport instance) Args: - mcp_endpoint_url: Full URL to the MCP endpoint (e.g., - "http://host/mcp/sse") scopes: OAuth scopes to request. Can be a - space-separated string or a list of strings. client_name: Name for this - client during registration token_storage_cache_dir: Directory for - FileTokenStorage additional_client_metadata: Extra fields for - OAuthClientMetadata + mcp_url: Full URL to the MCP endpoint (e.g., + "http://host/mcp/sse") + scopes: OAuth scopes to request. Can be a + space-separated string or a list of strings. + client_name: Name for this client during registration + token_storage_cache_dir: Directory for FileTokenStorage + additional_client_metadata: Extra fields for OAuthClientMetadata Returns: OAuthClientProvider """ - parsed_url = urlparse(mcp_endpoint_url) + parsed_url = urlparse(mcp_url) server_base_url = f"{parsed_url.scheme}://{parsed_url.netloc}" # Setup OAuth client @@ -347,7 +348,7 @@ def OAuth( # Define OAuth handlers async def redirect_handler(authorization_url: str) -> None: """Open browser for authorization.""" - logger.info(f"Opening browser for OAuth authorization: {authorization_url}") + logger.info(f"OAuth authorization URL: {authorization_url}") webbrowser.open(authorization_url) async def callback_handler() -> tuple[str, str | None]: @@ -363,19 +364,19 @@ def OAuth( ) # Run server until response is received with timeout logic - async with anyio.create_task_group() as tg: tg.start_soon(server.serve) logger.info( f"🎧 OAuth callback server started on http://127.0.0.1:{redirect_port}" ) + TIMEOUT = 300.0 # 5 minute timeout try: - with anyio.fail_after(300.0): # 5 minute timeout + with anyio.fail_after(TIMEOUT): auth_code, state = await response_future return auth_code, state except TimeoutError: - raise TimeoutError("OAuth callback timed out after 300 seconds") + raise TimeoutError(f"OAuth callback timed out after {TIMEOUT} seconds") finally: server.should_exit = True await asyncio.sleep(0.1) # Allow server to shutdown gracefully diff --git a/src/fastmcp/client/oauth_callback.py b/src/fastmcp/client/oauth_callback.py index 8376050a2..f9cecd16b 100644 --- a/src/fastmcp/client/oauth_callback.py +++ b/src/fastmcp/client/oauth_callback.py @@ -9,8 +9,10 @@ from __future__ import annotations import asyncio import socket +from dataclasses import dataclass from starlette.applications import Starlette +from starlette.requests import Request from starlette.responses import HTMLResponse from starlette.routing import Route from uvicorn import Config, Server @@ -184,6 +186,21 @@ def find_available_port() -> int: return s.getsockname()[1] +@dataclass +class CallbackResponse: + code: str | None = None + state: str | None = None + error: str | None = None + error_description: str | None = None + + @classmethod + def from_dict(cls, data: dict[str, str]) -> CallbackResponse: + return cls(**{k: v for k, v in data.items() if k in cls.__annotations__}) + + def to_dict(self) -> dict[str, str]: + return {k: v for k, v in self.__dict__.items() if v is not None} + + def create_oauth_callback_server( port: int, callback_path: str = "/callback", @@ -203,30 +220,31 @@ def create_oauth_callback_server( Configured uvicorn Server instance (not yet running) """ - async def callback_handler(request): + async def callback_handler(request: Request): """Handle OAuth callback requests with proper HTML responses.""" query_params = dict(request.query_params) - auth_code = query_params.get("code") - state = query_params.get("state") - error = query_params.get("error") + callback_response = CallbackResponse.from_dict(query_params) - if error: - error_desc = query_params.get("error_description", "Unknown error") + if callback_response.error: + error_desc = callback_response.error_description or "Unknown error" # Resolve future with exception if provided if response_future and not response_future.done(): response_future.set_exception( - RuntimeError(f"OAuth error: {error} - {error_desc}") + RuntimeError( + f"OAuth error: {callback_response.error} - {error_desc}" + ) ) return HTMLResponse( create_callback_html( - f"FastMCP OAuth Error: {error}
{error_desc}", is_success=False + f"FastMCP OAuth Error: {callback_response.error}
{error_desc}", + is_success=False, ), status_code=400, ) - if not auth_code: + if not callback_response.code: # Resolve future with exception if provided if response_future and not response_future.done(): response_future.set_exception( @@ -243,7 +261,9 @@ def create_oauth_callback_server( # Success case if response_future and not response_future.done(): - response_future.set_result((auth_code, state)) + response_future.set_result( + (callback_response.code, callback_response.state) + ) return HTMLResponse( create_callback_html("FastMCP OAuth login complete!", server_url=server_url) diff --git a/src/fastmcp/server/auth/in_memory_provider.py b/src/fastmcp/server/auth/in_memory_provider.py index 5887305c6..59ac0d2ad 100644 --- a/src/fastmcp/server/auth/in_memory_provider.py +++ b/src/fastmcp/server/auth/in_memory_provider.py @@ -25,8 +25,7 @@ from fastmcp.server.auth.auth import ( # Default expiration times (in seconds) DEFAULT_AUTH_CODE_EXPIRY_SECONDS = 5 * 60 # 5 minutes DEFAULT_ACCESS_TOKEN_EXPIRY_SECONDS = 60 * 60 # 1 hour -# Refresh tokens often have longer or no expiry; let's make them non-expiring for simplicity -DEFAULT_REFRESH_TOKEN_EXPIRY_SECONDS = None +DEFAULT_REFRESH_TOKEN_EXPIRY_SECONDS = None # No expiry class InMemoryOAuthProvider(OAuthProvider):