mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-24 06:24:18 +02:00
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:
parent
66d713b367
commit
d3ae0ed450
3 changed files with 185 additions and 129 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue