Simplify Provider transform architecture

Replace _get_all_transforms() with a .transforms property that returns
[*self._transforms, self._visibility]. This cleanly separates user transforms
from visibility filtering while keeping visibility applied last (outermost).

Also:
- AuthMiddleware now uses get_* instead of _get_* for proper component auth
- Remove redundant _is_component_enabled checks (visibility is a transform)
- Delete dead code (get_component method)
- Add versions field to FastMCPMeta
- Replace asserts with NotFoundError in component_service
This commit is contained in:
Jeremiah Lowin 2026-01-17 11:35:47 -05:00
commit 5ade5a4838
6 changed files with 48 additions and 95 deletions

View file

@ -188,12 +188,14 @@ class ComponentService:
if resource_keys:
self._server.enable(keys=resource_keys)
resource = await self._server.get_resource(uri)
assert resource is not None # We just enabled it
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)
assert template is not None # We just enabled it
if template is None:
raise NotFoundError(f"Template {uri!r} not found after enabling")
return template
# 2. Check mounted servers via FastMCPProvider
@ -236,13 +238,15 @@ class ComponentService:
if resource_keys:
# Get the highest version to return before disabling
resource = await self._server.get_resource(uri)
assert resource is not None # Keys exist so it must exist
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)
assert template is not None # Keys exist so it must exist
if template is None:
raise NotFoundError(f"Template {uri!r} not found")
self._server.disable(keys=template_keys)
return template

View file

@ -136,8 +136,8 @@ class AuthMiddleware(Middleware):
f"Authorization failed for tool '{tool_name}': missing context"
)
# Use _get_tool to resolve transformed names from MCP protocol
tool = await fastmcp.fastmcp._get_tool(tool_name)
# Resolve tool (includes component-level auth check)
tool = await fastmcp.fastmcp.get_tool(tool_name)
if tool is None:
raise AuthorizationError(
f"Authorization failed for tool '{tool_name}': tool not found"
@ -201,11 +201,10 @@ class AuthMiddleware(Middleware):
f"Authorization failed for resource '{uri}': missing context"
)
# Try concrete resource first, then template (for template-backed URIs)
# Use _get_* to resolve transformed URIs from MCP protocol
component = await fastmcp.fastmcp._get_resource(str(uri))
# Try concrete resource first, then template (includes component-level auth check)
component = await fastmcp.fastmcp.get_resource(str(uri))
if component is None:
component = await fastmcp.fastmcp._get_resource_template(str(uri))
component = await fastmcp.fastmcp.get_resource_template(str(uri))
if component is None:
raise AuthorizationError(
f"Authorization failed for resource '{uri}': resource not found"
@ -295,8 +294,8 @@ class AuthMiddleware(Middleware):
f"Authorization failed for prompt '{prompt_name}': missing context"
)
# Use _get_prompt to resolve transformed names from MCP protocol
prompt = await fastmcp.fastmcp._get_prompt(prompt_name)
# Resolve prompt (includes component-level auth check)
prompt = await fastmcp.fastmcp.get_prompt(prompt_name)
if prompt is None:
raise AuthorizationError(
f"Authorization failed for prompt '{prompt_name}': prompt not found"

View file

@ -65,18 +65,17 @@ class Provider:
"""
def __init__(self) -> None:
# Visibility is kept separate from user transforms and applied last (outermost)
# so that enable/disable works with transformed names, not raw names
self._visibility = Visibility()
self._transforms: list[Transform] = []
def _get_all_transforms(self) -> list[Transform]:
"""Get all transforms including visibility (applied last/outermost)."""
return [*self._transforms, self._visibility]
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.
@ -115,7 +114,7 @@ class Provider:
return await self.list_tools()
chain = base
for transform in self._get_all_transforms():
for transform in self.transforms:
chain = partial(transform.list_tools, call_next=chain)
return await chain()
@ -137,7 +136,7 @@ class Provider:
return await self.get_tool(n, version)
chain = base
for transform in self._get_all_transforms():
for transform in self.transforms:
chain = partial(transform.get_tool, call_next=chain)
return await chain(name, version=version)
@ -149,7 +148,7 @@ class Provider:
return await self.list_resources()
chain = base
for transform in self._get_all_transforms():
for transform in self.transforms:
chain = partial(transform.list_resources, call_next=chain)
return await chain()
@ -168,7 +167,7 @@ class Provider:
return await self.get_resource(u, version)
chain = base
for transform in self._get_all_transforms():
for transform in self.transforms:
chain = partial(transform.get_resource, call_next=chain)
return await chain(uri, version=version)
@ -180,7 +179,7 @@ class Provider:
return await self.list_resource_templates()
chain = base
for transform in self._get_all_transforms():
for transform in self.transforms:
chain = partial(transform.list_resource_templates, call_next=chain)
return await chain()
@ -201,7 +200,7 @@ class Provider:
return await self.get_resource_template(u, version)
chain = base
for transform in self._get_all_transforms():
for transform in self.transforms:
chain = partial(transform.get_resource_template, call_next=chain)
return await chain(uri, version=version)
@ -213,7 +212,7 @@ class Provider:
return await self.list_prompts()
chain = base
for transform in self._get_all_transforms():
for transform in self.transforms:
chain = partial(transform.list_prompts, call_next=chain)
return await chain()
@ -232,7 +231,7 @@ class Provider:
return await self.get_prompt(n, version)
chain = base
for transform in self._get_all_transforms():
for transform in self.transforms:
chain = partial(transform.get_prompt, call_next=chain)
return await chain(name, version=version)
@ -413,7 +412,7 @@ class Provider:
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

View file

@ -849,8 +849,6 @@ class FastMCP(Provider, Generic[LifespanResultT]):
# Tool Transforms
# -------------------------------------------------------------------------
# Note: _get_all_transforms() is inherited from Provider
def _collect_list_results(
self, results: list[Sequence[Any] | BaseException], operation: str
) -> list[Any]:
@ -942,7 +940,7 @@ class FastMCP(Provider, 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
@ -1084,10 +1082,6 @@ class FastMCP(Provider, 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.
@ -1191,7 +1185,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
f"{self._providers[i]}: {result}"
)
continue
if result is not None and self._is_component_enabled(result):
if result is not None:
valid.append(result)
if not valid:
@ -1201,7 +1195,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
# Build transform chain: server transforms applied over aggregation
chain = base
for transform in self._get_all_transforms():
for transform in self.transforms:
chain = partial(transform.get_tool, call_next=chain)
return await chain(name, version=version)
@ -1311,7 +1305,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
f"{self._providers[i]}: {result}"
)
continue
if result is not None and self._is_component_enabled(result):
if result is not None:
valid.append(result)
if not valid:
@ -1321,7 +1315,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
# Build transform chain: server transforms applied over aggregation
chain = base
for transform in self._get_all_transforms():
for transform in self.transforms:
chain = partial(transform.get_resource, call_next=chain)
return await chain(uri, version=version)
@ -1435,7 +1429,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
f"{self._providers[i]}: {result}"
)
continue
if result is not None and self._is_component_enabled(result):
if result is not None:
valid.append(result)
if not valid:
@ -1445,7 +1439,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
# Build transform chain: server transforms applied over aggregation
chain = base
for transform in self._get_all_transforms():
for transform in self.transforms:
chain = partial(transform.get_resource_template, call_next=chain)
return await chain(uri, version=version)
@ -1553,7 +1547,7 @@ class FastMCP(Provider, Generic[LifespanResultT]):
f"{self._providers[i]}: {result}"
)
continue
if result is not None and self._is_component_enabled(result):
if result is not None:
valid.append(result)
if not valid:
@ -1563,37 +1557,11 @@ class FastMCP(Provider, Generic[LifespanResultT]):
# Build transform chain: server transforms applied over aggregation
chain = base
for transform in self._get_all_transforms():
for transform in self.transforms:
chain = partial(transform.get_prompt, call_next=chain)
return await chain(name, version=version)
async def get_component(
self, key: str
) -> Tool | Resource | ResourceTemplate | Prompt | None:
"""Get a component by its prefixed key.
Uses _get_* methods to apply server-level transforms (including namespace
resolution for mounted servers). Task keys use transformed names, so this
ensures proper resolution.
Args:
key: The prefixed key (e.g., "tool:name", "resource:uri", "template:uri").
Returns:
The component if found, None otherwise.
"""
# Parse key and delegate to _get_* methods which apply server transforms
if key.startswith("tool:"):
return await self._get_tool(key[5:])
elif key.startswith("resource:"):
return await self._get_resource(key[9:])
elif key.startswith("template:"):
return await self._get_resource_template(key[9:])
elif key.startswith("prompt:"):
return await self._get_prompt(key[7:])
return None
@overload
async def call_tool(
self,

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

@ -838,9 +838,7 @@ class TestAsProxyKwarg:
provider = mcp._providers[1]
# With namespace, we get FastMCPProvider with a Namespace layer
assert isinstance(provider, FastMCPProvider)
assert (
len(provider._transforms) == 1
) # Just Namespace (Visibility is in _visibility)
assert len(provider._transforms) == 1 # Just Namespace
assert isinstance(provider._transforms[0], Namespace)
assert provider.server is sub
@ -854,9 +852,7 @@ class TestAsProxyKwarg:
provider = mcp._providers[1]
# With namespace, we get FastMCPProvider with a Namespace layer
assert isinstance(provider, FastMCPProvider)
assert (
len(provider._transforms) == 1
) # Just Namespace (Visibility is in _visibility)
assert len(provider._transforms) == 1 # Just Namespace
assert isinstance(provider._transforms[0], Namespace)
assert provider.server is sub
@ -870,9 +866,7 @@ class TestAsProxyKwarg:
provider = mcp._providers[1]
# With namespace, we get FastMCPProvider with a Namespace layer
assert isinstance(provider, FastMCPProvider)
assert (
len(provider._transforms) == 1
) # Just Namespace (Visibility is in _visibility)
assert len(provider._transforms) == 1 # Just Namespace
assert isinstance(provider._transforms[0], Namespace)
assert provider.server is not sub
assert isinstance(provider.server, FastMCPProxy)
@ -897,9 +891,7 @@ class TestAsProxyKwarg:
# Index 1 because LocalProvider is at index 0
provider = mcp._providers[1]
assert isinstance(provider, FastMCPProvider)
assert (
len(provider._transforms) == 1
) # Just Namespace (Visibility is in _visibility)
assert len(provider._transforms) == 1 # Just Namespace
assert isinstance(provider._transforms[0], Namespace)
assert provider.server is sub
@ -913,9 +905,7 @@ class TestAsProxyKwarg:
# Index 1 because LocalProvider is at index 0
provider = mcp._providers[1]
assert isinstance(provider, FastMCPProvider)
assert (
len(provider._transforms) == 1
) # Just Namespace (Visibility is in _visibility)
assert len(provider._transforms) == 1 # Just Namespace
assert isinstance(provider._transforms[0], Namespace)
assert provider.server is sub_proxy
@ -929,9 +919,7 @@ class TestAsProxyKwarg:
# Index 1 because LocalProvider is at index 0
provider = mcp._providers[1]
assert isinstance(provider, FastMCPProvider)
assert (
len(provider._transforms) == 1
) # Just Namespace (Visibility is in _visibility)
assert len(provider._transforms) == 1 # Just Namespace
assert isinstance(provider._transforms[0], Namespace)
assert provider.server is sub_proxy
@ -945,9 +933,7 @@ class TestAsProxyKwarg:
# Index 1 because LocalProvider is at index 0
provider = mcp._providers[1]
assert isinstance(provider, FastMCPProvider)
assert (
len(provider._transforms) == 1
) # Just Namespace (Visibility is in _visibility)
assert len(provider._transforms) == 1 # Just Namespace
assert isinstance(provider._transforms[0], Namespace)
assert provider.server is sub_proxy
@ -1183,9 +1169,7 @@ 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) == 1
) # Just Namespace (Visibility is in _visibility)
assert len(provider1._transforms) == 1 # Just Namespace
assert isinstance(provider1._transforms[0], Namespace)
assert provider1.server == sub_server1
assert provider1._transforms[0]._prefix == "sub1"
@ -1195,9 +1179,7 @@ class TestCustomRouteForwarding:
assert len(main_server._providers) == 3
provider2 = main_server._providers[2]
assert isinstance(provider2, FastMCPProvider)
assert (
len(provider2._transforms) == 1
) # Just Namespace (Visibility is in _visibility)
assert len(provider2._transforms) == 1 # Just Namespace
assert isinstance(provider2._transforms[0], Namespace)
assert provider2.server == sub_server2
assert provider2._transforms[0]._prefix == "sub2"