From d9541c9c2181fa4ad1e07e7c90c4db2f963e7924 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Wed, 25 Feb 2026 12:44:24 -0500 Subject: [PATCH] Lazy-load heavy imports to reduce import time Defer auth providers (JWTVerifier, OAuthProxy, OIDCProxy) and Client to avoid eagerly importing authlib, cryptography, key_value.aio, and beartype on every `from fastmcp import FastMCP`. --- scripts/benchmark_imports.py | 212 ++++++++++++++++++++++++++++ src/fastmcp/__init__.py | 24 +++- src/fastmcp/server/__init__.py | 1 - src/fastmcp/server/auth/__init__.py | 44 +++++- src/fastmcp/server/server.py | 3 +- 5 files changed, 275 insertions(+), 9 deletions(-) create mode 100644 scripts/benchmark_imports.py diff --git a/scripts/benchmark_imports.py b/scripts/benchmark_imports.py new file mode 100644 index 000000000..6ad1dfe57 --- /dev/null +++ b/scripts/benchmark_imports.py @@ -0,0 +1,212 @@ +#!/usr/bin/env python +"""Benchmark import times for fastmcp and its dependency chain. + +Each measurement runs in a fresh subprocess so there's no shared module cache. +Incremental costs are measured by pre-importing dependencies, so we can see +what each module truly adds. + +Usage: + uv run python scripts/benchmark_imports.py + uv run python scripts/benchmark_imports.py --runs 10 + uv run python scripts/benchmark_imports.py --json +""" + +from __future__ import annotations + +import argparse +import json +import subprocess +import sys +from dataclasses import dataclass + + +@dataclass +class BenchmarkCase: + label: str + stmt: str + prereqs: str = "" + group: str = "" + + +CASES = [ + # --- Floor --- + BenchmarkCase("pydantic", "import pydantic", group="floor"), + BenchmarkCase("mcp", "import mcp", group="floor"), + BenchmarkCase( + "mcp (server only)", "import mcp.server.lowlevel.server", group="floor" + ), + # --- Auth stack (incremental over mcp) --- + BenchmarkCase( + "authlib.jose", "import authlib.jose", prereqs="import mcp", group="auth" + ), + BenchmarkCase( + "cryptography.fernet", + "from cryptography.fernet import Fernet", + prereqs="import mcp", + group="auth", + ), + BenchmarkCase( + "authlib.integrations.httpx_client", + "from authlib.integrations.httpx_client import AsyncOAuth2Client", + prereqs="import mcp", + group="auth", + ), + BenchmarkCase( + "key_value.aio", "import key_value.aio", prereqs="import mcp", group="auth" + ), + BenchmarkCase( + "key_value.aio.stores.filetree", + "from key_value.aio.stores.filetree import FileTreeStore", + prereqs="import mcp", + group="auth", + ), + BenchmarkCase("beartype", "import beartype", prereqs="import mcp", group="auth"), + # --- Docket stack (incremental over mcp) --- + BenchmarkCase("redis", "import redis", prereqs="import mcp", group="docket"), + BenchmarkCase( + "opentelemetry.sdk.metrics", + "import opentelemetry.sdk.metrics", + prereqs="import mcp", + group="docket", + ), + BenchmarkCase("docket", "import docket", prereqs="import mcp", group="docket"), + BenchmarkCase("croniter", "import croniter", prereqs="import mcp", group="docket"), + # --- Other deps (incremental over mcp) --- + BenchmarkCase("httpx", "import httpx", prereqs="import mcp", group="other"), + BenchmarkCase( + "starlette", + "from starlette.applications import Starlette", + prereqs="import mcp", + group="other", + ), + BenchmarkCase( + "pydantic_settings", + "import pydantic_settings", + prereqs="import mcp", + group="other", + ), + BenchmarkCase( + "rich.console", "import rich.console", prereqs="import mcp", group="other" + ), + BenchmarkCase("jsonref", "import jsonref", prereqs="import mcp", group="other"), + BenchmarkCase("requests", "import requests", prereqs="import mcp", group="other"), + # --- FastMCP (total and incremental) --- + BenchmarkCase("fastmcp (total)", "from fastmcp import FastMCP", group="fastmcp"), + BenchmarkCase( + "fastmcp (over mcp)", + "from fastmcp import FastMCP", + prereqs="import mcp", + group="fastmcp", + ), + BenchmarkCase( + "fastmcp (over mcp+docket)", + "from fastmcp import FastMCP", + prereqs="import mcp; import docket", + group="fastmcp", + ), + BenchmarkCase( + "fastmcp (over mcp+docket+auth deps)", + "from fastmcp import FastMCP", + prereqs=( + "import mcp; import docket; import authlib.jose;" + " from cryptography.fernet import Fernet;" + " import key_value.aio" + ), + group="fastmcp", + ), +] + + +def measure_once(stmt: str, prereqs: str) -> float | None: + pre = prereqs + "; " if prereqs else "" + code = ( + f"{pre}" + "import time as _t; _s=_t.perf_counter(); " + f"{stmt}; " + "print(f'{(_t.perf_counter()-_s)*1000:.2f}')" + ) + r = subprocess.run( + [sys.executable, "-c", code], + capture_output=True, + text=True, + timeout=30, + ) + if r.returncode == 0 and r.stdout.strip(): + return float(r.stdout.strip()) + return None + + +def measure(case: BenchmarkCase, runs: int) -> dict[str, float | str | None]: + times: list[float] = [] + for _ in range(runs): + t = measure_once(case.stmt, case.prereqs) + if t is not None: + times.append(t) + + if not times: + return {"label": case.label, "group": case.group, "median_ms": None} + + times.sort() + median = times[len(times) // 2] + return { + "label": case.label, + "group": case.group, + "median_ms": round(median, 1), + "min_ms": round(times[0], 1), + "max_ms": round(times[-1], 1), + "runs": len(times), + } + + +def print_table(results: list[dict[str, float | str | None]]) -> None: + current_group = None + print(f"\n{'Module':<45} {'Median':>8} {'Min':>8} {'Max':>8}") + print("-" * 71) + for r in results: + if r["group"] != current_group: + current_group = r["group"] + group_labels = { + "floor": "--- Unavoidable floor ---", + "auth": "--- Auth stack (incremental over mcp) ---", + "docket": "--- Docket stack (incremental over mcp) ---", + "other": "--- Other deps (incremental over mcp) ---", + "fastmcp": "--- FastMCP totals ---", + } + print(f"\n{group_labels.get(current_group, current_group)}") + if r["median_ms"] is not None: + print( + f" {r['label']:<43} {r['median_ms']:>7.1f}ms" + f" {r['min_ms']:>7.1f}ms {r['max_ms']:>7.1f}ms" + ) + else: + print(f" {r['label']:<43} error") + + +def main() -> None: + parser = argparse.ArgumentParser(description="Benchmark fastmcp import times") + parser.add_argument( + "--runs", type=int, default=5, help="Number of runs per measurement (default 5)" + ) + parser.add_argument("--json", action="store_true", help="Output results as JSON") + args = parser.parse_args() + + print(f"Benchmarking import times ({args.runs} runs each)...") + print(f"Python: {sys.version.split()[0]}") + print(f"Executable: {sys.executable}") + + results = [] + for case in CASES: + r = measure(case, args.runs) + results.append(r) + if not args.json: + ms = f"{r['median_ms']:.1f}ms" if r["median_ms"] is not None else "error" + print(f" {case.label}: {ms}") + + if args.json: + print(json.dumps(results, indent=2)) + else: + print_table(results) + + +if __name__ == "__main__": + main() diff --git a/src/fastmcp/__init__.py b/src/fastmcp/__init__.py index 14c0abc5b..a524b402c 100644 --- a/src/fastmcp/__init__.py +++ b/src/fastmcp/__init__.py @@ -1,10 +1,16 @@ """FastMCP - An ergonomic MCP interface.""" +import importlib import warnings from importlib.metadata import version as _version +from typing import TYPE_CHECKING + from fastmcp.settings import Settings from fastmcp.utilities.logging import configure_logging as _configure_logging +if TYPE_CHECKING: + from fastmcp.client import Client as Client + settings = Settings() if settings.log_enabled: _configure_logging( @@ -16,9 +22,6 @@ from fastmcp.server.server import FastMCP from fastmcp.server.context import Context import fastmcp.server -from fastmcp.client import Client -from . import client - __version__ = _version("fastmcp") @@ -27,6 +30,21 @@ if settings.deprecation_warnings: warnings.simplefilter("default", DeprecationWarning) +# --- Lazy imports for performance (see #3292) --- +# Client and the client submodule are deferred so that server-only users +# don't pay for the client import chain. Do not convert back to top-level. + + +def __getattr__(name: str) -> object: + if name == "Client": + from fastmcp.client import Client + + return Client + if name == "client": + return importlib.import_module("fastmcp.client") + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + + __all__ = [ "Client", "Context", diff --git a/src/fastmcp/server/__init__.py b/src/fastmcp/server/__init__.py index fb9afd895..101adfcdb 100644 --- a/src/fastmcp/server/__init__.py +++ b/src/fastmcp/server/__init__.py @@ -1,6 +1,5 @@ from .context import Context from .server import FastMCP, create_proxy -from . import dependencies __all__ = ["Context", "FastMCP", "create_proxy"] diff --git a/src/fastmcp/server/auth/__init__.py b/src/fastmcp/server/auth/__init__.py index 94e23dca6..2c11d32b7 100644 --- a/src/fastmcp/server/auth/__init__.py +++ b/src/fastmcp/server/auth/__init__.py @@ -1,3 +1,5 @@ +from typing import TYPE_CHECKING + from .auth import ( OAuthProvider, TokenVerifier, @@ -12,10 +14,44 @@ from .authorization import ( restrict_tag, run_auth_checks, ) -from .providers.debug import DebugTokenVerifier -from .providers.jwt import JWTVerifier, StaticTokenVerifier -from .oauth_proxy import OAuthProxy -from .oidc_proxy import OIDCProxy + +if TYPE_CHECKING: + from .oauth_proxy import OAuthProxy as OAuthProxy + from .oidc_proxy import OIDCProxy as OIDCProxy + from .providers.debug import DebugTokenVerifier as DebugTokenVerifier + from .providers.jwt import JWTVerifier as JWTVerifier + from .providers.jwt import StaticTokenVerifier as StaticTokenVerifier + + +# --- Lazy imports for performance (see #3292) --- +# These providers pull in heavy deps (authlib, cryptography, key_value.aio, +# beartype) that most users never need. Keeping them behind __getattr__ +# avoids ~150ms+ of import overhead for the common server-only case. +# Do not convert these back to top-level imports. + + +def __getattr__(name: str) -> object: + if name == "DebugTokenVerifier": + from .providers.debug import DebugTokenVerifier + + return DebugTokenVerifier + if name == "JWTVerifier": + from .providers.jwt import JWTVerifier + + return JWTVerifier + if name == "StaticTokenVerifier": + from .providers.jwt import StaticTokenVerifier + + return StaticTokenVerifier + if name == "OAuthProxy": + from .oauth_proxy import OAuthProxy + + return OAuthProxy + if name == "OIDCProxy": + from .oidc_proxy import OIDCProxy + + return OIDCProxy + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") __all__ = [ diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index db558dca3..2bdecf803 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -62,7 +62,6 @@ from fastmcp.server.apps import ( resolve_ui_mime_type, ) from fastmcp.server.auth import AuthCheck, AuthContext, AuthProvider, run_auth_checks -from fastmcp.server.dependencies import get_access_token from fastmcp.server.lifespan import Lifespan from fastmcp.server.low_level import LowLevelServer from fastmcp.server.middleware import Middleware, MiddlewareContext @@ -162,6 +161,8 @@ def _get_auth_context() -> tuple[bool, Any]: is_stdio = _current_transport.get() == "stdio" if is_stdio: return (True, None) + from fastmcp.server.dependencies import get_access_token + return (False, get_access_token())