Merge pull request #2901 from jlowin/refactor-provider-inheritance-v2

This commit is contained in:
Jeremiah Lowin 2026-01-17 14:40:22 -05:00 committed by GitHub
commit 5d12afdbf6
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
15 changed files with 601 additions and 634 deletions

View file

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

View file

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

View file

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

View file

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

View file

@ -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()
]
# -------------------------------------------------------------------------

File diff suppressed because it is too large Load diff

View file

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

View file

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

View file

@ -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 = []

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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