Fix transform application and update tests for None return values

- Override _get_tool/resource/template/prompt in FastMCP to apply
  server transforms over provider aggregation
- Update FastMCPProvider get_* methods to check for None (not
  NotFoundError since get_* now returns None)
- Update versioning tests to expect None instead of NotFoundError
  when requesting filtered/nonexistent versions
This commit is contained in:
Jeremiah Lowin 2026-01-16 22:21:40 -05:00
commit d3ae0ed450
3 changed files with 185 additions and 129 deletions

View file

@ -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)

View file

@ -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

View file

@ -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: