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: