Refactor transform list methods to pure function pattern (#2942)

This commit is contained in:
Jeremiah Lowin 2026-01-19 16:21:35 -05:00 committed by GitHub
commit 4d2feb0c29
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
14 changed files with 143 additions and 334 deletions

View file

@ -4,6 +4,7 @@ from typing import Any
from fastmcp.server.providers.base import Provider
from fastmcp.server.tasks.config import TaskConfig
from fastmcp.server.transforms import Namespace
from fastmcp.tools.tool import Tool, ToolResult
@ -77,3 +78,14 @@ class TestBaseProviderGetTasks:
assert len(tasks) == 1
assert tasks[0].name == "enabled"
async def test_get_tasks_applies_transforms(self):
"""get_tasks should apply provider transforms to component names."""
tool = CustomTool(name="my_tool", description="A tool")
provider = SimpleProvider(tools=[tool])
provider.add_transform(Namespace("api"))
tasks = await provider.get_tasks()
assert len(tasks) == 1
assert tasks[0].name == "api_my_tool"

View file

@ -600,11 +600,9 @@ class TestProviderToolTransformations:
layer = ToolTransform({"my_tool": ToolTransformConfig(name="renamed_tool")})
provider.add_transform(layer)
# Use call_next pattern
async def get_tools():
return await provider.list_tools()
transformed_tools = await layer.list_tools(get_tools)
# Get tools and pass directly to transform
tools = await provider.list_tools()
transformed_tools = await layer.list_tools(tools)
assert len(transformed_tools) == 1
assert transformed_tools[0].name == "renamed_tool"
@ -676,11 +674,8 @@ class TestProviderToolTransformations:
original_tools = await provider._list_tools()
assert original_tools[0].name == "my_tool"
# Transform modifies them when applied via call_next
async def get_tools():
return original_tools
transformed_tools = await layer.list_tools(get_tools)
# Transform modifies them when applied directly
transformed_tools = await layer.list_tools(original_tools)
assert transformed_tools[0].name == "renamed"
def test_transform_layer_duplicate_target_name_raises_error(self):

View file

@ -23,11 +23,9 @@ class TestNamespaceTransform:
provider = FastMCPProvider(server)
layer = Namespace("ns")
# Use call_next pattern - create a callable that returns the tools
async def get_tools():
return await provider.list_tools()
transformed_tools = await layer.list_tools(get_tools)
# Get tools and pass directly to transform
tools = await provider.list_tools()
transformed_tools = await layer.list_tools(tools)
assert len(transformed_tools) == 1
assert transformed_tools[0].name == "ns_my_tool"
@ -43,10 +41,8 @@ class TestNamespaceTransform:
provider = FastMCPProvider(server)
layer = Namespace("ns")
async def get_prompts():
return await provider.list_prompts()
transformed_prompts = await layer.list_prompts(get_prompts)
prompts = await provider.list_prompts()
transformed_prompts = await layer.list_prompts(prompts)
assert len(transformed_prompts) == 1
assert transformed_prompts[0].name == "ns_my_prompt"
@ -62,10 +58,8 @@ class TestNamespaceTransform:
provider = FastMCPProvider(server)
layer = Namespace("ns")
async def get_resources():
return await provider.list_resources()
transformed_resources = await layer.list_resources(get_resources)
resources = await provider.list_resources()
transformed_resources = await layer.list_resources(resources)
assert len(transformed_resources) == 1
assert str(transformed_resources[0].uri) == "resource://ns/data"
@ -81,10 +75,8 @@ class TestNamespaceTransform:
provider = FastMCPProvider(server)
layer = Namespace("ns")
async def get_templates():
return await provider.list_resource_templates()
transformed_templates = await layer.list_resource_templates(get_templates)
templates = await provider.list_resource_templates()
transformed_templates = await layer.list_resource_templates(templates)
assert len(transformed_templates) == 1
assert transformed_templates[0].uri_template == "resource://ns/{name}/data"
@ -104,10 +96,8 @@ class TestToolTransformRenames:
provider = FastMCPProvider(server)
layer = ToolTransform({"verbose_tool_name": ToolTransformConfig(name="short")})
async def get_tools():
return await provider.list_tools()
transformed_tools = await layer.list_tools(get_tools)
tools = await provider.list_tools()
transformed_tools = await layer.list_tools(tools)
assert len(transformed_tools) == 1
assert transformed_tools[0].name == "short"
@ -239,14 +229,10 @@ class TestTransformStacking:
inner_layer = Namespace("inner")
outer_layer = Namespace("outer")
# Build chain: base -> inner -> outer
async def base():
return await provider.list_tools()
async def inner_chain():
return await inner_layer.list_tools(base)
tools = await outer_layer.list_tools(inner_chain)
# Apply transforms sequentially: base -> inner -> outer
tools = await provider.list_tools()
tools = await inner_layer.list_tools(tools)
tools = await outer_layer.list_tools(tools)
assert len(tools) == 1
assert tools[0].name == "outer_inner_my_tool"
@ -289,10 +275,8 @@ class TestNoTransformation:
provider = FastMCPProvider(server)
transform = Transform()
async def get_tools():
return await provider.list_tools()
transformed_tools = await transform.list_tools(get_tools)
tools = await provider.list_tools()
transformed_tools = await transform.list_tools(tools)
assert len(transformed_tools) == 1
assert transformed_tools[0].name == "my_tool"
@ -308,10 +292,8 @@ class TestNoTransformation:
provider = FastMCPProvider(server)
layer = ToolTransform({})
async def get_tools():
return await provider.list_tools()
transformed_tools = await layer.list_tools(get_tools)
tools = await provider.list_tools()
transformed_tools = await layer.list_tools(tools)
assert len(transformed_tools) == 1
assert transformed_tools[0].name == "my_tool"

View file

@ -236,10 +236,7 @@ class TestTransformChain:
"""list_tools applies marks to matching components."""
disable_internal = Enabled(False, tags=set({"internal"}))
async def base():
return tools
result = await disable_internal.list_tools(base)
result = await disable_internal.list_tools(tools)
assert len(result) == 3
assert is_enabled(result[0]) # public
@ -251,12 +248,8 @@ class TestTransformChain:
disable_internal = Enabled(False, tags=set({"internal"}))
enable_safe = Enabled(True, tags=set({"safe"}))
async def base():
return tools
async def after_disable():
return await disable_internal.list_tools(base)
# Apply transforms sequentially
after_disable = await disable_internal.list_tools(tools)
result = await enable_safe.list_tools(after_disable)
enabled = [t for t in result if is_enabled(t)]
@ -270,12 +263,8 @@ class TestTransformChain:
disable_all = Enabled(False, match_all=True)
enable_public = Enabled(True, tags=set({"public"}))
async def base():
return tools
async def after_disable():
return await disable_all.list_tools(base)
# Apply transforms sequentially
after_disable = await disable_all.list_tools(tools)
result = await enable_public.list_tools(after_disable)
enabled = [t for t in result if is_enabled(t)]