diff --git a/src/fastmcp/contrib/component_manager/component_service.py b/src/fastmcp/contrib/component_manager/component_service.py index 105994464..6d6e83878 100644 --- a/src/fastmcp/contrib/component_manager/component_service.py +++ b/src/fastmcp/contrib/component_manager/component_service.py @@ -188,10 +188,14 @@ class ComponentService: if resource_keys: self._server.enable(keys=resource_keys) resource = await self._server.get_resource(uri) + 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) + if template is None: + raise NotFoundError(f"Template {uri!r} not found after enabling") return template # 2. Check mounted servers via FastMCPProvider @@ -234,11 +238,15 @@ class ComponentService: if resource_keys: # Get the highest version to return before disabling resource = await self._server.get_resource(uri) + 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) + 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 f9f69ecd1..7b377b478 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,11 +136,16 @@ class AuthMiddleware(Middleware): f"Authorization failed for tool '{tool_name}': missing context" ) - 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" @@ -196,15 +201,18 @@ class AuthMiddleware(Middleware): f"Authorization failed for resource '{uri}': missing context" ) - # 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)) + # 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)) + 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" @@ -286,11 +294,16 @@ class AuthMiddleware(Middleware): f"Authorization failed for prompt '{prompt_name}': missing context" ) - 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/aggregate.py b/src/fastmcp/server/providers/aggregate.py index f080e58d5..2f4548043 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 @@ -28,12 +39,15 @@ 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 highest version is returned. Errors from individual providers are logged and skipped (graceful degradation). + + 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 3e5dce745..2ae7e973c 100644 --- a/src/fastmcp/server/providers/base.py +++ b/src/fastmcp/server/providers/base.py @@ -65,13 +65,17 @@ class Provider: """ def __init__(self) -> None: - # Visibility is the first (innermost) transform - closest to the base provider self._visibility = Visibility() - self._transforms: list[Transform] = [self._visibility] + self._transforms: list[Transform] = [] 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. @@ -110,7 +114,7 @@ class Provider: return await self.list_tools() chain = base - for transform in self._transforms: + for transform in self.transforms: chain = partial(transform.list_tools, call_next=chain) return await chain() @@ -132,7 +136,7 @@ class Provider: return await self.get_tool(n, version) chain = base - for transform in self._transforms: + for transform in self.transforms: chain = partial(transform.get_tool, call_next=chain) return await chain(name, version=version) @@ -144,7 +148,7 @@ class Provider: return await self.list_resources() chain = base - for transform in self._transforms: + for transform in self.transforms: chain = partial(transform.list_resources, call_next=chain) return await chain() @@ -163,7 +167,7 @@ class Provider: return await self.get_resource(u, version) chain = base - for transform in self._transforms: + for transform in self.transforms: chain = partial(transform.get_resource, call_next=chain) return await chain(uri, version=version) @@ -175,7 +179,7 @@ class Provider: return await self.list_resource_templates() chain = base - for transform in self._transforms: + for transform in self.transforms: chain = partial(transform.list_resource_templates, call_next=chain) return await chain() @@ -196,7 +200,7 @@ class Provider: return await self.get_resource_template(u, version) chain = base - for transform in self._transforms: + for transform in self.transforms: chain = partial(transform.get_resource_template, call_next=chain) return await chain(uri, version=version) @@ -208,7 +212,7 @@ class Provider: return await self.list_prompts() chain = base - for transform in self._transforms: + for transform in self.transforms: chain = partial(transform.list_prompts, call_next=chain) return await chain() @@ -227,7 +231,7 @@ class Provider: return await self.get_prompt(n, version) chain = base - for transform in self._transforms: + for transform in self.transforms: chain = partial(transform.get_prompt, call_next=chain) return await chain(name, version=version) @@ -402,13 +406,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.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 b7a56ab9a..a435bc0ac 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 @@ -486,9 +485,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,11 +498,11 @@ 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. """ - 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) @@ -514,9 +513,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] @@ -527,11 +526,11 @@ 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. """ - 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) @@ -542,6 +541,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. """ @@ -556,11 +556,11 @@ 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. """ - 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) @@ -571,6 +571,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. """ @@ -583,11 +584,11 @@ 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. """ - 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) @@ -599,12 +600,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)] @@ -631,7 +632,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 @@ -641,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 5fa10c11d..7917bfdc6 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 @@ -282,7 +281,7 @@ class StateValue(FastMCPBaseModel): value: Any -class FastMCP(Generic[LifespanResultT]): +class FastMCP(Provider, Generic[LifespanResultT]): def __init__( self, name: str | None = None, @@ -323,6 +322,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, @@ -353,9 +355,6 @@ 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 self._providers: list[Provider] = [ self._local_provider, @@ -409,10 +408,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.", @@ -581,11 +577,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: @@ -853,181 +848,72 @@ class FastMCP(Generic[LifespanResultT]): # Tool Transforms # ------------------------------------------------------------------------- - def _get_root_provider(self) -> AggregateProvider: - """Get the root provider (aggregate of all providers). - - 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). + async def list_tools(self) -> Sequence[Tool]: + """Aggregate tools from all sub-providers. - Returns all versions of all tools. Caller is responsible for deduplication - if showing to clients (keeping highest version per name). + This is the Provider interface implementation. The inherited _list_tools() + applies server-level transforms over this method. """ - root = self._get_root_provider() + results = await gather( + *[p._list_tools() for p in self._providers], + return_exceptions=True, + ) + return self._collect_list_results(results, "list_tools") - async def base() -> Sequence[Tool]: - return await root.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") - chain = base - for transform in self._get_all_transforms(): - chain = partial(transform.list_tools, call_next=chain) + 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") - return await chain() + 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 _source_get_tool( - self, name: str, version: VersionSpec | None = None - ) -> Tool | None: - """Get tool by name with all transforms applied (provider + server). + async def get_tasks(self) -> Sequence[FastMCPComponent]: + """Get task-eligible components with all transforms applied. - Args: - name: The tool name. - version: Optional version filter. If None, returns highest version. + Overrides Provider.get_tasks() to collect task-eligible components + from all sub-providers and apply server-level transforms. """ - root = self._get_root_provider() - - async def base(n: str, version: VersionSpec | None = None) -> Tool | None: - return await root._get_tool(n, version) - - 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 _source_list_resources(self) -> Sequence[Resource]: - """List resources with all transforms applied (provider + server). - - Returns all versions of all resources. Caller is responsible for deduplication - if showing to clients (keeping highest version per URI). - """ - 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, version: VersionSpec | None = None - ) -> Resource | None: - """Get resource by URI with all transforms applied (provider + server). - - Args: - uri: The resource URI. - version: Optional version filter. If None, returns highest version. - """ - root = self._get_root_provider() - - async def base(u: str, version: VersionSpec | None = None) -> Resource | None: - return await root._get_resource(u, version) - - 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 _source_list_resource_templates(self) -> Sequence[ResourceTemplate]: - """List resource templates with all transforms applied (provider + server). - - Returns all versions of all templates. Caller is responsible for deduplication - if showing to clients (keeping highest version per uri_template). - """ - 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, version: VersionSpec | None = None - ) -> ResourceTemplate | None: - """Get resource template by URI with all transforms applied (provider + server). - - Args: - uri: The template URI to match. - version: Optional version filter. If None, returns highest version. - """ - root = self._get_root_provider() - - async def base( - u: str, version: VersionSpec | None = None - ) -> ResourceTemplate | None: - return await root._get_resource_template(u, version) - - 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 _source_list_prompts(self) -> Sequence[Prompt]: - """List prompts with all transforms applied (provider + server). - - Returns all versions of all prompts. Caller is responsible for deduplication - if showing to clients (keeping highest version per name). - """ - 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, version: VersionSpec | None = None - ) -> Prompt | None: - """Get prompt by name with all transforms applied (provider + server). - - Args: - name: The prompt name. - version: Optional version filter. If None, returns highest version. - """ - root = self._get_root_provider() - - async def base(n: str, version: VersionSpec | None = None) -> Prompt | None: - return await root._get_prompt(n, version) - - 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 _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. - """ - root = self._get_root_provider() - components = list(await root.get_tasks()) + 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)] @@ -1053,7 +939,7 @@ class FastMCP(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 @@ -1195,10 +1081,6 @@ class FastMCP(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. @@ -1206,90 +1088,97 @@ 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() - # Filter by auth - authorized: list[Tool] = [] - for tool in tools: - 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): + # Filter by auth + authorized: list[Tool] = [] + for tool in tools: + 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: continue - except AuthorizationError: - continue - authorized.append(tool) + authorized.append(tool) - return _dedupe_with_versions(authorized, lambda t: t.name) + return _dedupe_with_versions(authorized, lambda t: t.name) async def get_tool( - self, name: str, version: VersionSpec | str | None = None - ) -> Tool: - """Get an enabled tool by name. + self, name: str, version: VersionSpec | None = None + ) -> Tool | None: + """Get a tool by name via aggregation from providers. - Queries providers with full transform chain (provider transforms + server transforms + visibility). - Returns only if enabled and authorized. + 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) + version: Version filter (None returns highest version). + + Returns: + The tool if found and authorized, None if not found or unauthorized. """ - # Convert string to VersionSpec for backward compatibility - if isinstance(version, str): - version_spec: VersionSpec | None = VersionSpec(eq=version) - else: - version_spec = version - tool = await self._source_get_tool(name, version_spec) + # 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, + ) - if tool is None: - if version is None: - raise NotFoundError(f"Unknown tool: {name!r}") - elif isinstance(version, str): - raise NotFoundError(f"Unknown tool: {name!r} version {version!r}") - else: - raise NotFoundError(f"Unknown tool: {name!r} matching {version!r}") + # 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) - # Check tool-level auth (skip for STDIO) + if not valid: + return None + + tool: Tool = max(valid, key=version_sort_key) # type: ignore[type-var] + + # 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): - if version is None: - raise NotFoundError(f"Unknown tool: {name!r}") - elif isinstance(version, str): - raise NotFoundError(f"Unknown tool: {name!r} version {version!r}") - else: - raise NotFoundError(f"Unknown tool: {name!r} matching {version!r}") + try: + if not run_auth_checks(tool.auth, ctx): + return None + except AuthorizationError: + return None return tool + # _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. @@ -1297,92 +1186,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={}, # 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() - # Filter by auth - authorized: list[Resource] = [] - for resource in resources: - 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): + # Filter by auth + authorized: list[Resource] = [] + for resource in resources: + 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: continue - except AuthorizationError: - continue - authorized.append(resource) + authorized.append(resource) - return _dedupe_with_versions(authorized, lambda r: str(r.uri)) + return _dedupe_with_versions(authorized, lambda r: str(r.uri)) async def get_resource( - self, uri: str, version: VersionSpec | str | None = None - ) -> Resource: - """Get an enabled resource by URI. + self, uri: str, version: VersionSpec | None = None + ) -> Resource | None: + """Get a resource by URI via aggregation from providers. - Queries providers with full transform chain (provider transforms + server transforms + visibility). - Returns only if enabled and authorized. + 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) + version: Version filter (None returns highest version). + + Returns: + The resource if found and authorized, None if not found or unauthorized. """ - # Convert string to VersionSpec for backward compatibility - if isinstance(version, str): - version_spec: VersionSpec | None = VersionSpec(eq=version) - else: - version_spec = version + # 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, + ) - resource = await self._source_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: + valid.append(result) - if resource is None: - if version is None: - raise NotFoundError(f"Unknown resource: {uri}") - elif isinstance(version, str): - raise NotFoundError(f"Unknown resource: {uri} version {version!r}") - else: - raise NotFoundError(f"Unknown resource: {uri} matching {version!r}") + if not valid: + return None - # Check resource-level auth (skip for STDIO) + resource: Resource = max(valid, key=version_sort_key) # type: ignore[type-var] + + # 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): - if version is None: - raise NotFoundError(f"Unknown resource: {uri}") - elif isinstance(version, str): - raise NotFoundError(f"Unknown resource: {uri} version {version!r}") - else: - raise NotFoundError(f"Unknown resource: {uri} matching {version!r}") + try: + if not run_auth_checks(resource.auth, ctx): + return None + except AuthorizationError: + return None return resource + # _get_resource is inherited from Provider - wraps get_resource() with transforms + async def get_resource_templates( self, *, run_middleware: bool = False ) -> list[ResourceTemplate]: @@ -1392,100 +1285,98 @@ 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() - # Filter by auth - authorized: list[ResourceTemplate] = [] - for template in templates: - 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): + # Filter by auth + authorized: list[ResourceTemplate] = [] + for template in templates: + 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: continue - except AuthorizationError: - continue - authorized.append(template) + authorized.append(template) - return _dedupe_with_versions(authorized, lambda t: t.uri_template) + return _dedupe_with_versions(authorized, lambda t: t.uri_template) async def get_resource_template( - self, uri: str, version: VersionSpec | str | None = None - ) -> ResourceTemplate: - """Get an enabled resource template that matches the given URI. + self, uri: str, version: VersionSpec | None = None + ) -> ResourceTemplate | None: + """Get a resource template by URI via aggregation from providers. - Queries providers with full transform chain (provider transforms + server transforms + visibility). - Returns only if enabled and authorized. + 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) + version: Version filter (None returns highest version). + + Returns: + The template if found and authorized, None if not found or unauthorized. """ - # Convert string to VersionSpec for backward compatibility - if isinstance(version, str): - version_spec: VersionSpec | None = VersionSpec(eq=version) - else: - version_spec = version + # 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, + ) - template = await self._source_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: + valid.append(result) - if template is None: - if version is None: - raise NotFoundError(f"Unknown resource template: {uri}") - elif isinstance(version, str): - raise NotFoundError( - f"Unknown resource template: {uri} version {version!r}" - ) - else: - raise NotFoundError( - f"Unknown resource template: {uri} matching {version!r}" - ) + if not valid: + return None - # Check template-level auth (skip for STDIO) + template: ResourceTemplate = max(valid, key=version_sort_key) # type: ignore[type-var] + + # 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): - if version is None: - raise NotFoundError(f"Unknown resource template: {uri}") - elif isinstance(version, str): - raise NotFoundError( - f"Unknown resource template: {uri} version {version!r}" - ) - else: - raise NotFoundError( - f"Unknown resource template: {uri} matching {version!r}" - ) + try: + if not run_auth_checks(template.auth, ctx): + return None + except AuthorizationError: + return None return template + # _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. @@ -1493,101 +1384,103 @@ 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() - # Filter by auth - authorized: list[Prompt] = [] - for prompt in prompts: - 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): + # Filter by auth + authorized: list[Prompt] = [] + for prompt in prompts: + 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: continue - except AuthorizationError: - continue - authorized.append(prompt) + authorized.append(prompt) - return _dedupe_with_versions(authorized, lambda p: p.name) + return _dedupe_with_versions(authorized, lambda p: p.name) async def get_prompt( - self, name: str, version: VersionSpec | str | None = None - ) -> Prompt: - """Get an enabled prompt by name. + self, name: str, version: VersionSpec | None = None + ) -> Prompt | None: + """Get a prompt by name via aggregation from providers. - Queries providers with full transform chain (provider transforms + server transforms + visibility). - Returns only if enabled and authorized. + 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) + version: Version filter (None returns highest version). + + Returns: + The prompt if found and authorized, None if not found or unauthorized. """ - # Convert string to VersionSpec for backward compatibility - if isinstance(version, str): - version_spec: VersionSpec | None = VersionSpec(eq=version) - else: - version_spec = version + # 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, + ) - prompt = await self._source_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: + valid.append(result) - if prompt is None: - if version is None: - raise NotFoundError(f"Unknown prompt: {name}") - elif isinstance(version, str): - raise NotFoundError(f"Unknown prompt: {name!r} version {version!r}") - else: - raise NotFoundError(f"Unknown prompt: {name!r} matching {version!r}") + if not valid: + return None - # Check prompt-level auth (skip for STDIO) + prompt: Prompt = max(valid, key=version_sort_key) # type: ignore[type-var] + + # 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): - if version is None: - raise NotFoundError(f"Unknown prompt: {name}") - elif isinstance(version, str): - raise NotFoundError(f"Unknown prompt: {name!r} version {version!r}") - else: - raise NotFoundError( - f"Unknown prompt: {name!r} matching {version!r}" - ) + try: + if not run_auth_checks(prompt.auth, ctx): + return None + except AuthorizationError: + return None return prompt + # _get_prompt is inherited from Provider - wraps get_prompt() with transforms + @overload async def call_tool( self, 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: ... @@ -1598,7 +1491,7 @@ class FastMCP(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: ... @@ -1608,7 +1501,7 @@ class FastMCP(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: @@ -1662,10 +1555,13 @@ class FastMCP(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()) if task_meta is not None and task_meta.fn_key is None: task_meta = replace(task_meta, fn_key=tool.key) @@ -1688,7 +1584,7 @@ class FastMCP(Generic[LifespanResultT]): self, uri: str, *, - version: str | None = None, + version: VersionSpec | None = None, run_middleware: bool = True, task_meta: None = None, ) -> ResourceResult: ... @@ -1698,7 +1594,7 @@ class FastMCP(Generic[LifespanResultT]): self, uri: str, *, - version: str | None = None, + version: VersionSpec | None = None, run_middleware: bool = True, task_meta: TaskMeta, ) -> mcp.types.CreateTaskResult: ... @@ -1707,7 +1603,7 @@ class FastMCP(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: @@ -1769,33 +1665,35 @@ class FastMCP(Generic[LifespanResultT]): uri, resource_uri=uri, ) as span: - # Try concrete resources first - try: - 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: 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 + try: + 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 - try: - template = await self.get_resource_template(uri, version=version) - except NotFoundError: + # 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}") from None + raise NotFoundError(f"Unknown resource: {uri!r}") raise NotFoundError( f"Unknown resource: {uri!r} version {version!r}" - ) from None + ) span.set_attributes(template.get_span_attributes()) params = template.matches(uri) assert params is not None @@ -1818,7 +1716,7 @@ class FastMCP(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: ... @@ -1829,7 +1727,7 @@ class FastMCP(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: ... @@ -1839,7 +1737,7 @@ class FastMCP(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: @@ -1889,10 +1787,13 @@ class FastMCP(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()) if task_meta is not None and task_meta.fn_key is None: task_meta = replace(task_meta, fn_key=prompt.key) @@ -2007,8 +1908,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, @@ -2023,8 +1934,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, @@ -2061,14 +1982,14 @@ class FastMCP(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 @@ -2077,6 +1998,7 @@ class FastMCP(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 ) @@ -2109,7 +2031,7 @@ class FastMCP(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 @@ -2117,7 +2039,7 @@ class FastMCP(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 @@ -2126,6 +2048,7 @@ class FastMCP(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 ) @@ -2161,14 +2084,14 @@ class FastMCP(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 @@ -2177,6 +2100,7 @@ class FastMCP(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/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/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/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..e0de6aa15 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,17 @@ 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_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" - with pytest.raises(NotFoundError): - 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() @@ -323,6 +324,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) @@ -334,6 +336,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 +350,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 +363,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 +379,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 +388,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 +407,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 +415,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 34937c594..666ecb62d 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,9 +351,6 @@ class TestPromptEnabled: prompts = await mcp.get_prompts() assert len(prompts) == 0 - with pytest.raises(NotFoundError, match="Unknown prompt"): - await mcp.get_prompt("sample_prompt") - async def test_prompt_toggle_enabled(self): mcp = FastMCP() @@ -381,8 +377,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() 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): mcp = FastMCP() @@ -391,15 +388,16 @@ 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 - with pytest.raises(NotFoundError, match="Unknown 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): mcp = FastMCP() @@ -410,8 +408,9 @@ class TestPromptEnabled: mcp.disable(keys=["prompt:sample_prompt@"]) - with pytest.raises(NotFoundError, match="Unknown 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 class TestPromptTags: @@ -455,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" - 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") + # _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_mount.py b/tests/server/test_mount.py index 29b0ad42b..fbc3f6f11 100644 --- a/tests/server/test_mount.py +++ b/tests/server/test_mount.py @@ -838,8 +838,8 @@ 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 + assert isinstance(provider._transforms[0], Namespace) assert provider.server is sub async def test_as_proxy_false(self): @@ -852,8 +852,8 @@ 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 + assert isinstance(provider._transforms[0], Namespace) assert provider.server is sub async def test_as_proxy_true(self): @@ -866,8 +866,8 @@ 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 + assert isinstance(provider._transforms[0], Namespace) assert provider.server is not sub assert isinstance(provider.server, FastMCPProxy) @@ -891,8 +891,8 @@ 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 + assert isinstance(provider._transforms[0], Namespace) assert provider.server is sub async def test_as_proxy_ignored_for_proxy_mounts_default(self): @@ -905,8 +905,8 @@ 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 + 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 +919,8 @@ 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 + 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 +933,8 @@ 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 + 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 +1169,20 @@ 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 + 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 + 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 diff --git a/tests/server/test_versioning.py b/tests/server/test_versioning.py index a7bf34461..f2b0713dc 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" @@ -482,18 +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", version="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 - 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", 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.""" @@ -524,21 +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", version="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 - 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", 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.""" @@ -588,10 +577,8 @@ 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.""" - import pytest + """get_tool() returns None if highest version is filtered out.""" - from fastmcp.exceptions import NotFoundError from fastmcp.server.transforms import VersionFilter mcp = FastMCP() @@ -602,9 +589,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 (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.""" @@ -845,13 +831,8 @@ class TestMountedVersionFiltering: tools = await parent.get_tools() 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") + # _get_tool should also return None (respects filter, applies transforms) + assert await parent._get_tool("child_high_version_tool") is None class TestMountedRangeFiltering: @@ -876,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" @@ -902,18 +884,14 @@ 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", version="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 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", VersionSpec(eq="3.0")) + assert result is None class TestUnversionedExemption: @@ -954,7 +932,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 @@ -1090,12 +1068,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 @@ -1116,7 +1098,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): @@ -1137,7 +1119,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" @@ -1154,7 +1136,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: 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",