From 85f32959b1bf7f3a348b7e9f642cc9a6a49363ca Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Fri, 16 Jan 2026 16:48:16 -0500 Subject: [PATCH 1/8] Refactor FastMCP to inherit from Provider FastMCP now properly inherits from Provider, eliminating ~200 lines of duplicated _source_* methods. Key changes: - get_tool/resource/prompt return None instead of raising NotFoundError - Visibility filter separated from transforms (applied last) - Nested server middleware runs on both list and execution operations - Resource auth failure doesn't fall back to templates - AggregateProvider kept as user-facing utility class --- .../server/middleware/authorization.py | 28 +- src/fastmcp/server/providers/__init__.py | 2 + src/fastmcp/server/providers/aggregate.py | 31 +- src/fastmcp/server/providers/base.py | 29 +- .../server/providers/fastmcp_provider.py | 22 +- src/fastmcp/server/server.py | 752 +++++++++--------- src/fastmcp/utilities/inspect.py | 10 +- tests/server/auth/test_authorization.py | 31 +- .../providers/test_local_provider_prompts.py | 29 +- tests/server/test_mount.py | 58 +- tests/server/test_tool_transformation.py | 2 +- 11 files changed, 535 insertions(+), 459 deletions(-) diff --git a/src/fastmcp/server/middleware/authorization.py b/src/fastmcp/server/middleware/authorization.py index f9f69ecd1..468284458 100644 --- a/src/fastmcp/server/middleware/authorization.py +++ b/src/fastmcp/server/middleware/authorization.py @@ -28,7 +28,7 @@ from collections.abc import Sequence import mcp.types as mt -from fastmcp.exceptions import AuthorizationError, NotFoundError +from fastmcp.exceptions import AuthorizationError from fastmcp.prompts.prompt import Prompt, PromptResult from fastmcp.resources.resource import Resource, ResourceResult from fastmcp.resources.template import ResourceTemplate @@ -136,7 +136,12 @@ class AuthMiddleware(Middleware): f"Authorization failed for tool '{tool_name}': missing context" ) - tool = await fastmcp.fastmcp.get_tool(tool_name) + # Use _get_tool to resolve transformed names from MCP protocol + tool = await fastmcp.fastmcp._get_tool(tool_name) + if tool is None: + raise AuthorizationError( + f"Authorization failed for tool '{tool_name}': tool not found" + ) token = get_access_token() ctx = AuthContext(token=token, component=tool) @@ -197,10 +202,14 @@ class AuthMiddleware(Middleware): ) # Try concrete resource first, then template (for template-backed URIs) - try: - component = await fastmcp.fastmcp.get_resource(str(uri)) - except NotFoundError: - component = await fastmcp.fastmcp.get_resource_template(str(uri)) + # Use _get_* to resolve transformed URIs from MCP protocol + component = await fastmcp.fastmcp._get_resource(str(uri)) + if component is None: + component = await fastmcp.fastmcp._get_resource_template(str(uri)) + if component is None: + raise AuthorizationError( + f"Authorization failed for resource '{uri}': resource not found" + ) token = get_access_token() ctx = AuthContext(token=token, component=component) @@ -286,7 +295,12 @@ class AuthMiddleware(Middleware): f"Authorization failed for prompt '{prompt_name}': missing context" ) - prompt = await fastmcp.fastmcp.get_prompt(prompt_name) + # Use _get_prompt to resolve transformed names from MCP protocol + prompt = await fastmcp.fastmcp._get_prompt(prompt_name) + if prompt is None: + raise AuthorizationError( + f"Authorization failed for prompt '{prompt_name}': prompt not found" + ) token = get_access_token() ctx = AuthContext(token=token, component=prompt) diff --git a/src/fastmcp/server/providers/__init__.py b/src/fastmcp/server/providers/__init__.py index 39ea4d0ff..8d997e5ff 100644 --- a/src/fastmcp/server/providers/__init__.py +++ b/src/fastmcp/server/providers/__init__.py @@ -27,6 +27,7 @@ Example: from typing import TYPE_CHECKING +from fastmcp.server.providers.aggregate import AggregateProvider from fastmcp.server.providers.base import Provider from fastmcp.server.providers.fastmcp_provider import FastMCPProvider from fastmcp.server.providers.filesystem import FileSystemProvider @@ -37,6 +38,7 @@ if TYPE_CHECKING: from fastmcp.server.providers.proxy import ProxyProvider as ProxyProvider __all__ = [ + "AggregateProvider", "FastMCPProvider", "FileSystemProvider", "LocalProvider", diff --git a/src/fastmcp/server/providers/aggregate.py b/src/fastmcp/server/providers/aggregate.py index 9beb6a917..4e0937d10 100644 --- a/src/fastmcp/server/providers/aggregate.py +++ b/src/fastmcp/server/providers/aggregate.py @@ -1,8 +1,19 @@ """AggregateProvider for combining multiple providers into one. -This module provides `AggregateProvider` which presents multiple providers -as a single unified provider. Used internally by FastMCP for aggregating -components from all providers. +This module provides `AggregateProvider`, a utility class that presents +multiple providers as a single unified provider. Useful when you want to +combine custom providers without creating a full FastMCP server. + +Example: + ```python + from fastmcp.server.providers import AggregateProvider + + # Combine multiple providers into one + combined = AggregateProvider([provider1, provider2, provider3]) + + # Use like any other provider + tools = await combined.list_tools() + ``` """ from __future__ import annotations @@ -27,13 +38,17 @@ T = TypeVar("T") class AggregateProvider(Provider): - """Presents multiple providers as a single provider. + """Utility provider that combines multiple providers into one. - Components are aggregated from all providers. For get_* operations, - providers are queried in parallel and the first non-None result is returned. + Components are aggregated from all providers. For list operations, results + from all providers are combined. For get operations, providers are queried + in parallel and the first non-None result is returned. - Errors from individual providers are logged and skipped (graceful degradation). - This matches the behavior of FastMCP's original provider iteration. + Errors from individual providers are logged and skipped (graceful degradation), + allowing the aggregate to continue working even if one provider fails. + + This is useful when you want to combine custom providers without creating + a full FastMCP server. """ def __init__(self, providers: Sequence[Provider]) -> None: diff --git a/src/fastmcp/server/providers/base.py b/src/fastmcp/server/providers/base.py index 239903cb3..55a54b83f 100644 --- a/src/fastmcp/server/providers/base.py +++ b/src/fastmcp/server/providers/base.py @@ -64,9 +64,14 @@ class Provider: """ def __init__(self) -> None: - # Visibility is the first (innermost) transform - closest to the base provider + # Visibility is kept separate from user transforms and applied last (outermost) + # so that enable/disable works with transformed names, not raw names self._visibility = Visibility() - self._transforms: list[Transform] = [self._visibility] + self._transforms: list[Transform] = [] + + def _get_all_transforms(self) -> list[Transform]: + """Get all transforms including visibility (applied last/outermost).""" + return [*self._transforms, self._visibility] def __repr__(self) -> str: return f"{self.__class__.__name__}()" @@ -109,7 +114,7 @@ class Provider: return await self.list_tools() chain = base - for transform in self._transforms: + for transform in self._get_all_transforms(): chain = partial(transform.list_tools, call_next=chain) return await chain() @@ -128,7 +133,7 @@ class Provider: return await self.get_tool(n) chain = base - for transform in self._transforms: + for transform in self._get_all_transforms(): chain = partial(transform.get_tool, call_next=chain) return await chain(name) @@ -140,7 +145,7 @@ class Provider: return await self.list_resources() chain = base - for transform in self._transforms: + for transform in self._get_all_transforms(): chain = partial(transform.list_resources, call_next=chain) return await chain() @@ -152,7 +157,7 @@ class Provider: return await self.get_resource(u) chain = base - for transform in self._transforms: + for transform in self._get_all_transforms(): chain = partial(transform.get_resource, call_next=chain) return await chain(uri) @@ -164,7 +169,7 @@ class Provider: return await self.list_resource_templates() chain = base - for transform in self._transforms: + for transform in self._get_all_transforms(): chain = partial(transform.list_resource_templates, call_next=chain) return await chain() @@ -176,7 +181,7 @@ class Provider: return await self.get_resource_template(u) chain = base - for transform in self._transforms: + for transform in self._get_all_transforms(): chain = partial(transform.get_resource_template, call_next=chain) return await chain(uri) @@ -188,7 +193,7 @@ class Provider: return await self.list_prompts() chain = base - for transform in self._transforms: + for transform in self._get_all_transforms(): chain = partial(transform.list_prompts, call_next=chain) return await chain() @@ -200,7 +205,7 @@ class Provider: return await self.get_prompt(n) chain = base - for transform in self._transforms: + for transform in self._get_all_transforms(): chain = partial(transform.get_prompt, call_next=chain) return await chain(name) @@ -330,13 +335,13 @@ class Provider: async def prompts_base() -> Sequence[Prompt]: return prompts - # Apply transforms in order (first is innermost) + # Apply transforms in order (visibility last/outermost) tools_chain = tools_base resources_chain = resources_base templates_chain = templates_base prompts_chain = prompts_base - for transform in self._transforms: + for transform in self._get_all_transforms(): tools_chain = partial(transform.list_tools, call_next=tools_chain) resources_chain = partial( transform.list_resources, call_next=resources_chain diff --git a/src/fastmcp/server/providers/fastmcp_provider.py b/src/fastmcp/server/providers/fastmcp_provider.py index 846775eb8..e6e801403 100644 --- a/src/fastmcp/server/providers/fastmcp_provider.py +++ b/src/fastmcp/server/providers/fastmcp_provider.py @@ -480,9 +480,9 @@ class FastMCPProvider(Provider): async def list_tools(self) -> Sequence[Tool]: """List all tools from the mounted server as FastMCPProviderTools. - Calls the nested server's middleware to list tools, then wraps - each tool as a FastMCPProviderTool that delegates execution to the - nested server's middleware. + Runs the mounted server's middleware so filtering/transformation applies. + Wraps each tool as a FastMCPProviderTool that delegates execution to + the nested server's middleware. """ raw_tools = await self.server.get_tools(run_middleware=True) return [FastMCPProviderTool.wrap(self.server, t) for t in raw_tools] @@ -499,9 +499,9 @@ class FastMCPProvider(Provider): async def list_resources(self) -> Sequence[Resource]: """List all resources from the mounted server as FastMCPProviderResources. - Calls the nested server's middleware to list resources, then wraps - each resource as a FastMCPProviderResource that delegates reading to the - nested server's middleware. + Runs the mounted server's middleware so filtering/transformation applies. + Wraps each resource as a FastMCPProviderResource that delegates reading + to the nested server's middleware. """ raw_resources = await self.server.get_resources(run_middleware=True) return [FastMCPProviderResource.wrap(self.server, r) for r in raw_resources] @@ -518,6 +518,7 @@ class FastMCPProvider(Provider): async def list_resource_templates(self) -> Sequence[ResourceTemplate]: """List all resource templates from the mounted server. + Runs the mounted server's middleware so filtering/transformation applies. Returns FastMCPProviderResourceTemplate instances that create FastMCPProviderResources when materialized. """ @@ -541,6 +542,7 @@ class FastMCPProvider(Provider): async def list_prompts(self) -> Sequence[Prompt]: """List all prompts from the mounted server as FastMCPProviderPrompts. + Runs the mounted server's middleware so filtering/transformation applies. Returns FastMCPProviderPrompt instances that delegate rendering to the wrapped server's middleware. """ @@ -560,12 +562,12 @@ class FastMCPProvider(Provider): """Return task-eligible components from the mounted server. Returns the child's ACTUAL components (not wrapped) so their actual - functions get registered with Docket. Uses _source_get_tasks() to get - components with child server's transforms applied, then applies this - provider's transforms for correct registration keys. + functions get registered with Docket. Gets components with child + server's transforms applied, then applies this provider's transforms + for correct registration keys. """ # Get tasks with child server's transforms already applied - components = list(await self.server._source_get_tasks()) + components = list(await self.server.get_tasks()) # Separate by type for this provider's transform application tools = [c for c in components if isinstance(c, Tool)] diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index 123fe582f..1d9e44a87 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -85,20 +85,19 @@ from fastmcp.server.lifespan import Lifespan from fastmcp.server.low_level import LowLevelServer from fastmcp.server.middleware import Middleware, MiddlewareContext from fastmcp.server.providers import LocalProvider, Provider -from fastmcp.server.providers.aggregate import AggregateProvider from fastmcp.server.tasks.config import TaskConfig, TaskMeta from fastmcp.server.telemetry import server_span from fastmcp.server.transforms import ( Namespace, ToolTransform, Transform, - Visibility, ) from fastmcp.settings import DuplicateBehavior as DuplicateBehaviorSetting from fastmcp.settings import Settings from fastmcp.tools.function_tool import FunctionTool from fastmcp.tools.tool import AuthCheckCallable, Tool, ToolResult from fastmcp.tools.tool_transform import ToolTransformConfig +from fastmcp.utilities.async_utils import gather from fastmcp.utilities.cli import log_server_banner from fastmcp.utilities.components import FastMCPComponent from fastmcp.utilities.logging import get_logger, temporary_log_level @@ -229,7 +228,7 @@ class StateValue(FastMCPBaseModel): value: Any -class FastMCP(Generic[LifespanResultT]): +class FastMCP(Provider, Generic[LifespanResultT]): def __init__( self, name: str | None = None, @@ -271,6 +270,9 @@ class FastMCP(Generic[LifespanResultT]): sampling_handler_behavior: Literal["always", "fallback"] | None = None, tool_transformations: Mapping[str, ToolTransformConfig] | None = None, ): + # Initialize Provider (sets up _transforms and _visibility) + super().__init__() + # Resolve on_duplicate from deprecated params (delete when removing deprecation) self._on_duplicate: DuplicateBehaviorSetting = _resolve_on_duplicate( on_duplicate, @@ -301,10 +303,8 @@ class FastMCP(Generic[LifespanResultT]): on_duplicate=self._on_duplicate ) - # Server-level transforms (applied after provider aggregation) - self._transforms: list[Transform] = [] - # Local provider is always first in the provider list + # Note: _transforms is initialized by Provider.__init__() self._providers: list[Provider] = [ self._local_provider, *(providers or []), @@ -357,10 +357,7 @@ class FastMCP(Generic[LifespanResultT]): tool = Tool.from_function(tool, serializer=self._tool_serializer) self.add_tool(tool) - # Server-level visibility for runtime enable/disable - self._visibility = Visibility() - - # Emit deprecation warnings for include_tags and exclude_tags + # Handle deprecated include_tags and exclude_tags parameters if include_tags is not None: warnings.warn( "include_tags is deprecated. Use server.enable(tags=..., only=True) instead.", @@ -535,11 +532,10 @@ class FastMCP(Generic[LifespanResultT]): return # Collect task-enabled components at startup with all transforms applied. - # Uses _source_get_tasks() to include both provider and server-level transforms. # Components must be available now to be registered with Docket workers; # dynamically added components after startup won't be registered. try: - task_components = list(await self._source_get_tasks()) + task_components = list(await self.get_tasks()) except Exception as e: logger.warning(f"Failed to get tasks: {e}") if fastmcp.settings.mounted_components_raise_on_load_error: @@ -807,135 +803,74 @@ class FastMCP(Generic[LifespanResultT]): # Tool Transforms # ------------------------------------------------------------------------- - def _get_root_provider(self) -> AggregateProvider: - """Get the root provider (aggregate of all providers). + # Note: _get_all_transforms() is inherited from Provider - Returns an AggregateProvider wrapping all providers. Each provider - applies its own transforms and provider-level visibility. - Server-level transforms and visibility are applied by _source_* methods. - """ - return AggregateProvider(self._providers) - - def _get_all_transforms(self) -> list[Transform]: - """Get all server-level transforms (including visibility as last).""" - return [*self._transforms, self._visibility] + def _collect_list_results( + self, results: list[Sequence[Any] | BaseException], operation: str + ) -> list[Any]: + """Collect successful list results, logging any exceptions.""" + collected: list[Any] = [] + for i, result in enumerate(results): + if isinstance(result, BaseException): + logger.debug( + f"Error during {operation} from provider " + f"{self._providers[i]}: {result}" + ) + continue + collected.extend(result) + return collected # ------------------------------------------------------------------------- - # Server-level transform chain building + # Provider interface overrides (aggregate from sub-providers) # ------------------------------------------------------------------------- - async def _source_list_tools(self) -> Sequence[Tool]: - """List tools with all transforms applied (provider + server).""" - root = self._get_root_provider() + async def list_tools(self) -> Sequence[Tool]: + """Aggregate tools from all sub-providers. - async def base() -> Sequence[Tool]: - return await root.list_tools() - - chain = base - for transform in self._get_all_transforms(): - chain = partial(transform.list_tools, call_next=chain) - - return await chain() - - async def _source_get_tool(self, name: str) -> Tool | None: - """Get tool by name with all transforms applied (provider + server).""" - root = self._get_root_provider() - - async def base(n: str) -> Tool | None: - return await root.get_tool(n) - - chain = base - for transform in self._get_all_transforms(): - chain = partial(transform.get_tool, call_next=chain) - - return await chain(name) - - async def _source_list_resources(self) -> Sequence[Resource]: - """List resources with all transforms applied (provider + server).""" - root = self._get_root_provider() - - async def base() -> Sequence[Resource]: - return await root.list_resources() - - chain = base - for transform in self._get_all_transforms(): - chain = partial(transform.list_resources, call_next=chain) - - return await chain() - - async def _source_get_resource(self, uri: str) -> Resource | None: - """Get resource by URI with all transforms applied (provider + server).""" - root = self._get_root_provider() - - async def base(u: str) -> Resource | None: - return await root.get_resource(u) - - chain = base - for transform in self._get_all_transforms(): - chain = partial(transform.get_resource, call_next=chain) - - return await chain(uri) - - async def _source_list_resource_templates(self) -> Sequence[ResourceTemplate]: - """List resource templates with all transforms applied (provider + server).""" - root = self._get_root_provider() - - async def base() -> Sequence[ResourceTemplate]: - return await root.list_resource_templates() - - chain = base - for transform in self._get_all_transforms(): - chain = partial(transform.list_resource_templates, call_next=chain) - - return await chain() - - async def _source_get_resource_template(self, uri: str) -> ResourceTemplate | None: - """Get resource template by URI with all transforms applied (provider + server).""" - root = self._get_root_provider() - - async def base(u: str) -> ResourceTemplate | None: - return await root.get_resource_template(u) - - chain = base - for transform in self._get_all_transforms(): - chain = partial(transform.get_resource_template, call_next=chain) - - return await chain(uri) - - async def _source_list_prompts(self) -> Sequence[Prompt]: - """List prompts with all transforms applied (provider + server).""" - root = self._get_root_provider() - - async def base() -> Sequence[Prompt]: - return await root.list_prompts() - - chain = base - for transform in self._get_all_transforms(): - chain = partial(transform.list_prompts, call_next=chain) - - return await chain() - - async def _source_get_prompt(self, name: str) -> Prompt | None: - """Get prompt by name with all transforms applied (provider + server).""" - root = self._get_root_provider() - - async def base(n: str) -> Prompt | None: - return await root.get_prompt(n) - - chain = base - for transform in self._get_all_transforms(): - chain = partial(transform.get_prompt, call_next=chain) - - return await chain(name) - - async def _source_get_tasks(self) -> Sequence[FastMCPComponent]: - """Get tasks with all transforms applied (provider + server). - - Collects task-eligible components from all providers and applies - server-level transforms to ensure task registration uses transformed names. + This is the Provider interface implementation. The inherited _list_tools() + applies server-level transforms over this method. """ - root = self._get_root_provider() - components = list(await root.get_tasks()) + results = await gather( + *[p._list_tools() for p in self._providers], + return_exceptions=True, + ) + return self._collect_list_results(results, "list_tools") + + async def list_resources(self) -> Sequence[Resource]: + """Aggregate resources from all sub-providers.""" + results = await gather( + *[p._list_resources() for p in self._providers], + return_exceptions=True, + ) + return self._collect_list_results(results, "list_resources") + + async def list_resource_templates(self) -> Sequence[ResourceTemplate]: + """Aggregate resource templates from all sub-providers.""" + results = await gather( + *[p._list_resource_templates() for p in self._providers], + return_exceptions=True, + ) + return self._collect_list_results(results, "list_resource_templates") + + async def list_prompts(self) -> Sequence[Prompt]: + """Aggregate prompts from all sub-providers.""" + results = await gather( + *[p._list_prompts() for p in self._providers], + return_exceptions=True, + ) + return self._collect_list_results(results, "list_prompts") + + async def get_tasks(self) -> Sequence[FastMCPComponent]: + """Get task-eligible components with all transforms applied. + + Overrides Provider.get_tasks() to collect task-eligible components + from all sub-providers and apply server-level transforms. + """ + results = await gather( + *[p.get_tasks() for p in self._providers], + return_exceptions=True, + ) + components = self._collect_list_results(results, "get_tasks") # Separate by component type for transform application tools = [c for c in components if isinstance(c, Tool)] @@ -1114,66 +1049,70 @@ class FastMCP(Generic[LifespanResultT]): server transforms, and visibility filtering). First provider wins for duplicate keys. Args: - run_middleware: If True, apply the middleware chain before - returning results. Used by MCP handlers and mounted servers. + run_middleware: If True, apply the middleware chain before returning. + Used by MCP handlers and FastMCPProvider for nested servers. """ - if run_middleware: - async with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx: + async with fastmcp.server.context.Context(fastmcp=self) as ctx: + if run_middleware: mw_context = MiddlewareContext( message=mcp.types.ListToolsRequest(method="tools/list"), source="client", type="request", method="tools/list", - fastmcp_context=fastmcp_ctx, + fastmcp_context=ctx, ) - return list( - await self._run_middleware( - context=mw_context, - call_next=lambda context: self.get_tools(run_middleware=False), - ) + return await self._run_middleware( + context=mw_context, + call_next=lambda context: self.get_tools(run_middleware=False), ) - # Query through full transform chain (provider transforms + server transforms + visibility) - tools = await self._source_list_tools() + # Query through full transform chain (provider transforms + server transforms + visibility) + tools = await self._list_tools() - # Get auth context (skip_auth=True for STDIO which has no auth concept) - skip_auth, token = _get_auth_context() + # Get auth context (skip_auth=True for STDIO which has no auth concept) + skip_auth, token = _get_auth_context() - # Deduplicate by key (first wins) and apply authorization checks - seen: dict[str, Tool] = {} - for tool in tools: - if tool.key in seen: - continue - # Check tool-level auth (skip for STDIO) - if not skip_auth and tool.auth is not None: - ctx = AuthContext(token=token, component=tool) - try: - if not run_auth_checks(tool.auth, ctx): - continue - except AuthorizationError: - # Treat auth errors as denials in list operations + # Deduplicate by key (first wins) and apply authorization checks + seen: dict[str, Tool] = {} + for tool in tools: + if tool.key in seen: continue - seen[tool.key] = tool - return list(seen.values()) + # Check tool-level auth (skip for STDIO) + if not skip_auth and tool.auth is not None: + auth_ctx = AuthContext(token=token, component=tool) + try: + if not run_auth_checks(tool.auth, auth_ctx): + continue + except AuthorizationError: + # Treat auth errors as denials in list operations + continue + seen[tool.key] = tool + return list(seen.values()) - async def get_tool(self, name: str) -> Tool: - """Get an enabled tool by name. + async def get_tool(self, name: str) -> Tool | None: + """Provider interface: aggregate tool lookup from all providers. - Queries providers with full transform chain (provider transforms + server transforms + visibility). - Returns only if enabled and authorized. + Returns None if not found or if the tool is disabled via visibility settings. """ - tool = await self._source_get_tool(name) - if tool is None: - raise NotFoundError(f"Unknown tool: {name!r}") + # Query all providers in parallel for efficient lookup + results = await gather( + *[p._get_tool(name) for p in self._providers], + return_exceptions=True, + ) - # Check tool-level auth (skip for STDIO) - skip_auth, token = _get_auth_context() - if not skip_auth and tool.auth is not None: - ctx = AuthContext(token=token, component=tool) - if not run_auth_checks(tool.auth, ctx): - raise NotFoundError(f"Unknown tool: {name!r}") + # Return first non-None, enabled result + for i, result in enumerate(results): + if isinstance(result, BaseException): + if not isinstance(result, NotFoundError): + logger.debug( + f"Error during get_tool({name!r}) from provider " + f"{self._providers[i]}: {result}" + ) + continue + if result is not None and self._is_component_enabled(result): + return result - return tool + return None async def get_resources(self, *, run_middleware: bool = False) -> list[Resource]: """Get all enabled resources from providers. @@ -1182,68 +1121,70 @@ class FastMCP(Generic[LifespanResultT]): server transforms, and visibility filtering). First provider wins for duplicate keys. Args: - run_middleware: If True, apply the middleware chain before - returning results. Used by MCP handlers and mounted servers. + run_middleware: If True, apply the middleware chain before returning. + Used by MCP handlers and FastMCPProvider for nested servers. """ - if run_middleware: - async with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx: + async with fastmcp.server.context.Context(fastmcp=self) as ctx: + if run_middleware: mw_context = MiddlewareContext( - message={}, # List resources doesn't have parameters + message={}, source="client", type="request", method="resources/list", - fastmcp_context=fastmcp_ctx, + fastmcp_context=ctx, ) - return list( - await self._run_middleware( - context=mw_context, - call_next=lambda context: self.get_resources( - run_middleware=False - ), - ) + return await self._run_middleware( + context=mw_context, + call_next=lambda context: self.get_resources(run_middleware=False), ) - # Query through full transform chain (provider transforms + server transforms + visibility) - resources = await self._source_list_resources() + # Query through full transform chain (provider transforms + server transforms + visibility) + resources = await self._list_resources() - # Get auth context (skip_auth=True for STDIO which has no auth concept) - skip_auth, token = _get_auth_context() + # Get auth context (skip_auth=True for STDIO which has no auth concept) + skip_auth, token = _get_auth_context() - # Deduplicate by key (first wins) and apply authorization checks - seen: dict[str, Resource] = {} - for resource in resources: - if resource.key in seen: - continue - # Check resource-level auth (skip for STDIO) - if not skip_auth and resource.auth is not None: - ctx = AuthContext(token=token, component=resource) - try: - if not run_auth_checks(resource.auth, ctx): - continue - except AuthorizationError: - # Treat auth errors as denials in list operations + # Deduplicate by key (first wins) and apply authorization checks + seen: dict[str, Resource] = {} + for resource in resources: + if resource.key in seen: continue - seen[resource.key] = resource - return list(seen.values()) + # Check resource-level auth (skip for STDIO) + if not skip_auth and resource.auth is not None: + auth_ctx = AuthContext(token=token, component=resource) + try: + if not run_auth_checks(resource.auth, auth_ctx): + continue + except AuthorizationError: + # Treat auth errors as denials in list operations + continue + seen[resource.key] = resource + return list(seen.values()) - async def get_resource(self, uri: str) -> Resource: - """Get an enabled resource by URI. + async def get_resource(self, uri: str) -> Resource | None: + """Provider interface: aggregate resource lookup from all providers. - Queries providers with full transform chain (provider transforms + server transforms + visibility). - Returns only if enabled and authorized. + Returns None if not found or if the resource is disabled via visibility settings. """ - resource = await self._source_get_resource(uri) - if resource is None: - raise NotFoundError(f"Unknown resource: {uri}") + # Query all providers in parallel for efficient lookup + results = await gather( + *[p._get_resource(uri) for p in self._providers], + return_exceptions=True, + ) - # Check resource-level auth (skip for STDIO) - skip_auth, token = _get_auth_context() - if not skip_auth and resource.auth is not None: - ctx = AuthContext(token=token, component=resource) - if not run_auth_checks(resource.auth, ctx): - raise NotFoundError(f"Unknown resource: {uri}") + # Return first non-None, enabled result + for i, result in enumerate(results): + if isinstance(result, BaseException): + if not isinstance(result, NotFoundError): + logger.debug( + f"Error during get_resource({uri!r}) from provider " + f"{self._providers[i]}: {result}" + ) + continue + if result is not None and self._is_component_enabled(result): + return result - return resource + return None async def get_resource_templates( self, *, run_middleware: bool = False @@ -1254,68 +1195,72 @@ class FastMCP(Generic[LifespanResultT]): server transforms, and visibility filtering). First provider wins for duplicate keys. Args: - run_middleware: If True, apply the middleware chain before - returning results. Used by MCP handlers and mounted servers. + run_middleware: If True, apply the middleware chain before returning. + Used by MCP handlers and FastMCPProvider for nested servers. """ - if run_middleware: - async with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx: + async with fastmcp.server.context.Context(fastmcp=self) as ctx: + if run_middleware: mw_context = MiddlewareContext( - message={}, # List resource templates doesn't have parameters + message={}, source="client", type="request", method="resources/templates/list", - fastmcp_context=fastmcp_ctx, + fastmcp_context=ctx, ) - return list( - await self._run_middleware( - context=mw_context, - call_next=lambda context: self.get_resource_templates( - run_middleware=False - ), - ) + return await self._run_middleware( + context=mw_context, + call_next=lambda context: self.get_resource_templates( + run_middleware=False + ), ) - # Query through full transform chain (provider transforms + server transforms + visibility) - templates = await self._source_list_resource_templates() + # Query through full transform chain (provider transforms + server transforms + visibility) + templates = await self._list_resource_templates() - # Get auth context (skip_auth=True for STDIO which has no auth concept) - skip_auth, token = _get_auth_context() + # Get auth context (skip_auth=True for STDIO which has no auth concept) + skip_auth, token = _get_auth_context() - # Deduplicate by key (first wins) and apply authorization checks - seen: dict[str, ResourceTemplate] = {} - for template in templates: - if template.key in seen: - continue - # Check template-level auth (skip for STDIO) - if not skip_auth and template.auth is not None: - ctx = AuthContext(token=token, component=template) - try: - if not run_auth_checks(template.auth, ctx): - continue - except AuthorizationError: - # Treat auth errors as denials in list operations + # Deduplicate by key (first wins) and apply authorization checks + seen: dict[str, ResourceTemplate] = {} + for template in templates: + if template.key in seen: continue - seen[template.key] = template - return list(seen.values()) + # Check template-level auth (skip for STDIO) + if not skip_auth and template.auth is not None: + auth_ctx = AuthContext(token=token, component=template) + try: + if not run_auth_checks(template.auth, auth_ctx): + continue + except AuthorizationError: + # Treat auth errors as denials in list operations + continue + seen[template.key] = template + return list(seen.values()) - async def get_resource_template(self, uri: str) -> ResourceTemplate: - """Get an enabled resource template that matches the given URI. + async def get_resource_template(self, uri: str) -> ResourceTemplate | None: + """Provider interface: aggregate template lookup from all providers. - Queries providers with full transform chain (provider transforms + server transforms + visibility). - Returns only if enabled and authorized. + Returns None if not found or if the template is disabled via visibility settings. """ - template = await self._source_get_resource_template(uri) - if template is None: - raise NotFoundError(f"Unknown resource template: {uri}") + # Query all providers in parallel for efficient lookup + results = await gather( + *[p._get_resource_template(uri) for p in self._providers], + return_exceptions=True, + ) - # Check template-level auth (skip for STDIO) - skip_auth, token = _get_auth_context() - if not skip_auth and template.auth is not None: - ctx = AuthContext(token=token, component=template) - if not run_auth_checks(template.auth, ctx): - raise NotFoundError(f"Unknown resource template: {uri}") + # Return first non-None, enabled result + for i, result in enumerate(results): + if isinstance(result, BaseException): + if not isinstance(result, NotFoundError): + logger.debug( + f"Error during get_resource_template({uri!r}) from provider " + f"{self._providers[i]}: {result}" + ) + continue + if result is not None and self._is_component_enabled(result): + return result - return template + return None async def get_prompts(self, *, run_middleware: bool = False) -> list[Prompt]: """Get all enabled prompts from providers. @@ -1324,95 +1269,96 @@ class FastMCP(Generic[LifespanResultT]): server transforms, and visibility filtering). First provider wins for duplicate keys. Args: - run_middleware: If True, apply the middleware chain before - returning results. Used by MCP handlers and mounted servers. + run_middleware: If True, apply the middleware chain before returning. + Used by MCP handlers and FastMCPProvider for nested servers. """ - if run_middleware: - async with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx: + async with fastmcp.server.context.Context(fastmcp=self) as ctx: + if run_middleware: mw_context = MiddlewareContext( - message=mcp.types.ListPromptsRequest(method="prompts/list"), + message={}, source="client", type="request", method="prompts/list", - fastmcp_context=fastmcp_ctx, + fastmcp_context=ctx, ) - return list( - await self._run_middleware( - context=mw_context, - call_next=lambda context: self.get_prompts( - run_middleware=False - ), - ) + return await self._run_middleware( + context=mw_context, + call_next=lambda context: self.get_prompts(run_middleware=False), ) - # Query through full transform chain (provider transforms + server transforms + visibility) - prompts = await self._source_list_prompts() + # Query through full transform chain (provider transforms + server transforms + visibility) + prompts = await self._list_prompts() - # Get auth context (skip_auth=True for STDIO which has no auth concept) - skip_auth, token = _get_auth_context() + # Get auth context (skip_auth=True for STDIO which has no auth concept) + skip_auth, token = _get_auth_context() - # Deduplicate by key (first wins) and apply authorization checks - seen: dict[str, Prompt] = {} - for prompt in prompts: - if prompt.key in seen: - continue - # Check prompt-level auth (skip for STDIO) - if not skip_auth and prompt.auth is not None: - ctx = AuthContext(token=token, component=prompt) - try: - if not run_auth_checks(prompt.auth, ctx): - continue - except AuthorizationError: - # Treat auth errors as denials in list operations + # Deduplicate by key (first wins) and apply authorization checks + seen: dict[str, Prompt] = {} + for prompt in prompts: + if prompt.key in seen: continue - seen[prompt.key] = prompt - return list(seen.values()) + # Check prompt-level auth (skip for STDIO) + if not skip_auth and prompt.auth is not None: + auth_ctx = AuthContext(token=token, component=prompt) + try: + if not run_auth_checks(prompt.auth, auth_ctx): + continue + except AuthorizationError: + # Treat auth errors as denials in list operations + continue + seen[prompt.key] = prompt + return list(seen.values()) - async def get_prompt(self, name: str) -> Prompt: - """Get an enabled prompt by name. + async def get_prompt(self, name: str) -> Prompt | None: + """Provider interface: aggregate prompt lookup from all providers. - Queries providers with full transform chain (provider transforms + server transforms + visibility). - Returns only if enabled and authorized. + Returns None if not found or if the prompt is disabled via visibility settings. """ - prompt = await self._source_get_prompt(name) - if prompt is None: - raise NotFoundError(f"Unknown prompt: {name}") + # Query all providers in parallel for efficient lookup + results = await gather( + *[p._get_prompt(name) for p in self._providers], + return_exceptions=True, + ) - # Check prompt-level auth (skip for STDIO) - skip_auth, token = _get_auth_context() - if not skip_auth and prompt.auth is not None: - ctx = AuthContext(token=token, component=prompt) - if not run_auth_checks(prompt.auth, ctx): - raise NotFoundError(f"Unknown prompt: {name}") + # Return first non-None, enabled result + for i, result in enumerate(results): + if isinstance(result, BaseException): + if not isinstance(result, NotFoundError): + logger.debug( + f"Error during get_prompt({name!r}) from provider " + f"{self._providers[i]}: {result}" + ) + continue + if result is not None and self._is_component_enabled(result): + return result - return prompt + return None async def get_component( self, key: str - ) -> Tool | Resource | ResourceTemplate | Prompt: + ) -> Tool | Resource | ResourceTemplate | Prompt | None: """Get a component by its prefixed key. - Routes to the appropriate get_* method which applies server-level layers. + Uses _get_* methods to apply server-level transforms (including namespace + resolution for mounted servers). Task keys use transformed names, so this + ensures proper resolution. Args: key: The prefixed key (e.g., "tool:name", "resource:uri", "template:uri"). Returns: - The component if found. - - Raises: - NotFoundError: If no component is found with the given key. + The component if found, None otherwise. """ - # Parse key and delegate to specific methods which apply layers + # Parse key and delegate to _get_* methods which apply server transforms if key.startswith("tool:"): - return await self.get_tool(key[5:]) + return await self._get_tool(key[5:]) elif key.startswith("resource:"): - return await self.get_resource(key[9:]) + return await self._get_resource(key[9:]) elif key.startswith("template:"): - return await self.get_resource_template(key[9:]) + return await self._get_resource_template(key[9:]) elif key.startswith("prompt:"): - return await self.get_prompt(key[7:]) - raise NotFoundError(f"Unknown component: {key}") + return await self._get_prompt(key[7:]) + return None @overload async def call_tool( @@ -1493,7 +1439,18 @@ class FastMCP(Generic[LifespanResultT]): with server_span( f"tools/call {name}", "tools/call", self.name, "tool", name ) as span: - tool = await self.get_tool(name) + # Use _get_tool() to apply server transforms including visibility + tool = await self._get_tool(name) + if tool is None: + raise NotFoundError(f"Unknown tool: {name!r}") + + # Check tool-level auth (skip for STDIO which has no auth concept) + skip_auth, token = _get_auth_context() + if not skip_auth and tool.auth is not None: + auth_ctx = AuthContext(token=token, component=tool) + if not run_auth_checks(tool.auth, auth_ctx): + raise NotFoundError(f"Unknown tool: {name!r}") + span.set_attributes(tool.get_span_attributes()) if task_meta is not None and task_meta.fn_key is None: task_meta = replace(task_meta, fn_key=tool.key) @@ -1592,29 +1549,49 @@ class FastMCP(Generic[LifespanResultT]): uri, resource_uri=uri, ) as span: - # Try concrete resources first - try: - resource = await self.get_resource(uri) - span.set_attributes(resource.get_span_attributes()) - if task_meta is not None and task_meta.fn_key is None: - task_meta = replace(task_meta, fn_key=resource.key) - return await resource._read(task_meta=task_meta) - except NotFoundError: - pass # Fall through to try templates - except (FastMCPError, McpError): - logger.exception(f"Error reading resource {uri!r}") - raise - except Exception as e: - logger.exception(f"Error reading resource {uri!r}") - if self._mask_error_details: - raise ResourceError(f"Error reading resource {uri!r}") from e - raise ResourceError(f"Error reading resource {uri!r}: {e}") from e + # Get auth context (skip_auth=True for STDIO which has no auth concept) + skip_auth, token = _get_auth_context() - # Try templates - try: - template = await self.get_resource_template(uri) - except NotFoundError: - raise NotFoundError(f"Unknown resource: {uri!r}") from None + # Try concrete resources first (use _get_resource for server transforms) + resource = await self._get_resource(uri) + if resource is not None: + # Check resource-level auth - auth failure on concrete resource + # does NOT fall back to templates (deny immediately) + if not skip_auth and resource.auth is not None: + auth_ctx = AuthContext(token=token, component=resource) + if not run_auth_checks(resource.auth, auth_ctx): + raise NotFoundError(f"Unknown resource: {uri!r}") + + # Resource found and authorized + try: + span.set_attributes(resource.get_span_attributes()) + if task_meta is not None and task_meta.fn_key is None: + task_meta = replace(task_meta, fn_key=resource.key) + return await resource._read(task_meta=task_meta) + except (FastMCPError, McpError): + logger.exception(f"Error reading resource {uri!r}") + raise + except Exception as e: + logger.exception(f"Error reading resource {uri!r}") + if self._mask_error_details: + raise ResourceError( + f"Error reading resource {uri!r}" + ) from e + raise ResourceError( + f"Error reading resource {uri!r}: {e}" + ) from e + + # Try templates (use _get_resource_template for server transforms) + template = await self._get_resource_template(uri) + if template is not None: + # Check template-level auth + if not skip_auth and template.auth is not None: + auth_ctx = AuthContext(token=token, component=template) + if not run_auth_checks(template.auth, auth_ctx): + template = None + + if template is None: + raise NotFoundError(f"Unknown resource: {uri!r}") span.set_attributes(template.get_span_attributes()) params = template.matches(uri) assert params is not None @@ -1706,7 +1683,18 @@ class FastMCP(Generic[LifespanResultT]): with server_span( f"prompts/get {name}", "prompts/get", self.name, "prompt", name ) as span: - prompt = await self.get_prompt(name) + # Use _get_prompt() to apply server transforms including visibility + prompt = await self._get_prompt(name) + if prompt is None: + raise NotFoundError(f"Unknown prompt: {name!r}") + + # Check prompt-level auth (skip for STDIO which has no auth concept) + skip_auth, token = _get_auth_context() + if not skip_auth and prompt.auth is not None: + auth_ctx = AuthContext(token=token, component=prompt) + if not run_auth_checks(prompt.auth, auth_ctx): + raise NotFoundError(f"Unknown prompt: {name!r}") + span.set_attributes(prompt.get_span_attributes()) if task_meta is not None and task_meta.fn_key is None: task_meta = replace(task_meta, fn_key=prompt.key) @@ -1789,15 +1777,14 @@ class FastMCP(Generic[LifespanResultT]): """ logger.debug(f"[{self.name}] Handler called: list_tools") - async with fastmcp.server.context.Context(fastmcp=self): - tools = await self.get_tools(run_middleware=True) - return [ - tool.to_mcp_tool( - name=tool.name, - include_fastmcp_meta=self.include_fastmcp_meta, - ) - for tool in tools - ] + tools = await self.get_tools(run_middleware=True) + return [ + tool.to_mcp_tool( + name=tool.name, + include_fastmcp_meta=self.include_fastmcp_meta, + ) + for tool in tools + ] async def _list_resources_mcp(self) -> list[SDKResource]: """ @@ -1806,15 +1793,14 @@ class FastMCP(Generic[LifespanResultT]): """ logger.debug(f"[{self.name}] Handler called: list_resources") - async with fastmcp.server.context.Context(fastmcp=self): - resources = await self.get_resources(run_middleware=True) - return [ - resource.to_mcp_resource( - uri=str(resource.uri), - include_fastmcp_meta=self.include_fastmcp_meta, - ) - for resource in resources - ] + resources = await self.get_resources(run_middleware=True) + return [ + resource.to_mcp_resource( + uri=str(resource.uri), + include_fastmcp_meta=self.include_fastmcp_meta, + ) + for resource in resources + ] async def _list_resource_templates_mcp(self) -> list[SDKResourceTemplate]: """ @@ -1823,8 +1809,18 @@ class FastMCP(Generic[LifespanResultT]): """ logger.debug(f"[{self.name}] Handler called: list_resource_templates") - async with fastmcp.server.context.Context(fastmcp=self): - templates = await self.get_resource_templates(run_middleware=True) + async with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx: + mw_context = MiddlewareContext( + message={}, + source="client", + type="request", + method="resources/templates/list", + fastmcp_context=fastmcp_ctx, + ) + templates = await self._run_middleware( + context=mw_context, + call_next=lambda context: self.get_resource_templates(), + ) return [ template.to_mcp_template( uriTemplate=template.uri_template, @@ -1840,8 +1836,18 @@ class FastMCP(Generic[LifespanResultT]): """ logger.debug(f"[{self.name}] Handler called: list_prompts") - async with fastmcp.server.context.Context(fastmcp=self): - prompts = await self.get_prompts(run_middleware=True) + async with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx: + mw_context = MiddlewareContext( + message={}, + source="client", + type="request", + method="prompts/list", + fastmcp_context=fastmcp_ctx, + ) + prompts = await self._run_middleware( + context=mw_context, + call_next=lambda context: self.get_prompts(), + ) return [ prompt.to_mcp_prompt( name=prompt.name, diff --git a/src/fastmcp/utilities/inspect.py b/src/fastmcp/utilities/inspect.py index 38002799f..a2234096a 100644 --- a/src/fastmcp/utilities/inspect.py +++ b/src/fastmcp/utilities/inspect.py @@ -106,11 +106,11 @@ async def inspect_fastmcp_v2(mcp: FastMCP[Any]) -> FastMCPInfo: Returns: FastMCPInfo dataclass containing the extracted information """ - # Get all components directly without middleware (auth, rate limiting, etc.) - tools_list = await mcp.get_tools(run_middleware=False) - prompts_list = await mcp.get_prompts(run_middleware=False) - resources_list = await mcp.get_resources(run_middleware=False) - templates_list = await mcp.get_resource_templates(run_middleware=False) + # Get all components + tools_list = await mcp.get_tools() + prompts_list = await mcp.get_prompts() + resources_list = await mcp.get_resources() + templates_list = await mcp.get_resource_templates() # Extract detailed tool information tool_infos = [] diff --git a/tests/server/auth/test_authorization.py b/tests/server/auth/test_authorization.py index 55acd5e7b..f137d638d 100644 --- a/tests/server/auth/test_authorization.py +++ b/tests/server/auth/test_authorization.py @@ -8,7 +8,6 @@ from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser from fastmcp import FastMCP from fastmcp.client import Client -from fastmcp.exceptions import NotFoundError from fastmcp.server.auth import ( AccessToken, AuthContext, @@ -302,15 +301,21 @@ class TestToolLevelAuth: finally: auth_context_var.reset(tok) - async def test_get_tool_returns_not_found_without_auth(self): + async def test_get_tool_returns_tool_without_auth_check(self): + """get_tool() is the Provider interface and doesn't check auth. + + Auth is checked in call_tool() during execution, not during lookup. + """ mcp = FastMCP() @mcp.tool(auth=require_auth) def protected_tool() -> str: return "protected" - with pytest.raises(NotFoundError): - await mcp.get_tool("protected_tool") + # get_tool() returns the tool without checking auth + tool = await mcp.get_tool("protected_tool") + assert tool is not None + assert tool.name == "protected_tool" async def test_get_tool_returns_tool_with_auth(self): mcp = FastMCP() @@ -334,6 +339,12 @@ class TestToolLevelAuth: class TestAuthMiddleware: + """Tests for middleware filtering via MCP handler layer. + + These tests call _list_tools_mcp() which applies middleware during list, + simulating what happens when a client calls list_tools over MCP. + """ + async def test_middleware_filters_tools_without_token(self): mcp = FastMCP(middleware=[AuthMiddleware(auth=require_auth)]) @@ -342,7 +353,7 @@ class TestAuthMiddleware: return "public" # No token - all tools filtered by middleware - tools = await mcp.get_tools(run_middleware=True) + tools = await mcp._list_tools_mcp() assert len(tools) == 0 async def test_middleware_allows_tools_with_token(self): @@ -355,7 +366,7 @@ class TestAuthMiddleware: token = make_token() tok = set_token(token) try: - tools = await mcp.get_tools(run_middleware=True) + tools = await mcp._list_tools_mcp() assert len(tools) == 1 finally: auth_context_var.reset(tok) @@ -371,7 +382,7 @@ class TestAuthMiddleware: token = make_token(scopes=["read"]) tok = set_token(token) try: - tools = await mcp.get_tools(run_middleware=True) + tools = await mcp._list_tools_mcp() assert len(tools) == 0 finally: auth_context_var.reset(tok) @@ -380,7 +391,7 @@ class TestAuthMiddleware: token = make_token(scopes=["api"]) tok = set_token(token) try: - tools = await mcp.get_tools(run_middleware=True) + tools = await mcp._list_tools_mcp() assert len(tools) == 1 finally: auth_context_var.reset(tok) @@ -399,7 +410,7 @@ class TestAuthMiddleware: return "admin" # No token - public tool allowed, admin tool blocked - tools = await mcp.get_tools(run_middleware=True) + tools = await mcp._list_tools_mcp() assert len(tools) == 1 assert tools[0].name == "public_tool" @@ -407,7 +418,7 @@ class TestAuthMiddleware: token = make_token(scopes=["admin"]) tok = set_token(token) try: - tools = await mcp.get_tools(run_middleware=True) + tools = await mcp._list_tools_mcp() assert len(tools) == 2 finally: auth_context_var.reset(tok) diff --git a/tests/server/providers/test_local_provider_prompts.py b/tests/server/providers/test_local_provider_prompts.py index 93ca93c1b..57fb9a055 100644 --- a/tests/server/providers/test_local_provider_prompts.py +++ b/tests/server/providers/test_local_provider_prompts.py @@ -9,7 +9,6 @@ import pytest from mcp.types import TextContent from fastmcp import Client, Context, FastMCP -from fastmcp.exceptions import NotFoundError from fastmcp.prompts.prompt import Prompt @@ -352,8 +351,9 @@ class TestPromptEnabled: prompts = await mcp.get_prompts() assert len(prompts) == 0 - with pytest.raises(NotFoundError, match="Unknown prompt"): - await mcp.get_prompt("sample_prompt") + # get_prompt() returns None for disabled prompts (Provider interface) + prompt = await mcp.get_prompt("sample_prompt") + assert prompt is None async def test_prompt_toggle_enabled(self): mcp = FastMCP() @@ -381,8 +381,9 @@ class TestPromptEnabled: prompts = await mcp.get_prompts() assert len(prompts) == 0 - with pytest.raises(NotFoundError, match="Unknown prompt"): - await mcp.get_prompt("sample_prompt") + # get_prompt() returns None for disabled prompts (Provider interface) + prompt = await mcp.get_prompt("sample_prompt") + assert prompt is None async def test_get_prompt_and_disable(self): mcp = FastMCP() @@ -398,8 +399,9 @@ class TestPromptEnabled: prompts = await mcp.get_prompts() assert len(prompts) == 0 - with pytest.raises(NotFoundError, match="Unknown prompt"): - await mcp.get_prompt("sample_prompt") + # get_prompt() returns None for disabled prompts (Provider interface) + prompt = await mcp.get_prompt("sample_prompt") + assert prompt is None async def test_cant_get_disabled_prompt(self): mcp = FastMCP() @@ -410,8 +412,9 @@ class TestPromptEnabled: mcp.disable(keys=["prompt:sample_prompt"]) - with pytest.raises(NotFoundError, match="Unknown prompt"): - await mcp.get_prompt("sample_prompt") + # get_prompt() returns None for disabled prompts (Provider interface) + prompt = await mcp.get_prompt("sample_prompt") + assert prompt is None class TestPromptTags: @@ -459,13 +462,13 @@ class TestPromptTags: result = await prompt.render({}) assert result.messages[0].content.text == "1" - with pytest.raises(NotFoundError, match="Unknown prompt"): - await mcp.get_prompt("prompt_2") + prompt = await mcp.get_prompt("prompt_2") + assert prompt is None async def test_read_prompt_excludes_tags(self): mcp = self.create_server(exclude_tags={"a"}) - with pytest.raises(NotFoundError, match="Unknown prompt"): - await mcp.get_prompt("prompt_1") + prompt = await mcp.get_prompt("prompt_1") + assert prompt is None prompt = await mcp.get_prompt("prompt_2") result = await prompt.render({}) diff --git a/tests/server/test_mount.py b/tests/server/test_mount.py index 29b0ad42b..c0c3774bd 100644 --- a/tests/server/test_mount.py +++ b/tests/server/test_mount.py @@ -838,8 +838,10 @@ class TestAsProxyKwarg: provider = mcp._providers[1] # With namespace, we get FastMCPProvider with a Namespace layer assert isinstance(provider, FastMCPProvider) - assert len(provider._transforms) == 2 # Visibility + Namespace - assert isinstance(provider._transforms[1], Namespace) + assert ( + len(provider._transforms) == 1 + ) # Just Namespace (Visibility is in _visibility) + assert isinstance(provider._transforms[0], Namespace) assert provider.server is sub async def test_as_proxy_false(self): @@ -852,8 +854,10 @@ class TestAsProxyKwarg: provider = mcp._providers[1] # With namespace, we get FastMCPProvider with a Namespace layer assert isinstance(provider, FastMCPProvider) - assert len(provider._transforms) == 2 # Visibility + Namespace - assert isinstance(provider._transforms[1], Namespace) + assert ( + len(provider._transforms) == 1 + ) # Just Namespace (Visibility is in _visibility) + assert isinstance(provider._transforms[0], Namespace) assert provider.server is sub async def test_as_proxy_true(self): @@ -866,8 +870,10 @@ class TestAsProxyKwarg: provider = mcp._providers[1] # With namespace, we get FastMCPProvider with a Namespace layer assert isinstance(provider, FastMCPProvider) - assert len(provider._transforms) == 2 # Visibility + Namespace - assert isinstance(provider._transforms[1], Namespace) + assert ( + len(provider._transforms) == 1 + ) # Just Namespace (Visibility is in _visibility) + assert isinstance(provider._transforms[0], Namespace) assert provider.server is not sub assert isinstance(provider.server, FastMCPProxy) @@ -891,8 +897,10 @@ class TestAsProxyKwarg: # Index 1 because LocalProvider is at index 0 provider = mcp._providers[1] assert isinstance(provider, FastMCPProvider) - assert len(provider._transforms) == 2 # Visibility + Namespace - assert isinstance(provider._transforms[1], Namespace) + assert ( + len(provider._transforms) == 1 + ) # Just Namespace (Visibility is in _visibility) + assert isinstance(provider._transforms[0], Namespace) assert provider.server is sub async def test_as_proxy_ignored_for_proxy_mounts_default(self): @@ -905,8 +913,10 @@ class TestAsProxyKwarg: # Index 1 because LocalProvider is at index 0 provider = mcp._providers[1] assert isinstance(provider, FastMCPProvider) - assert len(provider._transforms) == 2 # Visibility + Namespace - assert isinstance(provider._transforms[1], Namespace) + assert ( + len(provider._transforms) == 1 + ) # Just Namespace (Visibility is in _visibility) + assert isinstance(provider._transforms[0], Namespace) assert provider.server is sub_proxy async def test_as_proxy_ignored_for_proxy_mounts_false(self): @@ -919,8 +929,10 @@ class TestAsProxyKwarg: # Index 1 because LocalProvider is at index 0 provider = mcp._providers[1] assert isinstance(provider, FastMCPProvider) - assert len(provider._transforms) == 2 # Visibility + Namespace - assert isinstance(provider._transforms[1], Namespace) + assert ( + len(provider._transforms) == 1 + ) # Just Namespace (Visibility is in _visibility) + assert isinstance(provider._transforms[0], Namespace) assert provider.server is sub_proxy async def test_as_proxy_ignored_for_proxy_mounts_true(self): @@ -933,8 +945,10 @@ class TestAsProxyKwarg: # Index 1 because LocalProvider is at index 0 provider = mcp._providers[1] assert isinstance(provider, FastMCPProvider) - assert len(provider._transforms) == 2 # Visibility + Namespace - assert isinstance(provider._transforms[1], Namespace) + assert ( + len(provider._transforms) == 1 + ) # Just Namespace (Visibility is in _visibility) + assert isinstance(provider._transforms[0], Namespace) assert provider.server is sub_proxy async def test_as_proxy_mounts_still_have_live_link(self): @@ -1169,20 +1183,24 @@ class TestCustomRouteForwarding: # LocalProvider is at index 0, mounted provider at index 1 provider1 = main_server._providers[1] assert isinstance(provider1, FastMCPProvider) - assert len(provider1._transforms) == 2 # Visibility + Namespace - assert isinstance(provider1._transforms[1], Namespace) + assert ( + len(provider1._transforms) == 1 + ) # Just Namespace (Visibility is in _visibility) + assert isinstance(provider1._transforms[0], Namespace) assert provider1.server == sub_server1 - assert provider1._transforms[1]._prefix == "sub1" + assert provider1._transforms[0]._prefix == "sub1" # Mount second server main_server.mount(sub_server2, "sub2") assert len(main_server._providers) == 3 provider2 = main_server._providers[2] assert isinstance(provider2, FastMCPProvider) - assert len(provider2._transforms) == 2 # Visibility + Namespace - assert isinstance(provider2._transforms[1], Namespace) + assert ( + len(provider2._transforms) == 1 + ) # Just Namespace (Visibility is in _visibility) + assert isinstance(provider2._transforms[0], Namespace) assert provider2.server == sub_server2 - assert provider2._transforms[1]._prefix == "sub2" + assert provider2._transforms[0]._prefix == "sub2" async def test_multiple_routes_same_server(self): """Test that multiple custom routes from same server are all included.""" diff --git a/tests/server/test_tool_transformation.py b/tests/server/test_tool_transformation.py index c41b7651d..606dcfc96 100644 --- a/tests/server/test_tool_transformation.py +++ b/tests/server/test_tool_transformation.py @@ -45,7 +45,7 @@ async def test_transformed_tool_filtering(): # Enable only tools with the enabled_tools tag mcp.enable(tags={"enabled_tools"}, only=True) - tools = await mcp.get_tools(run_middleware=True) + tools = await mcp.get_tools() # With transformation applied, the tool now has the enabled_tools tag assert len(tools) == 1 From d3ae0ed45076ab1981d8ad629ec70f3cab50c4d1 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Fri, 16 Jan 2026 22:21:40 -0500 Subject: [PATCH 2/8] Fix transform application and update tests for None return values - Override _get_tool/resource/template/prompt in FastMCP to apply server transforms over provider aggregation - Update FastMCPProvider get_* methods to check for None (not NotFoundError since get_* now returns None) - Update versioning tests to expect None instead of NotFoundError when requesting filtered/nonexistent versions --- .../server/providers/fastmcp_provider.py | 20 +- src/fastmcp/server/server.py | 250 ++++++++++++------ tests/server/test_versioning.py | 44 +-- 3 files changed, 185 insertions(+), 129 deletions(-) diff --git a/src/fastmcp/server/providers/fastmcp_provider.py b/src/fastmcp/server/providers/fastmcp_provider.py index 3d42bed9a..2b1dc8ba7 100644 --- a/src/fastmcp/server/providers/fastmcp_provider.py +++ b/src/fastmcp/server/providers/fastmcp_provider.py @@ -501,9 +501,8 @@ class FastMCPProvider(Provider): Passes the full VersionSpec to the nested server, which handles both exact version matching and range filtering. """ - try: - raw_tool = await self.server.get_tool(name, version) - except NotFoundError: + raw_tool = await self.server.get_tool(name, version) + if raw_tool is None: return None return FastMCPProviderTool.wrap(self.server, raw_tool) @@ -529,9 +528,8 @@ class FastMCPProvider(Provider): Passes the full VersionSpec to the nested server, which handles both exact version matching and range filtering. """ - try: - raw_resource = await self.server.get_resource(uri, version) - except NotFoundError: + raw_resource = await self.server.get_resource(uri, version) + if raw_resource is None: return None return FastMCPProviderResource.wrap(self.server, raw_resource) @@ -559,9 +557,8 @@ class FastMCPProvider(Provider): Passes the full VersionSpec to the nested server, which handles both exact version matching and range filtering. """ - try: - raw_template = await self.server.get_resource_template(uri, version) - except NotFoundError: + raw_template = await self.server.get_resource_template(uri, version) + if raw_template is None: return None return FastMCPProviderResourceTemplate.wrap(self.server, raw_template) @@ -587,9 +584,8 @@ class FastMCPProvider(Provider): Passes the full VersionSpec to the nested server, which handles both exact version matching and range filtering. """ - try: - raw_prompt = await self.server.get_prompt(name, version) - except NotFoundError: + raw_prompt = await self.server.get_prompt(name, version) + if raw_prompt is None: return None return FastMCPProviderPrompt.wrap(self.server, raw_prompt) diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index 0705a0929..f7463cbf3 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -1087,7 +1087,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): async def get_tool( self, name: str, version: VersionSpec | str | None = None ) -> Tool | None: - """Provider interface: aggregate tool lookup from all providers. + """Get a tool by name with all server transforms applied. Returns None if not found or if the tool is disabled via visibility settings. @@ -1104,29 +1104,49 @@ class FastMCP(Provider, Generic[LifespanResultT]): else: version_spec = version - # Query all providers in parallel for efficient lookup - results = await gather( - *[p._get_tool(name, version_spec) for p in self._providers], - return_exceptions=True, - ) + # Use _get_tool which applies server transforms + return await self._get_tool(name, version_spec) - # Collect valid results, pick highest version - valid: list[Tool] = [] - for i, result in enumerate(results): - if isinstance(result, BaseException): - if not isinstance(result, NotFoundError): - logger.debug( - f"Error during get_tool({name!r}) from provider " - f"{self._providers[i]}: {result}" - ) - continue - if result is not None and self._is_component_enabled(result): - valid.append(result) + async def _get_tool( + self, name: str, version: VersionSpec | None = None + ) -> Tool | None: + """Get tool with all transforms applied (server transforms + aggregation). - if not valid: - return None + Overrides Provider._get_tool to aggregate from providers after applying + server-level transforms. + """ - return max(valid, key=version_sort_key) + async def base(n: str, *, version: VersionSpec | None = None) -> Tool | None: + # Aggregate from all sub-providers + results = await gather( + *[p._get_tool(n, version) for p in self._providers], + return_exceptions=True, + ) + + # Collect valid results, pick highest version + valid: list[Tool] = [] + for i, result in enumerate(results): + if isinstance(result, BaseException): + if not isinstance(result, NotFoundError): + logger.debug( + f"Error during get_tool({n!r}) from provider " + f"{self._providers[i]}: {result}" + ) + continue + if result is not None and self._is_component_enabled(result): + valid.append(result) + + if not valid: + return None + + return max(valid, key=version_sort_key) + + # Build transform chain: server transforms applied over aggregation + chain = base + for transform in self._get_all_transforms(): + chain = partial(transform.get_tool, call_next=chain) + + return await chain(name, version=version) async def get_resources(self, *, run_middleware: bool = False) -> list[Resource]: """Get all enabled resources from providers. @@ -1182,7 +1202,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): async def get_resource( self, uri: str, version: VersionSpec | str | None = None ) -> Resource | None: - """Provider interface: aggregate resource lookup from all providers. + """Get a resource by URI with all server transforms applied. Returns None if not found or if the resource is disabled via visibility settings. @@ -1199,29 +1219,49 @@ class FastMCP(Provider, Generic[LifespanResultT]): else: version_spec = version - # Query all providers in parallel for efficient lookup - results = await gather( - *[p._get_resource(uri, version_spec) for p in self._providers], - return_exceptions=True, - ) + # Use _get_resource which applies server transforms + return await self._get_resource(uri, version_spec) - # Collect valid results, pick highest version - valid: list[Resource] = [] - for i, result in enumerate(results): - if isinstance(result, BaseException): - if not isinstance(result, NotFoundError): - logger.debug( - f"Error during get_resource({uri!r}) from provider " - f"{self._providers[i]}: {result}" - ) - continue - if result is not None and self._is_component_enabled(result): - valid.append(result) + async def _get_resource( + self, uri: str, version: VersionSpec | None = None + ) -> Resource | None: + """Get resource with all transforms applied (server transforms + aggregation). - if not valid: - return None + Overrides Provider._get_resource to aggregate from providers after applying + server-level transforms. + """ - return max(valid, key=version_sort_key) + async def base(u: str, *, version: VersionSpec | None = None) -> Resource | None: + # Aggregate from all sub-providers + results = await gather( + *[p._get_resource(u, version) for p in self._providers], + return_exceptions=True, + ) + + # Collect valid results, pick highest version + valid: list[Resource] = [] + for i, result in enumerate(results): + if isinstance(result, BaseException): + if not isinstance(result, NotFoundError): + logger.debug( + f"Error during get_resource({u!r}) from provider " + f"{self._providers[i]}: {result}" + ) + continue + if result is not None and self._is_component_enabled(result): + valid.append(result) + + if not valid: + return None + + return max(valid, key=version_sort_key) + + # Build transform chain: server transforms applied over aggregation + chain = base + for transform in self._get_all_transforms(): + chain = partial(transform.get_resource, call_next=chain) + + return await chain(uri, version=version) async def get_resource_templates( self, *, run_middleware: bool = False @@ -1280,7 +1320,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): async def get_resource_template( self, uri: str, version: VersionSpec | str | None = None ) -> ResourceTemplate | None: - """Provider interface: aggregate template lookup from all providers. + """Get a resource template by URI with all server transforms applied. Returns None if not found or if the template is disabled via visibility settings. @@ -1297,29 +1337,51 @@ class FastMCP(Provider, Generic[LifespanResultT]): else: version_spec = version - # Query all providers in parallel for efficient lookup - results = await gather( - *[p._get_resource_template(uri, version_spec) for p in self._providers], - return_exceptions=True, - ) + # Use _get_resource_template which applies server transforms + return await self._get_resource_template(uri, version_spec) - # Collect valid results, pick highest version - valid: list[ResourceTemplate] = [] - for i, result in enumerate(results): - if isinstance(result, BaseException): - if not isinstance(result, NotFoundError): - logger.debug( - f"Error during get_resource_template({uri!r}) from provider " - f"{self._providers[i]}: {result}" - ) - continue - if result is not None and self._is_component_enabled(result): - valid.append(result) + async def _get_resource_template( + self, uri: str, version: VersionSpec | None = None + ) -> ResourceTemplate | None: + """Get resource template with all transforms applied. - if not valid: - return None + Overrides Provider._get_resource_template to aggregate from providers after + applying server-level transforms. + """ - return max(valid, key=version_sort_key) + async def base( + u: str, *, version: VersionSpec | None = None + ) -> ResourceTemplate | None: + # Aggregate from all sub-providers + results = await gather( + *[p._get_resource_template(u, version) for p in self._providers], + return_exceptions=True, + ) + + # Collect valid results, pick highest version + valid: list[ResourceTemplate] = [] + for i, result in enumerate(results): + if isinstance(result, BaseException): + if not isinstance(result, NotFoundError): + logger.debug( + f"Error during get_resource_template({u!r}) from provider " + f"{self._providers[i]}: {result}" + ) + continue + if result is not None and self._is_component_enabled(result): + valid.append(result) + + if not valid: + return None + + return max(valid, key=version_sort_key) + + # Build transform chain: server transforms applied over aggregation + chain = base + for transform in self._get_all_transforms(): + chain = partial(transform.get_resource_template, call_next=chain) + + return await chain(uri, version=version) async def get_prompts(self, *, run_middleware: bool = False) -> list[Prompt]: """Get all enabled prompts from providers. @@ -1374,7 +1436,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): async def get_prompt( self, name: str, version: VersionSpec | str | None = None ) -> Prompt | None: - """Provider interface: aggregate prompt lookup from all providers. + """Get a prompt by name with all server transforms applied. Returns None if not found or if the prompt is disabled via visibility settings. @@ -1391,29 +1453,49 @@ class FastMCP(Provider, Generic[LifespanResultT]): else: version_spec = version - # Query all providers in parallel for efficient lookup - results = await gather( - *[p._get_prompt(name, version_spec) for p in self._providers], - return_exceptions=True, - ) + # Use _get_prompt which applies server transforms + return await self._get_prompt(name, version_spec) - # Collect valid results, pick highest version - valid: list[Prompt] = [] - for i, result in enumerate(results): - if isinstance(result, BaseException): - if not isinstance(result, NotFoundError): - logger.debug( - f"Error during get_prompt({name!r}) from provider " - f"{self._providers[i]}: {result}" - ) - continue - if result is not None and self._is_component_enabled(result): - valid.append(result) + async def _get_prompt( + self, name: str, version: VersionSpec | None = None + ) -> Prompt | None: + """Get prompt with all transforms applied (server transforms + aggregation). - if not valid: - return None + Overrides Provider._get_prompt to aggregate from providers after applying + server-level transforms. + """ - return max(valid, key=version_sort_key) + async def base(n: str, *, version: VersionSpec | None = None) -> Prompt | None: + # Aggregate from all sub-providers + results = await gather( + *[p._get_prompt(n, version) for p in self._providers], + return_exceptions=True, + ) + + # Collect valid results, pick highest version + valid: list[Prompt] = [] + for i, result in enumerate(results): + if isinstance(result, BaseException): + if not isinstance(result, NotFoundError): + logger.debug( + f"Error during get_prompt({n!r}) from provider " + f"{self._providers[i]}: {result}" + ) + continue + if result is not None and self._is_component_enabled(result): + valid.append(result) + + if not valid: + return None + + return max(valid, key=version_sort_key) + + # Build transform chain: server transforms applied over aggregation + chain = base + for transform in self._get_all_transforms(): + chain = partial(transform.get_prompt, call_next=chain) + + return await chain(name, version=version) async def get_component( self, key: str diff --git a/tests/server/test_versioning.py b/tests/server/test_versioning.py index 7742ae485..2d68fedef 100644 --- a/tests/server/test_versioning.py +++ b/tests/server/test_versioning.py @@ -487,13 +487,8 @@ class TestVersionFilter: assert tool_v2 is not None assert tool_v2.version == "2.0" - # Cannot request version outside range - import pytest - - from fastmcp.exceptions import NotFoundError - - with pytest.raises(NotFoundError): - await mcp.get_tool("add", version="1.0") + # Cannot request version outside range - returns None + assert await mcp.get_tool("add", version="1.0") is None async def test_version_range(self): """VersionFilter(version_gte='2.0', version_lt='3.0') shows only v2.x.""" @@ -529,16 +524,9 @@ class TestVersionFilter: assert tool_v2 is not None assert tool_v2.version == "2.0" - # Versions outside range are not accessible - import pytest - - from fastmcp.exceptions import NotFoundError - - with pytest.raises(NotFoundError): - await mcp.get_tool("calc", version="1.0") - - with pytest.raises(NotFoundError): - await mcp.get_tool("calc", version="3.0") + # Versions outside range are not accessible - return None + assert await mcp.get_tool("calc", version="1.0") is None + assert await mcp.get_tool("calc", version="3.0") is None async def test_unversioned_always_passes(self): """Unversioned components pass through any filter.""" @@ -602,9 +590,8 @@ class TestVersionFilter: mcp.add_transform(VersionFilter(version_lt="3.0")) - # Tool exists but is filtered out - with pytest.raises(NotFoundError): - await mcp.get_tool("only_v5") + # Tool exists but is filtered out - returns None + assert await mcp.get_tool("only_v5") is None async def test_must_specify_at_least_one(self): """VersionFilter() with no args raises ValueError.""" @@ -846,12 +833,7 @@ class TestMountedVersionFiltering: assert len(tools) == 0 # get_tool should also return None (respects filter) - import pytest - - from fastmcp.exceptions import NotFoundError - - with pytest.raises(NotFoundError): - await parent.get_tool("child_high_version_tool") + assert await parent.get_tool("child_high_version_tool") is None class TestMountedRangeFiltering: @@ -907,13 +889,9 @@ class TestMountedRangeFiltering: assert tool is not None assert tool.version == "1.0" - # Request version outside range should fail - import pytest - - from fastmcp.exceptions import NotFoundError - - with pytest.raises(NotFoundError): - await parent.get_tool("child_calc", version="3.0") + # Request version outside range should return None + result = await parent.get_tool("child_calc", version="3.0") + assert result is None class TestUnversionedExemption: From 6ba5e5d1944a1314fbd05f95187af916cfc5923d Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Sat, 17 Jan 2026 07:23:29 -0500 Subject: [PATCH 3/8] Fix type errors and add type ignore comments - Add None checks in auth and tool transform tests - Add assertions in component_service.py for None returns - Add type ignore comments for max() with version_sort_key --- .../component_manager/component_service.py | 4 ++++ .../server/providers/fastmcp_provider.py | 1 - src/fastmcp/server/server.py | 17 ++++++++++------- tests/server/auth/test_authorization.py | 1 + tests/server/test_versioning.py | 2 -- tests/tools/test_tool_transform.py | 1 + 6 files changed, 16 insertions(+), 10 deletions(-) diff --git a/src/fastmcp/contrib/component_manager/component_service.py b/src/fastmcp/contrib/component_manager/component_service.py index 105994464..cb829521a 100644 --- a/src/fastmcp/contrib/component_manager/component_service.py +++ b/src/fastmcp/contrib/component_manager/component_service.py @@ -188,10 +188,12 @@ class ComponentService: if resource_keys: self._server.enable(keys=resource_keys) resource = await self._server.get_resource(uri) + assert resource is not None # We just enabled it return resource if template_keys: self._server.enable(keys=template_keys) template = await self._server.get_resource_template(uri) + assert template is not None # We just enabled it return template # 2. Check mounted servers via FastMCPProvider @@ -234,11 +236,13 @@ class ComponentService: if resource_keys: # Get the highest version to return before disabling resource = await self._server.get_resource(uri) + assert resource is not None # Keys exist so it must exist self._server.disable(keys=resource_keys) return resource if template_keys: # Get the highest version to return before disabling template = await self._server.get_resource_template(uri) + assert template is not None # Keys exist so it must exist self._server.disable(keys=template_keys) return template diff --git a/src/fastmcp/server/providers/fastmcp_provider.py b/src/fastmcp/server/providers/fastmcp_provider.py index 2b1dc8ba7..e2d55804a 100644 --- a/src/fastmcp/server/providers/fastmcp_provider.py +++ b/src/fastmcp/server/providers/fastmcp_provider.py @@ -19,7 +19,6 @@ from typing import TYPE_CHECKING, Any, overload import mcp.types from mcp.types import AnyUrl -from fastmcp.exceptions import NotFoundError from fastmcp.prompts.prompt import Prompt, PromptResult from fastmcp.resources.resource import Resource, ResourceResult from fastmcp.resources.template import ResourceTemplate diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index f7463cbf3..dfaed6922 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -1080,7 +1080,9 @@ class FastMCP(Provider, Generic[LifespanResultT]): continue # Keep highest version per name existing = by_name.get(tool.name) - if existing is None or version_sort_key(tool) > version_sort_key(existing): + if existing is None or version_sort_key(tool) > version_sort_key( + existing + ): by_name[tool.name] = tool return list(by_name.values()) @@ -1139,7 +1141,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): if not valid: return None - return max(valid, key=version_sort_key) + return max(valid, key=version_sort_key) # type: ignore[type-var] # Build transform chain: server transforms applied over aggregation chain = base @@ -1231,7 +1233,9 @@ class FastMCP(Provider, Generic[LifespanResultT]): server-level transforms. """ - async def base(u: str, *, version: VersionSpec | None = None) -> Resource | None: + async def base( + u: str, *, version: VersionSpec | None = None + ) -> Resource | None: # Aggregate from all sub-providers results = await gather( *[p._get_resource(u, version) for p in self._providers], @@ -1254,7 +1258,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): if not valid: return None - return max(valid, key=version_sort_key) + return max(valid, key=version_sort_key) # type: ignore[type-var] # Build transform chain: server transforms applied over aggregation chain = base @@ -1374,7 +1378,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): if not valid: return None - return max(valid, key=version_sort_key) + return max(valid, key=version_sort_key) # type: ignore[type-var] # Build transform chain: server transforms applied over aggregation chain = base @@ -1488,7 +1492,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): if not valid: return None - return max(valid, key=version_sort_key) + return max(valid, key=version_sort_key) # type: ignore[type-var] # Build transform chain: server transforms applied over aggregation chain = base @@ -1523,7 +1527,6 @@ class FastMCP(Provider, Generic[LifespanResultT]): return await self._get_prompt(key[7:]) return None - @overload async def call_tool( self, diff --git a/tests/server/auth/test_authorization.py b/tests/server/auth/test_authorization.py index f137d638d..ffdaa185a 100644 --- a/tests/server/auth/test_authorization.py +++ b/tests/server/auth/test_authorization.py @@ -328,6 +328,7 @@ class TestToolLevelAuth: tok = set_token(token) try: tool = await mcp.get_tool("protected_tool") + assert tool is not None assert tool.name == "protected_tool" finally: auth_context_var.reset(tok) diff --git a/tests/server/test_versioning.py b/tests/server/test_versioning.py index 2d68fedef..ebc8e8c1c 100644 --- a/tests/server/test_versioning.py +++ b/tests/server/test_versioning.py @@ -577,9 +577,7 @@ class TestVersionFilter: async def test_get_tool_respects_filter(self): """get_tool() raises NotFoundError if highest version is filtered out.""" - import pytest - from fastmcp.exceptions import NotFoundError from fastmcp.server.transforms import VersionFilter mcp = FastMCP() diff --git a/tests/tools/test_tool_transform.py b/tests/tools/test_tool_transform.py index af81e3beb..2b12b99fe 100644 --- a/tests/tools/test_tool_transform.py +++ b/tests/tools/test_tool_transform.py @@ -806,6 +806,7 @@ class TestProxy: # when adding transformed tools to proxy servers. Needs separate investigation. add_tool = await proxy_server.get_tool("add") + assert add_tool is not None new_add_tool = Tool.from_tool( add_tool, name="add_transformed", From 5ade5a48381c2eb233d391ce2275cfbfc7427239 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Sat, 17 Jan 2026 11:35:47 -0500 Subject: [PATCH 4/8] Simplify Provider transform architecture Replace _get_all_transforms() with a .transforms property that returns [*self._transforms, self._visibility]. This cleanly separates user transforms from visibility filtering while keeping visibility applied last (outermost). Also: - AuthMiddleware now uses get_* instead of _get_* for proper component auth - Remove redundant _is_component_enabled checks (visibility is a transform) - Delete dead code (get_component method) - Add versions field to FastMCPMeta - Replace asserts with NotFoundError in component_service --- .../component_manager/component_service.py | 12 +++-- .../server/middleware/authorization.py | 15 +++--- src/fastmcp/server/providers/base.py | 29 ++++++----- src/fastmcp/server/server.py | 50 ++++--------------- src/fastmcp/utilities/components.py | 1 + tests/server/test_mount.py | 36 ++++--------- 6 files changed, 48 insertions(+), 95 deletions(-) diff --git a/src/fastmcp/contrib/component_manager/component_service.py b/src/fastmcp/contrib/component_manager/component_service.py index cb829521a..6d6e83878 100644 --- a/src/fastmcp/contrib/component_manager/component_service.py +++ b/src/fastmcp/contrib/component_manager/component_service.py @@ -188,12 +188,14 @@ class ComponentService: if resource_keys: self._server.enable(keys=resource_keys) resource = await self._server.get_resource(uri) - assert resource is not None # We just enabled it + if resource is None: + raise NotFoundError(f"Resource {uri!r} not found after enabling") return resource if template_keys: self._server.enable(keys=template_keys) template = await self._server.get_resource_template(uri) - assert template is not None # We just enabled it + if template is None: + raise NotFoundError(f"Template {uri!r} not found after enabling") return template # 2. Check mounted servers via FastMCPProvider @@ -236,13 +238,15 @@ class ComponentService: if resource_keys: # Get the highest version to return before disabling resource = await self._server.get_resource(uri) - assert resource is not None # Keys exist so it must exist + if resource is None: + raise NotFoundError(f"Resource {uri!r} not found") self._server.disable(keys=resource_keys) return resource if template_keys: # Get the highest version to return before disabling template = await self._server.get_resource_template(uri) - assert template is not None # Keys exist so it must exist + if template is None: + raise NotFoundError(f"Template {uri!r} not found") self._server.disable(keys=template_keys) return template diff --git a/src/fastmcp/server/middleware/authorization.py b/src/fastmcp/server/middleware/authorization.py index 468284458..451c509b5 100644 --- a/src/fastmcp/server/middleware/authorization.py +++ b/src/fastmcp/server/middleware/authorization.py @@ -136,8 +136,8 @@ class AuthMiddleware(Middleware): f"Authorization failed for tool '{tool_name}': missing context" ) - # Use _get_tool to resolve transformed names from MCP protocol - tool = await fastmcp.fastmcp._get_tool(tool_name) + # Resolve tool (includes component-level auth check) + tool = await fastmcp.fastmcp.get_tool(tool_name) if tool is None: raise AuthorizationError( f"Authorization failed for tool '{tool_name}': tool not found" @@ -201,11 +201,10 @@ class AuthMiddleware(Middleware): f"Authorization failed for resource '{uri}': missing context" ) - # Try concrete resource first, then template (for template-backed URIs) - # Use _get_* to resolve transformed URIs from MCP protocol - component = await fastmcp.fastmcp._get_resource(str(uri)) + # Try concrete resource first, then template (includes component-level auth check) + component = await fastmcp.fastmcp.get_resource(str(uri)) if component is None: - component = await fastmcp.fastmcp._get_resource_template(str(uri)) + component = await fastmcp.fastmcp.get_resource_template(str(uri)) if component is None: raise AuthorizationError( f"Authorization failed for resource '{uri}': resource not found" @@ -295,8 +294,8 @@ class AuthMiddleware(Middleware): f"Authorization failed for prompt '{prompt_name}': missing context" ) - # Use _get_prompt to resolve transformed names from MCP protocol - prompt = await fastmcp.fastmcp._get_prompt(prompt_name) + # Resolve prompt (includes component-level auth check) + prompt = await fastmcp.fastmcp.get_prompt(prompt_name) if prompt is None: raise AuthorizationError( f"Authorization failed for prompt '{prompt_name}': prompt not found" diff --git a/src/fastmcp/server/providers/base.py b/src/fastmcp/server/providers/base.py index 40f9eda31..2ae7e973c 100644 --- a/src/fastmcp/server/providers/base.py +++ b/src/fastmcp/server/providers/base.py @@ -65,18 +65,17 @@ class Provider: """ def __init__(self) -> None: - # Visibility is kept separate from user transforms and applied last (outermost) - # so that enable/disable works with transformed names, not raw names self._visibility = Visibility() self._transforms: list[Transform] = [] - def _get_all_transforms(self) -> list[Transform]: - """Get all transforms including visibility (applied last/outermost).""" - return [*self._transforms, self._visibility] - def __repr__(self) -> str: return f"{self.__class__.__name__}()" + @property + def transforms(self) -> list[Transform]: + """All transforms including visibility (applied last/outermost).""" + return [*self._transforms, self._visibility] + def add_transform(self, transform: Transform) -> None: """Add a transform to this provider. @@ -115,7 +114,7 @@ class Provider: return await self.list_tools() chain = base - for transform in self._get_all_transforms(): + for transform in self.transforms: chain = partial(transform.list_tools, call_next=chain) return await chain() @@ -137,7 +136,7 @@ class Provider: return await self.get_tool(n, version) chain = base - for transform in self._get_all_transforms(): + for transform in self.transforms: chain = partial(transform.get_tool, call_next=chain) return await chain(name, version=version) @@ -149,7 +148,7 @@ class Provider: return await self.list_resources() chain = base - for transform in self._get_all_transforms(): + for transform in self.transforms: chain = partial(transform.list_resources, call_next=chain) return await chain() @@ -168,7 +167,7 @@ class Provider: return await self.get_resource(u, version) chain = base - for transform in self._get_all_transforms(): + for transform in self.transforms: chain = partial(transform.get_resource, call_next=chain) return await chain(uri, version=version) @@ -180,7 +179,7 @@ class Provider: return await self.list_resource_templates() chain = base - for transform in self._get_all_transforms(): + for transform in self.transforms: chain = partial(transform.list_resource_templates, call_next=chain) return await chain() @@ -201,7 +200,7 @@ class Provider: return await self.get_resource_template(u, version) chain = base - for transform in self._get_all_transforms(): + for transform in self.transforms: chain = partial(transform.get_resource_template, call_next=chain) return await chain(uri, version=version) @@ -213,7 +212,7 @@ class Provider: return await self.list_prompts() chain = base - for transform in self._get_all_transforms(): + for transform in self.transforms: chain = partial(transform.list_prompts, call_next=chain) return await chain() @@ -232,7 +231,7 @@ class Provider: return await self.get_prompt(n, version) chain = base - for transform in self._get_all_transforms(): + for transform in self.transforms: chain = partial(transform.get_prompt, call_next=chain) return await chain(name, version=version) @@ -413,7 +412,7 @@ class Provider: templates_chain = templates_base prompts_chain = prompts_base - for transform in self._get_all_transforms(): + for transform in self.transforms: tools_chain = partial(transform.list_tools, call_next=tools_chain) resources_chain = partial( transform.list_resources, call_next=resources_chain diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index 8dc0fde40..b3b61b8a2 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -849,8 +849,6 @@ class FastMCP(Provider, Generic[LifespanResultT]): # Tool Transforms # ------------------------------------------------------------------------- - # Note: _get_all_transforms() is inherited from Provider - def _collect_list_results( self, results: list[Sequence[Any] | BaseException], operation: str ) -> list[Any]: @@ -942,7 +940,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): templates_chain = templates_base prompts_chain = prompts_base - for transform in self._get_all_transforms(): + for transform in self.transforms: tools_chain = partial(transform.list_tools, call_next=tools_chain) resources_chain = partial( transform.list_resources, call_next=resources_chain @@ -1084,10 +1082,6 @@ class FastMCP(Provider, Generic[LifespanResultT]): """ self._visibility.disable(keys=keys, tags=tags) - def _is_component_enabled(self, component: FastMCPComponent) -> bool: - """Check if a component is enabled (not in blocklist, passes allowlist).""" - return self._visibility.is_enabled(component) - async def get_tools(self, *, run_middleware: bool = False) -> list[Tool]: """Get all enabled tools from providers. @@ -1191,7 +1185,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): f"{self._providers[i]}: {result}" ) continue - if result is not None and self._is_component_enabled(result): + if result is not None: valid.append(result) if not valid: @@ -1201,7 +1195,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): # Build transform chain: server transforms applied over aggregation chain = base - for transform in self._get_all_transforms(): + for transform in self.transforms: chain = partial(transform.get_tool, call_next=chain) return await chain(name, version=version) @@ -1311,7 +1305,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): f"{self._providers[i]}: {result}" ) continue - if result is not None and self._is_component_enabled(result): + if result is not None: valid.append(result) if not valid: @@ -1321,7 +1315,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): # Build transform chain: server transforms applied over aggregation chain = base - for transform in self._get_all_transforms(): + for transform in self.transforms: chain = partial(transform.get_resource, call_next=chain) return await chain(uri, version=version) @@ -1435,7 +1429,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): f"{self._providers[i]}: {result}" ) continue - if result is not None and self._is_component_enabled(result): + if result is not None: valid.append(result) if not valid: @@ -1445,7 +1439,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): # Build transform chain: server transforms applied over aggregation chain = base - for transform in self._get_all_transforms(): + for transform in self.transforms: chain = partial(transform.get_resource_template, call_next=chain) return await chain(uri, version=version) @@ -1553,7 +1547,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): f"{self._providers[i]}: {result}" ) continue - if result is not None and self._is_component_enabled(result): + if result is not None: valid.append(result) if not valid: @@ -1563,37 +1557,11 @@ class FastMCP(Provider, Generic[LifespanResultT]): # Build transform chain: server transforms applied over aggregation chain = base - for transform in self._get_all_transforms(): + for transform in self.transforms: chain = partial(transform.get_prompt, call_next=chain) return await chain(name, version=version) - async def get_component( - self, key: str - ) -> Tool | Resource | ResourceTemplate | Prompt | None: - """Get a component by its prefixed key. - - Uses _get_* methods to apply server-level transforms (including namespace - resolution for mounted servers). Task keys use transformed names, so this - ensures proper resolution. - - Args: - key: The prefixed key (e.g., "tool:name", "resource:uri", "template:uri"). - - Returns: - The component if found, None otherwise. - """ - # Parse key and delegate to _get_* methods which apply server transforms - if key.startswith("tool:"): - return await self._get_tool(key[5:]) - elif key.startswith("resource:"): - return await self._get_resource(key[9:]) - elif key.startswith("template:"): - return await self._get_resource_template(key[9:]) - elif key.startswith("prompt:"): - return await self._get_prompt(key[7:]) - return None - @overload async def call_tool( self, diff --git a/src/fastmcp/utilities/components.py b/src/fastmcp/utilities/components.py index 53bcc17d6..99fef71e4 100644 --- a/src/fastmcp/utilities/components.py +++ b/src/fastmcp/utilities/components.py @@ -20,6 +20,7 @@ T = TypeVar("T", default=Any) class FastMCPMeta(TypedDict, total=False): tags: list[str] version: str + versions: list[str] def get_fastmcp_metadata(meta: dict[str, Any] | None) -> FastMCPMeta: diff --git a/tests/server/test_mount.py b/tests/server/test_mount.py index c0c3774bd..fbc3f6f11 100644 --- a/tests/server/test_mount.py +++ b/tests/server/test_mount.py @@ -838,9 +838,7 @@ class TestAsProxyKwarg: provider = mcp._providers[1] # With namespace, we get FastMCPProvider with a Namespace layer assert isinstance(provider, FastMCPProvider) - assert ( - len(provider._transforms) == 1 - ) # Just Namespace (Visibility is in _visibility) + assert len(provider._transforms) == 1 # Just Namespace assert isinstance(provider._transforms[0], Namespace) assert provider.server is sub @@ -854,9 +852,7 @@ class TestAsProxyKwarg: provider = mcp._providers[1] # With namespace, we get FastMCPProvider with a Namespace layer assert isinstance(provider, FastMCPProvider) - assert ( - len(provider._transforms) == 1 - ) # Just Namespace (Visibility is in _visibility) + assert len(provider._transforms) == 1 # Just Namespace assert isinstance(provider._transforms[0], Namespace) assert provider.server is sub @@ -870,9 +866,7 @@ class TestAsProxyKwarg: provider = mcp._providers[1] # With namespace, we get FastMCPProvider with a Namespace layer assert isinstance(provider, FastMCPProvider) - assert ( - len(provider._transforms) == 1 - ) # Just Namespace (Visibility is in _visibility) + assert len(provider._transforms) == 1 # Just Namespace assert isinstance(provider._transforms[0], Namespace) assert provider.server is not sub assert isinstance(provider.server, FastMCPProxy) @@ -897,9 +891,7 @@ class TestAsProxyKwarg: # Index 1 because LocalProvider is at index 0 provider = mcp._providers[1] assert isinstance(provider, FastMCPProvider) - assert ( - len(provider._transforms) == 1 - ) # Just Namespace (Visibility is in _visibility) + assert len(provider._transforms) == 1 # Just Namespace assert isinstance(provider._transforms[0], Namespace) assert provider.server is sub @@ -913,9 +905,7 @@ class TestAsProxyKwarg: # Index 1 because LocalProvider is at index 0 provider = mcp._providers[1] assert isinstance(provider, FastMCPProvider) - assert ( - len(provider._transforms) == 1 - ) # Just Namespace (Visibility is in _visibility) + assert len(provider._transforms) == 1 # Just Namespace assert isinstance(provider._transforms[0], Namespace) assert provider.server is sub_proxy @@ -929,9 +919,7 @@ class TestAsProxyKwarg: # Index 1 because LocalProvider is at index 0 provider = mcp._providers[1] assert isinstance(provider, FastMCPProvider) - assert ( - len(provider._transforms) == 1 - ) # Just Namespace (Visibility is in _visibility) + assert len(provider._transforms) == 1 # Just Namespace assert isinstance(provider._transforms[0], Namespace) assert provider.server is sub_proxy @@ -945,9 +933,7 @@ class TestAsProxyKwarg: # Index 1 because LocalProvider is at index 0 provider = mcp._providers[1] assert isinstance(provider, FastMCPProvider) - assert ( - len(provider._transforms) == 1 - ) # Just Namespace (Visibility is in _visibility) + assert len(provider._transforms) == 1 # Just Namespace assert isinstance(provider._transforms[0], Namespace) assert provider.server is sub_proxy @@ -1183,9 +1169,7 @@ class TestCustomRouteForwarding: # LocalProvider is at index 0, mounted provider at index 1 provider1 = main_server._providers[1] assert isinstance(provider1, FastMCPProvider) - assert ( - len(provider1._transforms) == 1 - ) # Just Namespace (Visibility is in _visibility) + assert len(provider1._transforms) == 1 # Just Namespace assert isinstance(provider1._transforms[0], Namespace) assert provider1.server == sub_server1 assert provider1._transforms[0]._prefix == "sub1" @@ -1195,9 +1179,7 @@ class TestCustomRouteForwarding: assert len(main_server._providers) == 3 provider2 = main_server._providers[2] assert isinstance(provider2, FastMCPProvider) - assert ( - len(provider2._transforms) == 1 - ) # Just Namespace (Visibility is in _visibility) + assert len(provider2._transforms) == 1 # Just Namespace assert isinstance(provider2._transforms[0], Namespace) assert provider2.server == sub_server2 assert provider2._transforms[0]._prefix == "sub2" From 5b24ea393daeb13774fac4e4288c507517d887cf Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Sat, 17 Jan 2026 12:15:01 -0500 Subject: [PATCH 5/8] Refactor FastMCP to use inherited _get_* methods from Provider - get_*() now does aggregation + component auth (raises AuthorizationError) - Deleted _get_*() overrides - inherited from Provider applies transforms - Simplified AuthMiddleware to global auth only - Changed version params to VersionSpec | None (not str | None) - Updated tests to use _get_*() where visibility filtering is expected --- .../server/middleware/authorization.py | 20 +- .../server/providers/fastmcp_provider.py | 2 +- src/fastmcp/server/server.py | 418 +++++++----------- src/fastmcp/server/tasks/requests.py | 13 +- tests/server/auth/test_authorization.py | 12 +- .../providers/test_local_provider_prompts.py | 28 +- tests/server/test_versioning.py | 35 +- 7 files changed, 226 insertions(+), 302 deletions(-) diff --git a/src/fastmcp/server/middleware/authorization.py b/src/fastmcp/server/middleware/authorization.py index 451c509b5..7b377b478 100644 --- a/src/fastmcp/server/middleware/authorization.py +++ b/src/fastmcp/server/middleware/authorization.py @@ -136,16 +136,16 @@ class AuthMiddleware(Middleware): f"Authorization failed for tool '{tool_name}': missing context" ) - # Resolve tool (includes component-level auth check) - tool = await fastmcp.fastmcp.get_tool(tool_name) + # Get tool (component auth is checked in get_tool, raises if unauthorized) + tool = await fastmcp.fastmcp._get_tool(tool_name) if tool is None: raise AuthorizationError( f"Authorization failed for tool '{tool_name}': tool not found" ) + # Global auth check token = get_access_token() ctx = AuthContext(token=token, component=tool) - if not run_auth_checks(self.auth, ctx): raise AuthorizationError( f"Authorization failed for tool '{tool_name}': insufficient permissions" @@ -201,18 +201,18 @@ class AuthMiddleware(Middleware): f"Authorization failed for resource '{uri}': missing context" ) - # Try concrete resource first, then template (includes component-level auth check) - component = await fastmcp.fastmcp.get_resource(str(uri)) + # Get resource/template (component auth is checked in get_*, raises if unauthorized) + component = await fastmcp.fastmcp._get_resource(str(uri)) if component is None: - component = await fastmcp.fastmcp.get_resource_template(str(uri)) + component = await fastmcp.fastmcp._get_resource_template(str(uri)) if component is None: raise AuthorizationError( f"Authorization failed for resource '{uri}': resource not found" ) + # Global auth check token = get_access_token() ctx = AuthContext(token=token, component=component) - if not run_auth_checks(self.auth, ctx): raise AuthorizationError( f"Authorization failed for resource '{uri}': insufficient permissions" @@ -294,16 +294,16 @@ class AuthMiddleware(Middleware): f"Authorization failed for prompt '{prompt_name}': missing context" ) - # Resolve prompt (includes component-level auth check) - prompt = await fastmcp.fastmcp.get_prompt(prompt_name) + # Get prompt (component auth is checked in get_prompt, raises if unauthorized) + prompt = await fastmcp.fastmcp._get_prompt(prompt_name) if prompt is None: raise AuthorizationError( f"Authorization failed for prompt '{prompt_name}': prompt not found" ) + # Global auth check token = get_access_token() ctx = AuthContext(token=token, component=prompt) - if not run_auth_checks(self.auth, ctx): raise AuthorizationError( f"Authorization failed for prompt '{prompt_name}': insufficient permissions" diff --git a/src/fastmcp/server/providers/fastmcp_provider.py b/src/fastmcp/server/providers/fastmcp_provider.py index e2d55804a..ed4f12ec4 100644 --- a/src/fastmcp/server/providers/fastmcp_provider.py +++ b/src/fastmcp/server/providers/fastmcp_provider.py @@ -628,7 +628,7 @@ class FastMCPProvider(Provider): templates_chain = templates_base prompts_chain = prompts_base - for transform in self._transforms: + for transform in self.transforms: tools_chain = partial(transform.list_tools, call_next=tools_chain) resources_chain = partial( transform.list_resources, call_next=resources_chain diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index b3b61b8a2..df8f051cb 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -356,7 +356,6 @@ class FastMCP(Provider, Generic[LifespanResultT]): ) # Local provider is always first in the provider list - # Note: _transforms is initialized by Provider.__init__() self._providers: list[Provider] = [ self._local_provider, *(providers or []), @@ -1127,78 +1126,58 @@ class FastMCP(Provider, Generic[LifespanResultT]): return _dedupe_with_versions(authorized, lambda t: t.name) async def get_tool( - self, name: str, version: VersionSpec | str | None = None + self, name: str, version: VersionSpec | None = None ) -> Tool | None: - """Get a tool by name with all server transforms applied. + """Get a tool by name via aggregation from providers. - Returns None if not found or if the tool is disabled via visibility settings. + This is the raw lookup that Provider._get_tool() wraps with transforms. + Aggregates from all sub-providers and applies component-level auth. Args: name: The tool name. - version: Version filter. Can be: - - None: returns highest version - - str: returns exact version match - - VersionSpec: returns best match within spec (highest matching) - """ - # Convert string to VersionSpec for backward compatibility - if isinstance(version, str): - version_spec: VersionSpec | None = VersionSpec(eq=version) - else: - version_spec = version + version: Version filter (None returns highest version). - # Use _get_tool which applies server transforms - tool = await self._get_tool(name, version_spec) - if tool is None: + Returns: + The tool if found and authorized, None if not found. + + Raises: + AuthorizationError: If component-level auth fails. + """ + + # Aggregate from all sub-providers (each applies their own transforms) + results = await gather( + *[p._get_tool(name, version) for p in self._providers], + return_exceptions=True, + ) + + # Collect valid results, pick highest version + valid: list[Tool] = [] + for i, result in enumerate(results): + if isinstance(result, BaseException): + if not isinstance(result, NotFoundError): + logger.debug( + f"Error during get_tool({name!r}) from provider " + f"{self._providers[i]}: {result}" + ) + continue + if result is not None: + valid.append(result) + + if not valid: return None - # Auth check + tool: Tool = max(valid, key=version_sort_key) # type: ignore[type-var] + + # Component auth - raises if unauthorized skip_auth, token = _get_auth_context() if not skip_auth and tool.auth is not None: ctx = AuthContext(token=token, component=tool) if not run_auth_checks(tool.auth, ctx): - return None + raise AuthorizationError(f"Unauthorized access to tool: {name!r}") + return tool - async def _get_tool( - self, name: str, version: VersionSpec | None = None - ) -> Tool | None: - """Get tool with all transforms applied (server transforms + aggregation). - - Overrides Provider._get_tool to aggregate from providers after applying - server-level transforms. - """ - - async def base(n: str, *, version: VersionSpec | None = None) -> Tool | None: - # Aggregate from all sub-providers - results = await gather( - *[p._get_tool(n, version) for p in self._providers], - return_exceptions=True, - ) - - # Collect valid results, pick highest version - valid: list[Tool] = [] - for i, result in enumerate(results): - if isinstance(result, BaseException): - if not isinstance(result, NotFoundError): - logger.debug( - f"Error during get_tool({n!r}) from provider " - f"{self._providers[i]}: {result}" - ) - continue - if result is not None: - valid.append(result) - - if not valid: - return None - - return max(valid, key=version_sort_key) # type: ignore[type-var] - - # Build transform chain: server transforms applied over aggregation - chain = base - for transform in self.transforms: - chain = partial(transform.get_tool, call_next=chain) - - return await chain(name, version=version) + # _get_tool is inherited from Provider - wraps get_tool() with transforms async def get_resources(self, *, run_middleware: bool = False) -> list[Resource]: """Get all enabled resources from providers. @@ -1245,80 +1224,57 @@ class FastMCP(Provider, Generic[LifespanResultT]): return _dedupe_with_versions(authorized, lambda r: str(r.uri)) async def get_resource( - self, uri: str, version: VersionSpec | str | None = None + self, uri: str, version: VersionSpec | None = None ) -> Resource | None: - """Get a resource by URI with all server transforms applied. + """Get a resource by URI via aggregation from providers. - Returns None if not found or if the resource is disabled via visibility settings. + This is the raw lookup that Provider._get_resource() wraps with transforms. + Aggregates from all sub-providers and applies component-level auth. Args: uri: The resource URI. - version: Version filter. Can be: - - None: returns highest version - - str: returns exact version match - - VersionSpec: returns best match within spec (highest matching) - """ - # Convert string to VersionSpec for backward compatibility - if isinstance(version, str): - version_spec: VersionSpec | None = VersionSpec(eq=version) - else: - version_spec = version + version: Version filter (None returns highest version). - # Use _get_resource which applies server transforms - resource = await self._get_resource(uri, version_spec) - if resource is None: + Returns: + The resource if found and authorized, None if not found. + + Raises: + AuthorizationError: If component-level auth fails. + """ + # Aggregate from all sub-providers (each applies their own transforms) + results = await gather( + *[p._get_resource(uri, version) for p in self._providers], + return_exceptions=True, + ) + + # Collect valid results, pick highest version + valid: list[Resource] = [] + for i, result in enumerate(results): + if isinstance(result, BaseException): + if not isinstance(result, NotFoundError): + logger.debug( + f"Error during get_resource({uri!r}) from provider " + f"{self._providers[i]}: {result}" + ) + continue + if result is not None: + valid.append(result) + + if not valid: return None - # Auth check + resource: Resource = max(valid, key=version_sort_key) # type: ignore[type-var] + + # Component auth - raises if unauthorized skip_auth, token = _get_auth_context() if not skip_auth and resource.auth is not None: ctx = AuthContext(token=token, component=resource) if not run_auth_checks(resource.auth, ctx): - return None + raise AuthorizationError(f"Unauthorized access to resource: {uri!r}") + return resource - async def _get_resource( - self, uri: str, version: VersionSpec | None = None - ) -> Resource | None: - """Get resource with all transforms applied (server transforms + aggregation). - - Overrides Provider._get_resource to aggregate from providers after applying - server-level transforms. - """ - - async def base( - u: str, *, version: VersionSpec | None = None - ) -> Resource | None: - # Aggregate from all sub-providers - results = await gather( - *[p._get_resource(u, version) for p in self._providers], - return_exceptions=True, - ) - - # Collect valid results, pick highest version - valid: list[Resource] = [] - for i, result in enumerate(results): - if isinstance(result, BaseException): - if not isinstance(result, NotFoundError): - logger.debug( - f"Error during get_resource({u!r}) from provider " - f"{self._providers[i]}: {result}" - ) - continue - if result is not None: - valid.append(result) - - if not valid: - return None - - return max(valid, key=version_sort_key) # type: ignore[type-var] - - # Build transform chain: server transforms applied over aggregation - chain = base - for transform in self.transforms: - chain = partial(transform.get_resource, call_next=chain) - - return await chain(uri, version=version) + # _get_resource is inherited from Provider - wraps get_resource() with transforms async def get_resource_templates( self, *, run_middleware: bool = False @@ -1369,80 +1325,57 @@ class FastMCP(Provider, Generic[LifespanResultT]): return _dedupe_with_versions(authorized, lambda t: t.uri_template) async def get_resource_template( - self, uri: str, version: VersionSpec | str | None = None + self, uri: str, version: VersionSpec | None = None ) -> ResourceTemplate | None: - """Get a resource template by URI with all server transforms applied. + """Get a resource template by URI via aggregation from providers. - Returns None if not found or if the template is disabled via visibility settings. + This is the raw lookup that Provider._get_resource_template() wraps with transforms. + Aggregates from all sub-providers and applies component-level auth. Args: uri: The template URI to match. - version: Version filter. Can be: - - None: returns highest version - - str: returns exact version match - - VersionSpec: returns best match within spec (highest matching) - """ - # Convert string to VersionSpec for backward compatibility - if isinstance(version, str): - version_spec: VersionSpec | None = VersionSpec(eq=version) - else: - version_spec = version + version: Version filter (None returns highest version). - # Use _get_resource_template which applies server transforms - template = await self._get_resource_template(uri, version_spec) - if template is None: + Returns: + The template if found and authorized, None if not found. + + Raises: + AuthorizationError: If component-level auth fails. + """ + # Aggregate from all sub-providers (each applies their own transforms) + results = await gather( + *[p._get_resource_template(uri, version) for p in self._providers], + return_exceptions=True, + ) + + # Collect valid results, pick highest version + valid: list[ResourceTemplate] = [] + for i, result in enumerate(results): + if isinstance(result, BaseException): + if not isinstance(result, NotFoundError): + logger.debug( + f"Error during get_resource_template({uri!r}) from provider " + f"{self._providers[i]}: {result}" + ) + continue + if result is not None: + valid.append(result) + + if not valid: return None - # Auth check + template: ResourceTemplate = max(valid, key=version_sort_key) # type: ignore[type-var] + + # Component auth - raises if unauthorized skip_auth, token = _get_auth_context() if not skip_auth and template.auth is not None: ctx = AuthContext(token=token, component=template) if not run_auth_checks(template.auth, ctx): - return None + raise AuthorizationError(f"Unauthorized access to template: {uri!r}") + return template - async def _get_resource_template( - self, uri: str, version: VersionSpec | None = None - ) -> ResourceTemplate | None: - """Get resource template with all transforms applied. - - Overrides Provider._get_resource_template to aggregate from providers after - applying server-level transforms. - """ - - async def base( - u: str, *, version: VersionSpec | None = None - ) -> ResourceTemplate | None: - # Aggregate from all sub-providers - results = await gather( - *[p._get_resource_template(u, version) for p in self._providers], - return_exceptions=True, - ) - - # Collect valid results, pick highest version - valid: list[ResourceTemplate] = [] - for i, result in enumerate(results): - if isinstance(result, BaseException): - if not isinstance(result, NotFoundError): - logger.debug( - f"Error during get_resource_template({u!r}) from provider " - f"{self._providers[i]}: {result}" - ) - continue - if result is not None: - valid.append(result) - - if not valid: - return None - - return max(valid, key=version_sort_key) # type: ignore[type-var] - - # Build transform chain: server transforms applied over aggregation - chain = base - for transform in self.transforms: - chain = partial(transform.get_resource_template, call_next=chain) - - return await chain(uri, version=version) + # _get_resource_template is inherited from Provider - wraps get_resource_template() with transforms async def get_prompts(self, *, run_middleware: bool = False) -> list[Prompt]: """Get all enabled prompts from providers. @@ -1489,78 +1422,57 @@ class FastMCP(Provider, Generic[LifespanResultT]): return _dedupe_with_versions(authorized, lambda p: p.name) async def get_prompt( - self, name: str, version: VersionSpec | str | None = None + self, name: str, version: VersionSpec | None = None ) -> Prompt | None: - """Get a prompt by name with all server transforms applied. + """Get a prompt by name via aggregation from providers. - Returns None if not found or if the prompt is disabled via visibility settings. + This is the raw lookup that Provider._get_prompt() wraps with transforms. + Aggregates from all sub-providers and applies component-level auth. Args: name: The prompt name. - version: Version filter. Can be: - - None: returns highest version - - str: returns exact version match - - VersionSpec: returns best match within spec (highest matching) - """ - # Convert string to VersionSpec for backward compatibility - if isinstance(version, str): - version_spec: VersionSpec | None = VersionSpec(eq=version) - else: - version_spec = version + version: Version filter (None returns highest version). - # Use _get_prompt which applies server transforms - prompt = await self._get_prompt(name, version_spec) - if prompt is None: + Returns: + The prompt if found and authorized, None if not found. + + Raises: + AuthorizationError: If component-level auth fails. + """ + # Aggregate from all sub-providers (each applies their own transforms) + results = await gather( + *[p._get_prompt(name, version) for p in self._providers], + return_exceptions=True, + ) + + # Collect valid results, pick highest version + valid: list[Prompt] = [] + for i, result in enumerate(results): + if isinstance(result, BaseException): + if not isinstance(result, NotFoundError): + logger.debug( + f"Error during get_prompt({name!r}) from provider " + f"{self._providers[i]}: {result}" + ) + continue + if result is not None: + valid.append(result) + + if not valid: return None - # Auth check + prompt: Prompt = max(valid, key=version_sort_key) # type: ignore[type-var] + + # Component auth - raises if unauthorized skip_auth, token = _get_auth_context() if not skip_auth and prompt.auth is not None: ctx = AuthContext(token=token, component=prompt) if not run_auth_checks(prompt.auth, ctx): - return None + raise AuthorizationError(f"Unauthorized access to prompt: {name!r}") + return prompt - async def _get_prompt( - self, name: str, version: VersionSpec | None = None - ) -> Prompt | None: - """Get prompt with all transforms applied (server transforms + aggregation). - - Overrides Provider._get_prompt to aggregate from providers after applying - server-level transforms. - """ - - async def base(n: str, *, version: VersionSpec | None = None) -> Prompt | None: - # Aggregate from all sub-providers - results = await gather( - *[p._get_prompt(n, version) for p in self._providers], - return_exceptions=True, - ) - - # Collect valid results, pick highest version - valid: list[Prompt] = [] - for i, result in enumerate(results): - if isinstance(result, BaseException): - if not isinstance(result, NotFoundError): - logger.debug( - f"Error during get_prompt({n!r}) from provider " - f"{self._providers[i]}: {result}" - ) - continue - if result is not None: - valid.append(result) - - if not valid: - return None - - return max(valid, key=version_sort_key) # type: ignore[type-var] - - # Build transform chain: server transforms applied over aggregation - chain = base - for transform in self.transforms: - chain = partial(transform.get_prompt, call_next=chain) - - return await chain(name, version=version) + # _get_prompt is inherited from Provider - wraps get_prompt() with transforms @overload async def call_tool( @@ -1568,7 +1480,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): name: str, arguments: dict[str, Any] | None = None, *, - version: str | None = None, + version: VersionSpec | None = None, run_middleware: bool = True, task_meta: None = None, ) -> ToolResult: ... @@ -1579,7 +1491,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): name: str, arguments: dict[str, Any] | None = None, *, - version: str | None = None, + version: VersionSpec | None = None, run_middleware: bool = True, task_meta: TaskMeta, ) -> mcp.types.CreateTaskResult: ... @@ -1589,7 +1501,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): name: str, arguments: dict[str, Any] | None = None, *, - version: str | None = None, + version: VersionSpec | None = None, run_middleware: bool = True, task_meta: TaskMeta | None = None, ) -> ToolResult | mcp.types.CreateTaskResult: @@ -1643,10 +1555,11 @@ class FastMCP(Provider, Generic[LifespanResultT]): ) # Core logic: find and execute tool (providers queried in parallel) + # Use _get_tool to apply transforms (including visibility) with server_span( f"tools/call {name}", "tools/call", self.name, "tool", name ) as span: - tool = await self.get_tool(name, version=version) + tool = await self._get_tool(name, version=version) if tool is None: raise NotFoundError(f"Unknown tool: {name!r}") span.set_attributes(tool.get_span_attributes()) @@ -1671,7 +1584,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): self, uri: str, *, - version: str | None = None, + version: VersionSpec | None = None, run_middleware: bool = True, task_meta: None = None, ) -> ResourceResult: ... @@ -1681,7 +1594,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): self, uri: str, *, - version: str | None = None, + version: VersionSpec | None = None, run_middleware: bool = True, task_meta: TaskMeta, ) -> mcp.types.CreateTaskResult: ... @@ -1690,7 +1603,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): self, uri: str, *, - version: str | None = None, + version: VersionSpec | None = None, run_middleware: bool = True, task_meta: TaskMeta | None = None, ) -> ResourceResult | mcp.types.CreateTaskResult: @@ -1752,8 +1665,8 @@ class FastMCP(Provider, Generic[LifespanResultT]): uri, resource_uri=uri, ) as span: - # Try concrete resources first (auth checked in get_resource) - resource = await self.get_resource(uri, version=version) + # Try concrete resources first (transforms + auth via _get_resource) + resource = await self._get_resource(uri, version=version) if resource is not None: span.set_attributes(resource.get_span_attributes()) if task_meta is not None and task_meta.fn_key is None: @@ -1773,8 +1686,8 @@ class FastMCP(Provider, Generic[LifespanResultT]): f"Error reading resource {uri!r}: {e}" ) from e - # Try templates (auth checked in get_resource_template) - template = await self.get_resource_template(uri, version=version) + # Try templates (transforms + auth via _get_resource_template) + template = await self._get_resource_template(uri, version=version) if template is None: if version is None: raise NotFoundError(f"Unknown resource: {uri!r}") @@ -1803,7 +1716,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): name: str, arguments: dict[str, Any] | None = None, *, - version: str | None = None, + version: VersionSpec | None = None, run_middleware: bool = True, task_meta: None = None, ) -> PromptResult: ... @@ -1814,7 +1727,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): name: str, arguments: dict[str, Any] | None = None, *, - version: str | None = None, + version: VersionSpec | None = None, run_middleware: bool = True, task_meta: TaskMeta, ) -> mcp.types.CreateTaskResult: ... @@ -1824,7 +1737,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): name: str, arguments: dict[str, Any] | None = None, *, - version: str | None = None, + version: VersionSpec | None = None, run_middleware: bool = True, task_meta: TaskMeta | None = None, ) -> PromptResult | mcp.types.CreateTaskResult: @@ -1874,10 +1787,11 @@ class FastMCP(Provider, Generic[LifespanResultT]): ) # Core logic: find and render prompt (providers queried in parallel) + # Use _get_prompt to apply transforms (including visibility) with server_span( f"prompts/get {name}", "prompts/get", self.name, "prompt", name ) as span: - prompt = await self.get_prompt(name, version=version) + prompt = await self._get_prompt(name, version=version) if prompt is None: raise NotFoundError(f"Unknown prompt: {name!r}") span.set_attributes(prompt.get_span_attributes()) diff --git a/src/fastmcp/server/tasks/requests.py b/src/fastmcp/server/tasks/requests.py index ff23f2f6c..61286d831 100644 --- a/src/fastmcp/server/tasks/requests.py +++ b/src/fastmcp/server/tasks/requests.py @@ -31,6 +31,7 @@ from fastmcp.resources.template import ResourceTemplate from fastmcp.server.tasks.config import DEFAULT_POLL_INTERVAL_MS, DEFAULT_TTL_MS from fastmcp.server.tasks.keys import parse_task_key from fastmcp.tools.tool import Tool +from fastmcp.utilities.versions import VersionSpec if TYPE_CHECKING: from fastmcp.server.server import FastMCP @@ -313,16 +314,20 @@ async def tasks_result_handler(server: FastMCP, params: dict[str, Any]) -> Any: component: Tool | Resource | ResourceTemplate | Prompt | None = None try: if component_key.startswith("tool:"): - name, version = _parse_key_version(component_key[5:]) + name, version_str = _parse_key_version(component_key[5:]) + version = VersionSpec(eq=version_str) if version_str else None component = await server.get_tool(name, version) elif component_key.startswith("resource:"): - uri, version = _parse_key_version(component_key[9:]) + uri, version_str = _parse_key_version(component_key[9:]) + version = VersionSpec(eq=version_str) if version_str else None component = await server.get_resource(uri, version) elif component_key.startswith("template:"): - uri, version = _parse_key_version(component_key[9:]) + uri, version_str = _parse_key_version(component_key[9:]) + version = VersionSpec(eq=version_str) if version_str else None component = await server.get_resource_template(uri, version) elif component_key.startswith("prompt:"): - name, version = _parse_key_version(component_key[7:]) + name, version_str = _parse_key_version(component_key[7:]) + version = VersionSpec(eq=version_str) if version_str else None component = await server.get_prompt(name, version) except NotFoundError: component = None diff --git a/tests/server/auth/test_authorization.py b/tests/server/auth/test_authorization.py index 79d4e5e3d..6a64dff3f 100644 --- a/tests/server/auth/test_authorization.py +++ b/tests/server/auth/test_authorization.py @@ -301,17 +301,19 @@ class TestToolLevelAuth: finally: auth_context_var.reset(tok) - async def test_get_tool_returns_none_without_auth(self): - """get_tool() checks auth and returns None for unauthorized tools.""" + async def test_get_tool_raises_without_auth(self): + """get_tool() checks auth and raises AuthorizationError for unauthorized tools.""" + from fastmcp.exceptions import AuthorizationError + mcp = FastMCP() @mcp.tool(auth=require_auth) def protected_tool() -> str: return "protected" - # get_tool() returns None for unauthorized tools - tool = await mcp.get_tool("protected_tool") - assert tool is None + # get_tool() raises AuthorizationError for unauthorized tools + with pytest.raises(AuthorizationError, match="Unauthorized access to tool"): + await mcp.get_tool("protected_tool") async def test_get_tool_returns_tool_with_auth(self): mcp = FastMCP() diff --git a/tests/server/providers/test_local_provider_prompts.py b/tests/server/providers/test_local_provider_prompts.py index 0445401bb..666ecb62d 100644 --- a/tests/server/providers/test_local_provider_prompts.py +++ b/tests/server/providers/test_local_provider_prompts.py @@ -351,10 +351,6 @@ class TestPromptEnabled: prompts = await mcp.get_prompts() assert len(prompts) == 0 - # get_prompt() returns None for disabled prompts (Provider interface) - prompt = await mcp.get_prompt("sample_prompt") - assert prompt is None - async def test_prompt_toggle_enabled(self): mcp = FastMCP() @@ -381,8 +377,8 @@ class TestPromptEnabled: prompts = await mcp.get_prompts() assert len(prompts) == 0 - # get_prompt() returns None for disabled prompts (Provider interface) - prompt = await mcp.get_prompt("sample_prompt") + # _get_prompt() applies visibility transform, returns None for disabled + prompt = await mcp._get_prompt("sample_prompt") assert prompt is None async def test_get_prompt_and_disable(self): @@ -392,15 +388,15 @@ class TestPromptEnabled: def sample_prompt() -> str: return "Hello, world!" - prompt = await mcp.get_prompt("sample_prompt") + prompt = await mcp._get_prompt("sample_prompt") assert prompt is not None mcp.disable(keys=["prompt:sample_prompt@"]) prompts = await mcp.get_prompts() assert len(prompts) == 0 - # get_prompt() returns None for disabled prompts (Provider interface) - prompt = await mcp.get_prompt("sample_prompt") + # _get_prompt() applies visibility transform, returns None for disabled + prompt = await mcp._get_prompt("sample_prompt") assert prompt is None async def test_cant_get_disabled_prompt(self): @@ -412,8 +408,8 @@ class TestPromptEnabled: mcp.disable(keys=["prompt:sample_prompt@"]) - # get_prompt() returns None for disabled prompts (Provider interface) - prompt = await mcp.get_prompt("sample_prompt") + # _get_prompt() applies visibility transform, returns None for disabled + prompt = await mcp._get_prompt("sample_prompt") assert prompt is None @@ -458,18 +454,20 @@ class TestPromptTags: async def test_read_prompt_includes_tags(self): mcp = self.create_server(include_tags={"a"}) - prompt = await mcp.get_prompt("prompt_1") + # _get_prompt applies visibility transform (tag filtering) + prompt = await mcp._get_prompt("prompt_1") result = await prompt.render({}) assert result.messages[0].content.text == "1" - prompt = await mcp.get_prompt("prompt_2") + prompt = await mcp._get_prompt("prompt_2") assert prompt is None async def test_read_prompt_excludes_tags(self): mcp = self.create_server(exclude_tags={"a"}) - prompt = await mcp.get_prompt("prompt_1") + # _get_prompt applies visibility transform (tag filtering) + prompt = await mcp._get_prompt("prompt_1") assert prompt is None - prompt = await mcp.get_prompt("prompt_2") + prompt = await mcp._get_prompt("prompt_2") result = await prompt.render({}) assert result.messages[0].content.text == "2" diff --git a/tests/server/test_versioning.py b/tests/server/test_versioning.py index a6e451c81..d6c61f78d 100644 --- a/tests/server/test_versioning.py +++ b/tests/server/test_versioning.py @@ -8,6 +8,7 @@ from mcp.types import TextContent from fastmcp import FastMCP from fastmcp.utilities.versions import ( VersionKey, + VersionSpec, compare_versions, is_version_greater, ) @@ -377,7 +378,7 @@ class TestMountedServerVersioning: assert tool.version == "2.0" # Get specific version - tool_v1 = await parent.get_tool("child_calc", version="1.0") + tool_v1 = await parent.get_tool("child_calc", VersionSpec(eq="1.0")) assert tool_v1 is not None assert tool_v1.version == "1.0" @@ -483,12 +484,12 @@ class TestVersionFilter: assert tools[0].version == "3.0" # Can request specific versions in range - tool_v2 = await mcp.get_tool("add", version="2.0") + tool_v2 = await mcp.get_tool("add", VersionSpec(eq="2.0")) assert tool_v2 is not None assert tool_v2.version == "2.0" # Cannot request version outside range - returns None - assert await mcp.get_tool("add", version="1.0") is None + assert await mcp.get_tool("add", VersionSpec(eq="1.0")) is None async def test_version_range(self): """VersionFilter(version_gte='2.0', version_lt='3.0') shows only v2.x.""" @@ -520,13 +521,13 @@ class TestVersionFilter: assert tools[0].version == "2.5" # Can request specific versions in range - tool_v2 = await mcp.get_tool("calc", version="2.0") + tool_v2 = await mcp.get_tool("calc", VersionSpec(eq="2.0")) assert tool_v2 is not None assert tool_v2.version == "2.0" # Versions outside range are not accessible - return None - assert await mcp.get_tool("calc", version="1.0") is None - assert await mcp.get_tool("calc", version="3.0") is None + assert await mcp.get_tool("calc", VersionSpec(eq="1.0")) is None + assert await mcp.get_tool("calc", VersionSpec(eq="3.0")) is None async def test_unversioned_always_passes(self): """Unversioned components pass through any filter.""" @@ -576,7 +577,7 @@ class TestVersionFilter: assert tools[0].version == "2025-01-01" async def test_get_tool_respects_filter(self): - """get_tool() raises NotFoundError if highest version is filtered out.""" + """get_tool() returns None if highest version is filtered out.""" from fastmcp.server.transforms import VersionFilter @@ -883,12 +884,12 @@ class TestMountedRangeFiltering: parent.add_transform(VersionFilter(version_gte="1.0", version_lt="3.0")) # Request specific version within range - tool = await parent.get_tool("child_calc", version="1.0") + tool = await parent.get_tool("child_calc", VersionSpec(eq="1.0")) assert tool is not None assert tool.version == "1.0" # Request version outside range should return None - result = await parent.get_tool("child_calc", version="3.0") + result = await parent.get_tool("child_calc", VersionSpec(eq="3.0")) assert result is None @@ -930,7 +931,7 @@ class TestUnversionedExemption: # Even with explicit version request, unversioned tool is returned # (it's the only version that exists, and unversioned matches any spec) - tool = await mcp.get_tool("my_tool", version="1.0") + tool = await mcp.get_tool("my_tool", VersionSpec(eq="1.0")) assert tool is not None assert tool.version is None @@ -1066,12 +1067,16 @@ class TestVersionedCalls: assert result.structured_content["result"] == 12 # Explicit v1.0 (addition) - result = await mcp.call_tool("calculate", {"x": 3, "y": 4}, version="1.0") + result = await mcp.call_tool( + "calculate", {"x": 3, "y": 4}, version=VersionSpec(eq="1.0") + ) assert result.structured_content is not None assert result.structured_content["result"] == 7 # Explicit v2.0 (multiplication) - result = await mcp.call_tool("calculate", {"x": 3, "y": 4}, version="2.0") + result = await mcp.call_tool( + "calculate", {"x": 3, "y": 4}, version=VersionSpec(eq="2.0") + ) assert result.structured_content is not None assert result.structured_content["result"] == 12 @@ -1092,7 +1097,7 @@ class TestVersionedCalls: assert result.contents[0].content == "config v2" # Explicit v1.0 - result = await mcp.read_resource("data://config", version="1.0") + result = await mcp.read_resource("data://config", version=VersionSpec(eq="1.0")) assert result.contents[0].content == "config v1" async def test_render_prompt_with_version(self): @@ -1113,7 +1118,7 @@ class TestVersionedCalls: assert isinstance(content, TextContent) and content.text == "Hello from v2" # Explicit v1.0 - result = await mcp.render_prompt("greet", version="1.0") + result = await mcp.render_prompt("greet", version=VersionSpec(eq="1.0")) content = result.messages[0].content assert isinstance(content, TextContent) and content.text == "Hello from v1" @@ -1130,7 +1135,7 @@ class TestVersionedCalls: return "v1" with pytest.raises(NotFoundError): - await mcp.call_tool("mytool", {}, version="999.0") + await mcp.call_tool("mytool", {}, version=VersionSpec(eq="999.0")) class TestClientVersionSelection: From 78ab933dc530a060af3694abf1c25dd7cb851898 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Sat, 17 Jan 2026 12:37:37 -0500 Subject: [PATCH 6/8] Fix transform application in FastMCPProvider and MCP handlers FastMCPProvider now calls _get_* methods instead of get_* to ensure nested server transforms are applied during lookups. Also converts string versions to VersionSpec in MCP handlers. --- .../server/providers/fastmcp_provider.py | 20 +++++++----- src/fastmcp/server/server.py | 15 +++++---- tests/server/test_versioning.py | 31 ++++++++++--------- 3 files changed, 37 insertions(+), 29 deletions(-) diff --git a/src/fastmcp/server/providers/fastmcp_provider.py b/src/fastmcp/server/providers/fastmcp_provider.py index ed4f12ec4..4ad388ebd 100644 --- a/src/fastmcp/server/providers/fastmcp_provider.py +++ b/src/fastmcp/server/providers/fastmcp_provider.py @@ -498,9 +498,10 @@ class FastMCPProvider(Provider): """Get a tool by name as a FastMCPProviderTool. Passes the full VersionSpec to the nested server, which handles both - exact version matching and range filtering. + exact version matching and range filtering. Uses _get_tool to ensure + the nested server's transforms are applied. """ - raw_tool = await self.server.get_tool(name, version) + raw_tool = await self.server._get_tool(name, version) if raw_tool is None: return None return FastMCPProviderTool.wrap(self.server, raw_tool) @@ -525,9 +526,10 @@ class FastMCPProvider(Provider): """Get a concrete resource by URI as a FastMCPProviderResource. Passes the full VersionSpec to the nested server, which handles both - exact version matching and range filtering. + exact version matching and range filtering. Uses _get_resource to ensure + the nested server's transforms are applied. """ - raw_resource = await self.server.get_resource(uri, version) + raw_resource = await self.server._get_resource(uri, version) if raw_resource is None: return None return FastMCPProviderResource.wrap(self.server, raw_resource) @@ -554,9 +556,10 @@ class FastMCPProvider(Provider): """Get a resource template that matches the given URI. Passes the full VersionSpec to the nested server, which handles both - exact version matching and range filtering. + exact version matching and range filtering. Uses _get_resource_template + to ensure the nested server's transforms are applied. """ - raw_template = await self.server.get_resource_template(uri, version) + raw_template = await self.server._get_resource_template(uri, version) if raw_template is None: return None return FastMCPProviderResourceTemplate.wrap(self.server, raw_template) @@ -581,9 +584,10 @@ class FastMCPProvider(Provider): """Get a prompt by name as a FastMCPProviderPrompt. Passes the full VersionSpec to the nested server, which handles both - exact version matching and range filtering. + exact version matching and range filtering. Uses _get_prompt to ensure + the nested server's transforms are applied. """ - raw_prompt = await self.server.get_prompt(name, version) + raw_prompt = await self.server._get_prompt(name, version) if raw_prompt is None: return None return FastMCPProviderPrompt.wrap(self.server, raw_prompt) diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index df8f051cb..7054cd2dd 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -1982,14 +1982,14 @@ class FastMCP(Provider, Generic[LifespanResultT]): try: # Extract version and task metadata from request context. # fn_key is set by call_tool() after finding the tool. - version: str | None = None + version_str: str | None = None task_meta: TaskMeta | None = None try: ctx = self._mcp_server.request_context # Extract version from request-level _meta.fastmcp.version if ctx.meta: meta_dict = ctx.meta.model_dump(exclude_none=True) - version = meta_dict.get("fastmcp", {}).get("version") + version_str = meta_dict.get("fastmcp", {}).get("version") # Extract SEP-1686 task metadata if ctx.experimental.is_task: mcp_task_meta = ctx.experimental.task_metadata @@ -1998,6 +1998,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): except (AttributeError, LookupError): pass + version = VersionSpec(eq=version_str) if version_str else None result = await self.call_tool( key, arguments, version=version, task_meta=task_meta ) @@ -2030,7 +2031,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): try: # Extract version and task metadata from request context. - version: str | None = None + version_str: str | None = None task_meta: TaskMeta | None = None try: ctx = self._mcp_server.request_context @@ -2038,7 +2039,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): if ctx.meta: meta_dict = ctx.meta.model_dump(exclude_none=True) fastmcp_meta = meta_dict.get("fastmcp") or {} - version = fastmcp_meta.get("version") + version_str = fastmcp_meta.get("version") # Extract SEP-1686 task metadata if ctx.experimental.is_task: mcp_task_meta = ctx.experimental.task_metadata @@ -2047,6 +2048,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): except (AttributeError, LookupError): pass + version = VersionSpec(eq=version_str) if version_str else None result = await self.read_resource( str(uri), version=version, task_meta=task_meta ) @@ -2082,14 +2084,14 @@ class FastMCP(Provider, Generic[LifespanResultT]): try: # Extract version and task metadata from request context. # fn_key is set by render_prompt() after finding the prompt. - version: str | None = None + version_str: str | None = None task_meta: TaskMeta | None = None try: ctx = self._mcp_server.request_context # Extract version from request-level _meta.fastmcp.version if ctx.meta: meta_dict = ctx.meta.model_dump(exclude_none=True) - version = meta_dict.get("fastmcp", {}).get("version") + version_str = meta_dict.get("fastmcp", {}).get("version") # Extract SEP-1686 task metadata if ctx.experimental.is_task: mcp_task_meta = ctx.experimental.task_metadata @@ -2098,6 +2100,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): except (AttributeError, LookupError): pass + version = VersionSpec(eq=version_str) if version_str else None result = await self.render_prompt( name, arguments, version=version, task_meta=task_meta ) diff --git a/tests/server/test_versioning.py b/tests/server/test_versioning.py index d6c61f78d..f2b0713dc 100644 --- a/tests/server/test_versioning.py +++ b/tests/server/test_versioning.py @@ -483,13 +483,13 @@ class TestVersionFilter: assert len(tools) == 1 assert tools[0].version == "3.0" - # Can request specific versions in range - tool_v2 = await mcp.get_tool("add", VersionSpec(eq="2.0")) + # Can request specific versions in range (use _get_tool to apply transforms) + tool_v2 = await mcp._get_tool("add", VersionSpec(eq="2.0")) assert tool_v2 is not None assert tool_v2.version == "2.0" # Cannot request version outside range - returns None - assert await mcp.get_tool("add", VersionSpec(eq="1.0")) is None + assert await mcp._get_tool("add", VersionSpec(eq="1.0")) is None async def test_version_range(self): """VersionFilter(version_gte='2.0', version_lt='3.0') shows only v2.x.""" @@ -520,14 +520,14 @@ class TestVersionFilter: assert len(tools) == 1 assert tools[0].version == "2.5" - # Can request specific versions in range - tool_v2 = await mcp.get_tool("calc", VersionSpec(eq="2.0")) + # Can request specific versions in range (use _get_tool to apply transforms) + tool_v2 = await mcp._get_tool("calc", VersionSpec(eq="2.0")) assert tool_v2 is not None assert tool_v2.version == "2.0" # Versions outside range are not accessible - return None - assert await mcp.get_tool("calc", VersionSpec(eq="1.0")) is None - assert await mcp.get_tool("calc", VersionSpec(eq="3.0")) is None + assert await mcp._get_tool("calc", VersionSpec(eq="1.0")) is None + assert await mcp._get_tool("calc", VersionSpec(eq="3.0")) is None async def test_unversioned_always_passes(self): """Unversioned components pass through any filter.""" @@ -589,8 +589,8 @@ class TestVersionFilter: mcp.add_transform(VersionFilter(version_lt="3.0")) - # Tool exists but is filtered out - returns None - assert await mcp.get_tool("only_v5") is None + # Tool exists but is filtered out - returns None (use _get_tool to apply transforms) + assert await mcp._get_tool("only_v5") is None async def test_must_specify_at_least_one(self): """VersionFilter() with no args raises ValueError.""" @@ -831,8 +831,8 @@ class TestMountedVersionFiltering: tools = await parent.get_tools() assert len(tools) == 0 - # get_tool should also return None (respects filter) - assert await parent.get_tool("child_high_version_tool") is None + # _get_tool should also return None (respects filter, applies transforms) + assert await parent._get_tool("child_high_version_tool") is None class TestMountedRangeFiltering: @@ -857,7 +857,8 @@ class TestMountedRangeFiltering: parent.add_transform(VersionFilter(version_lt="2.0")) # Should return v1.0 (the highest version that matches <2.0) - tool = await parent.get_tool("child_calc") + # Use _get_tool to apply transforms + tool = await parent._get_tool("child_calc") assert tool is not None assert tool.version == "1.0" @@ -883,13 +884,13 @@ class TestMountedRangeFiltering: parent.mount(child, "child") parent.add_transform(VersionFilter(version_gte="1.0", version_lt="3.0")) - # Request specific version within range - tool = await parent.get_tool("child_calc", VersionSpec(eq="1.0")) + # Request specific version within range (use _get_tool to apply transforms) + tool = await parent._get_tool("child_calc", VersionSpec(eq="1.0")) assert tool is not None assert tool.version == "1.0" # Request version outside range should return None - result = await parent.get_tool("child_calc", VersionSpec(eq="3.0")) + result = await parent._get_tool("child_calc", VersionSpec(eq="3.0")) assert result is None From 6dd1de62e4d369d9fa2066d2a658af73ca5b2802 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Sat, 17 Jan 2026 14:20:10 -0500 Subject: [PATCH 7/8] Address PR review feedback - Filter task-eligible components in FastMCPProvider.get_tasks() - Catch AuthorizationError in get_* methods and return None for consistency - Remove AggregateProvider from top-level exports --- src/fastmcp/server/providers/__init__.py | 2 - .../server/providers/fastmcp_provider.py | 13 +++-- src/fastmcp/server/server.py | 56 +++++++++---------- 3 files changed, 37 insertions(+), 34 deletions(-) diff --git a/src/fastmcp/server/providers/__init__.py b/src/fastmcp/server/providers/__init__.py index 8d997e5ff..39ea4d0ff 100644 --- a/src/fastmcp/server/providers/__init__.py +++ b/src/fastmcp/server/providers/__init__.py @@ -27,7 +27,6 @@ Example: from typing import TYPE_CHECKING -from fastmcp.server.providers.aggregate import AggregateProvider from fastmcp.server.providers.base import Provider from fastmcp.server.providers.fastmcp_provider import FastMCPProvider from fastmcp.server.providers.filesystem import FileSystemProvider @@ -38,7 +37,6 @@ if TYPE_CHECKING: from fastmcp.server.providers.proxy import ProxyProvider as ProxyProvider __all__ = [ - "AggregateProvider", "FastMCPProvider", "FileSystemProvider", "LocalProvider", diff --git a/src/fastmcp/server/providers/fastmcp_provider.py b/src/fastmcp/server/providers/fastmcp_provider.py index 4ad388ebd..a435bc0ac 100644 --- a/src/fastmcp/server/providers/fastmcp_provider.py +++ b/src/fastmcp/server/providers/fastmcp_provider.py @@ -642,11 +642,16 @@ class FastMCPProvider(Provider): ) prompts_chain = partial(transform.list_prompts, call_next=prompts_chain) + # Filter to only task-eligible components (same as base Provider) return [ - *await tools_chain(), - *await resources_chain(), - *await templates_chain(), - *await prompts_chain(), + c + for c in [ + *await tools_chain(), + *await resources_chain(), + *await templates_chain(), + *await prompts_chain(), + ] + if c.task_config.supports_tasks() ] # ------------------------------------------------------------------------- diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index 7054cd2dd..7917bfdc6 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -1138,10 +1138,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): version: Version filter (None returns highest version). Returns: - The tool if found and authorized, None if not found. - - Raises: - AuthorizationError: If component-level auth fails. + The tool if found and authorized, None if not found or unauthorized. """ # Aggregate from all sub-providers (each applies their own transforms) @@ -1168,12 +1165,15 @@ class FastMCP(Provider, Generic[LifespanResultT]): tool: Tool = max(valid, key=version_sort_key) # type: ignore[type-var] - # Component auth - raises if unauthorized + # Component auth - return None if unauthorized (consistent with list filtering) skip_auth, token = _get_auth_context() if not skip_auth and tool.auth is not None: ctx = AuthContext(token=token, component=tool) - if not run_auth_checks(tool.auth, ctx): - raise AuthorizationError(f"Unauthorized access to tool: {name!r}") + try: + if not run_auth_checks(tool.auth, ctx): + return None + except AuthorizationError: + return None return tool @@ -1236,10 +1236,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): version: Version filter (None returns highest version). Returns: - The resource if found and authorized, None if not found. - - Raises: - AuthorizationError: If component-level auth fails. + The resource if found and authorized, None if not found or unauthorized. """ # Aggregate from all sub-providers (each applies their own transforms) results = await gather( @@ -1265,12 +1262,15 @@ class FastMCP(Provider, Generic[LifespanResultT]): resource: Resource = max(valid, key=version_sort_key) # type: ignore[type-var] - # Component auth - raises if unauthorized + # Component auth - return None if unauthorized (consistent with list filtering) skip_auth, token = _get_auth_context() if not skip_auth and resource.auth is not None: ctx = AuthContext(token=token, component=resource) - if not run_auth_checks(resource.auth, ctx): - raise AuthorizationError(f"Unauthorized access to resource: {uri!r}") + try: + if not run_auth_checks(resource.auth, ctx): + return None + except AuthorizationError: + return None return resource @@ -1337,10 +1337,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): version: Version filter (None returns highest version). Returns: - The template if found and authorized, None if not found. - - Raises: - AuthorizationError: If component-level auth fails. + The template if found and authorized, None if not found or unauthorized. """ # Aggregate from all sub-providers (each applies their own transforms) results = await gather( @@ -1366,12 +1363,15 @@ class FastMCP(Provider, Generic[LifespanResultT]): template: ResourceTemplate = max(valid, key=version_sort_key) # type: ignore[type-var] - # Component auth - raises if unauthorized + # Component auth - return None if unauthorized (consistent with list filtering) skip_auth, token = _get_auth_context() if not skip_auth and template.auth is not None: ctx = AuthContext(token=token, component=template) - if not run_auth_checks(template.auth, ctx): - raise AuthorizationError(f"Unauthorized access to template: {uri!r}") + try: + if not run_auth_checks(template.auth, ctx): + return None + except AuthorizationError: + return None return template @@ -1434,10 +1434,7 @@ class FastMCP(Provider, Generic[LifespanResultT]): version: Version filter (None returns highest version). Returns: - The prompt if found and authorized, None if not found. - - Raises: - AuthorizationError: If component-level auth fails. + The prompt if found and authorized, None if not found or unauthorized. """ # Aggregate from all sub-providers (each applies their own transforms) results = await gather( @@ -1463,12 +1460,15 @@ class FastMCP(Provider, Generic[LifespanResultT]): prompt: Prompt = max(valid, key=version_sort_key) # type: ignore[type-var] - # Component auth - raises if unauthorized + # Component auth - return None if unauthorized (consistent with list filtering) skip_auth, token = _get_auth_context() if not skip_auth and prompt.auth is not None: ctx = AuthContext(token=token, component=prompt) - if not run_auth_checks(prompt.auth, ctx): - raise AuthorizationError(f"Unauthorized access to prompt: {name!r}") + try: + if not run_auth_checks(prompt.auth, ctx): + return None + except AuthorizationError: + return None return prompt From 3d988629953483bd7ff1b4b812adcd4c15d00e90 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Sat, 17 Jan 2026 14:25:29 -0500 Subject: [PATCH 8/8] Fix auth test to expect None instead of AuthorizationError --- tests/server/auth/test_authorization.py | 12 +++++------- 1 file changed, 5 insertions(+), 7 deletions(-) diff --git a/tests/server/auth/test_authorization.py b/tests/server/auth/test_authorization.py index 6a64dff3f..e0de6aa15 100644 --- a/tests/server/auth/test_authorization.py +++ b/tests/server/auth/test_authorization.py @@ -301,19 +301,17 @@ class TestToolLevelAuth: finally: auth_context_var.reset(tok) - async def test_get_tool_raises_without_auth(self): - """get_tool() checks auth and raises AuthorizationError for unauthorized tools.""" - from fastmcp.exceptions import AuthorizationError - + async def test_get_tool_returns_none_without_auth(self): + """get_tool() returns None for unauthorized tools (consistent with list filtering).""" mcp = FastMCP() @mcp.tool(auth=require_auth) def protected_tool() -> str: return "protected" - # get_tool() raises AuthorizationError for unauthorized tools - with pytest.raises(AuthorizationError, match="Unauthorized access to tool"): - await mcp.get_tool("protected_tool") + # get_tool() returns None for unauthorized tools + tool = await mcp.get_tool("protected_tool") + assert tool is None async def test_get_tool_returns_tool_with_auth(self): mcp = FastMCP()