Add middleware for all current handlers

This commit is contained in:
Jeremiah Lowin 2025-06-18 21:41:28 -04:00
commit a42c0c40b0
3 changed files with 440 additions and 116 deletions

View file

@ -148,7 +148,7 @@ class MCPMiddleware:
handler = partial(self.on_list_tools, call_next=handler)
case "resources/list":
handler = partial(self.on_list_resources, call_next=handler)
case "resource-templates/list":
case "resources/templates/list":
handler = partial(self.on_list_resource_templates, call_next=handler)
case "prompts/list":
handler = partial(self.on_list_prompts, call_next=handler)

View file

@ -352,34 +352,8 @@ class FastMCP(Generic[LifespanResultT]):
async def get_resources(self) -> dict[str, Resource]:
"""Get all registered resources, indexed by registered key."""
if (resources := self._cache.get("resources")) is self._cache.NOT_FOUND:
resources: dict[str, Resource] = {}
# iterate such that new mounts overwrite older ones
for mounted_server in self._mounted_servers:
try:
server_resources = await mounted_server.server.get_resources()
# Apply prefix to each resource key if prefix exists
if mounted_server.prefix:
for resource in server_resources.values():
resource = resource.with_key(
add_resource_prefix(
resource.key,
mounted_server.prefix,
self.resource_prefix_format,
)
)
resources[resource.key] = resource
else:
resources.update(server_resources)
except Exception as e:
logger.warning(
f"Failed to get resources from mounted server '{mounted_server.prefix}': {e}"
)
continue
resources.update(self._resource_manager.get_resources())
self._cache.set("resources", resources)
return resources
resources = await self._list_resources(apply_middleware=False)
return {resource.key: resource for resource in resources}
async def get_resource(self, key: str) -> Resource:
resources = await self.get_resources()
@ -389,39 +363,8 @@ class FastMCP(Generic[LifespanResultT]):
async def get_resource_templates(self) -> dict[str, ResourceTemplate]:
"""Get all registered resource templates, indexed by registered key."""
if (
templates := self._cache.get("resource_templates")
) is self._cache.NOT_FOUND:
templates: dict[str, ResourceTemplate] = {}
# iterate such that new mounts overwrite older ones
for mounted_server in self._mounted_servers:
try:
server_templates = (
await mounted_server.server.get_resource_templates()
)
# Apply prefix to each template key if prefix exists
if mounted_server.prefix:
for template in server_templates.values():
template = template.with_key(
add_resource_prefix(
template.key,
mounted_server.prefix,
self.resource_prefix_format,
)
)
templates[template.key] = template
else:
templates.update(server_templates)
except Exception as e:
logger.warning(
"Failed to get resource templates from mounted server "
f"'{mounted_server.prefix}': {e}"
)
continue
templates.update(self._resource_manager.get_templates())
self._cache.set("resource_templates", templates)
return templates
templates = await self._list_resource_templates(apply_middleware=False)
return {template.key: template for template in templates}
async def get_resource_template(self, key: str) -> ResourceTemplate:
templates = await self.get_resource_templates()
@ -434,30 +377,8 @@ class FastMCP(Generic[LifespanResultT]):
List all available prompts.
"""
if (prompts := self._cache.get("prompts")) is self._cache.NOT_FOUND:
prompts: dict[str, Prompt] = {}
# iterate such that new mounts overwrite older ones
for mounted_server in self._mounted_servers:
try:
server_prompts = await mounted_server.server.get_prompts()
# Apply prefix to each prompt key if prefix exists
if mounted_server.prefix:
for prompt in server_prompts.values():
prompt = prompt.with_key(
f"{mounted_server.prefix}_{prompt.key}"
)
prompts[prompt.key] = prompt
else:
prompts.update(server_prompts)
except Exception as e:
logger.warning(
f"Failed to get prompts from mounted server '{mounted_server.prefix}': {e}"
)
continue
prompts.update(self._prompt_manager.get_prompts())
self._cache.set("prompts", prompts)
return prompts
prompts = await self._list_prompts(apply_middleware=False)
return {prompt.key: prompt for prompt in prompts}
async def get_prompt(self, key: str) -> Prompt:
prompts = await self.get_prompts()
@ -582,21 +503,31 @@ class FastMCP(Generic[LifespanResultT]):
return tools
async def _mcp_list_resources(self) -> list[MCPResource]:
logger.debug("Handler called: list_resources")
with fastmcp.server.context.Context(fastmcp=self):
resources = await self._middleware_list_resources()
return [
resource.to_mcp_resource(uri=resource.key) for resource in resources
]
async def _middleware_list_resources(self) -> list[Resource]:
"""
List all available resources, in the format expected by the low-level MCP
server.
"""
logger.debug("Handler called: list_resources")
async def _final_handler(
async def _handler(
context: MiddlewareContext[dict[str, Any]],
) -> list[MCPResource]:
resources = await self.get_resources()
mcp_resources: list[MCPResource] = []
for key, resource in resources.items():
) -> list[Resource]:
resources = await self._list_resources()
mcp_resources: list[Resource] = []
for resource in resources:
if self._should_enable_component(resource):
mcp_resources.append(resource.to_mcp_resource(uri=key))
mcp_resources.append(resource)
return mcp_resources
with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx:
@ -610,24 +541,74 @@ class FastMCP(Generic[LifespanResultT]):
)
# Apply the middleware chain.
return await self._apply_middleware(mw_context, _final_handler)
return await self._apply_middleware(mw_context, _handler)
async def _list_resources(self, apply_middleware: bool = True) -> list[Resource]:
"""
List all available resources.
"""
if (resources := self._cache.get("resources")) is self._cache.NOT_FOUND:
resources: list[Resource] = []
# iterate such that new mounts overwrite older ones
for mounted_server in self._mounted_servers:
try:
if apply_middleware:
server_resources = (
await mounted_server.server._middleware_list_resources()
)
else:
server_resources = await mounted_server.server._list_resources()
# Apply prefix to each resource key if prefix exists
if mounted_server.prefix:
for resource in server_resources:
resource = resource.with_key(
add_resource_prefix(
resource.key,
mounted_server.prefix,
self.resource_prefix_format,
)
)
resources.append(resource)
else:
resources.extend(server_resources)
except Exception as e:
logger.warning(
f"Failed to get resources from mounted server '{mounted_server.prefix}': {e}"
)
continue
resources.extend(self._resource_manager.get_resources().values())
self._cache.set("resources", resources)
return resources
async def _mcp_list_resource_templates(self) -> list[MCPResourceTemplate]:
"""
List all available resource templates, in the format expected by the low-level
MCP server.
"""
logger.debug("Handler called: list_resource_templates")
async def _final_handler(
with fastmcp.server.context.Context(fastmcp=self):
templates = await self._middleware_list_resource_templates()
return [
template.to_mcp_template(uriTemplate=template.key)
for template in templates
]
async def _middleware_list_resource_templates(self) -> list[ResourceTemplate]:
"""
List all available resource templates, in the format expected by the low-level MCP
server.
"""
async def _handler(
context: MiddlewareContext[dict[str, Any]],
) -> list[MCPResourceTemplate]:
templates = await self.get_resource_templates()
mcp_templates: list[MCPResourceTemplate] = []
for key, template in templates.items():
) -> list[ResourceTemplate]:
templates = await self._list_resource_templates()
mcp_templates: list[ResourceTemplate] = []
for template in templates:
if self._should_enable_component(template):
mcp_templates.append(template.to_mcp_template(uriTemplate=key))
mcp_templates.append(template)
return mcp_templates
with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx:
@ -636,29 +617,81 @@ class FastMCP(Generic[LifespanResultT]):
message={}, # List resource templates doesn't have parameters
source="client",
type="request",
method="resources/list_templates",
method="resources/templates/list",
fastmcp_context=fastmcp_ctx,
)
# Apply the middleware chain.
return await self._apply_middleware(mw_context, _final_handler)
return await self._apply_middleware(mw_context, _handler)
async def _list_resource_templates(
self, apply_middleware: bool = True
) -> list[ResourceTemplate]:
"""
List all available resource templates.
"""
if (
templates := self._cache.get("resource_templates")
) is self._cache.NOT_FOUND:
templates: list[ResourceTemplate] = []
# iterate such that new mounts overwrite older ones
for mounted_server in self._mounted_servers:
try:
if apply_middleware:
server_templates = await mounted_server.server._middleware_list_resource_templates()
else:
server_templates = (
await mounted_server.server._list_resource_templates()
)
# Apply prefix to each template key if prefix exists
if mounted_server.prefix:
for template in server_templates:
template = template.with_key(
add_resource_prefix(
template.key,
mounted_server.prefix,
self.resource_prefix_format,
)
)
templates.append(template)
else:
templates.extend(server_templates)
except Exception as e:
logger.warning(
"Failed to get resource templates from mounted server "
f"'{mounted_server.prefix}': {e}"
)
continue
templates.extend(self._resource_manager.get_templates().values())
self._cache.set("resource_templates", templates)
return templates
async def _mcp_list_prompts(self) -> list[MCPPrompt]:
logger.debug("Handler called: list_prompts")
with fastmcp.server.context.Context(fastmcp=self):
prompts = await self._middleware_list_prompts()
return [prompt.to_mcp_prompt(name=prompt.key) for prompt in prompts]
async def _middleware_list_prompts(self) -> list[Prompt]:
"""
List all available prompts, in the format expected by the low-level MCP
server.
"""
logger.debug("Handler called: list_prompts")
async def _final_handler(
async def _handler(
context: MiddlewareContext[dict[str, Any]],
) -> list[MCPPrompt]:
prompts = await self.get_prompts()
mcp_prompts: list[MCPPrompt] = []
for key, prompt in prompts.items():
) -> list[Prompt]:
prompts = await self._list_prompts()
mcp_prompts: list[Prompt] = []
for prompt in prompts:
if self._should_enable_component(prompt):
mcp_prompts.append(prompt.to_mcp_prompt(name=key))
mcp_prompts.append(prompt)
return mcp_prompts
with fastmcp.server.context.Context(fastmcp=self) as fastmcp_ctx:
@ -672,7 +705,42 @@ class FastMCP(Generic[LifespanResultT]):
)
# Apply the middleware chain.
return await self._apply_middleware(mw_context, _final_handler)
return await self._apply_middleware(mw_context, _handler)
async def _list_prompts(self, apply_middleware: bool = True) -> list[Prompt]:
"""
List all available prompts.
"""
if (prompts := self._cache.get("prompts")) is self._cache.NOT_FOUND:
prompts: list[Prompt] = []
# iterate such that new mounts overwrite older ones
for mounted_server in self._mounted_servers:
try:
if apply_middleware:
server_prompts = (
await mounted_server.server._middleware_list_prompts()
)
else:
server_prompts = await mounted_server.server._list_prompts()
# Apply prefix to each prompt key if prefix exists
if mounted_server.prefix:
for prompt in server_prompts:
prompt = prompt.with_key(
f"{mounted_server.prefix}_{prompt.key}"
)
prompts.append(prompt)
else:
prompts.extend(server_prompts)
except Exception as e:
logger.warning(
f"Failed to get prompts from mounted server '{mounted_server.prefix}': {e}"
)
continue
prompts.extend(self._prompt_manager.get_prompts().values())
self._cache.set("prompts", prompts)
return prompts
async def _mcp_call_tool(
self, key: str, arguments: dict[str, Any]

View file

@ -21,9 +21,10 @@ class Recording:
class RecordingMiddleware(MCPMiddleware):
"""A middleware that automatically records all method calls."""
def __init__(self):
def __init__(self, name: str | None = None):
super().__init__()
self.calls: list[Recording] = []
self.name = name
def __getattribute__(self, name: str) -> Callable:
"""Dynamically create recording methods for any on_* method."""
@ -95,7 +96,7 @@ class RecordingMiddleware(MCPMiddleware):
@pytest.fixture
def recording_middleware():
"""Fixture that provides a recording middleware instance."""
middleware = RecordingMiddleware()
middleware = RecordingMiddleware(name="recording_middleware")
yield middleware
@ -215,7 +216,7 @@ class TestMiddlewareHooks:
assert recording_middleware.assert_called(times=3)
assert recording_middleware.assert_called(
method="resource-templates/list", times=3
method="resources/templates/list", times=3
)
assert recording_middleware.assert_called(hook="on_message", times=1)
assert recording_middleware.assert_called(hook="on_request", times=1)
@ -234,3 +235,258 @@ class TestMiddlewareHooks:
assert recording_middleware.assert_called(hook="on_message", times=1)
assert recording_middleware.assert_called(hook="on_request", times=1)
assert recording_middleware.assert_called(hook="on_list_prompts", times=1)
class TestNestedMiddlewareHooks:
@pytest.fixture
@staticmethod
def nested_middleware():
return RecordingMiddleware(name="nested_middleware")
@pytest.fixture
def nested_mcp_server(self, nested_middleware: RecordingMiddleware):
mcp = FastMCP(name="Nested MCP")
@mcp.tool
def add(a: int, b: int) -> int:
return a + b
@mcp.resource("resource://test")
def test_resource() -> str:
return "test resource"
@mcp.resource("resource://test-template/{x}")
def test_resource_with_path(x: int) -> str:
return f"test resource with {x}"
@mcp.prompt
def test_prompt(x: str) -> str:
return f"test prompt with {x}"
@mcp.tool
async def progress_tool(context: Context) -> None:
await context.report_progress(progress=1, total=10, message="test")
@mcp.tool
async def log_tool(context: Context) -> None:
await context.info(message="test log")
@mcp.tool
async def sample_tool(context: Context) -> None:
await context.sample("hello")
mcp.add_middleware(nested_middleware)
return mcp
async def test_call_tool_on_parent_server(
self,
mcp_server: FastMCP,
nested_mcp_server: FastMCP,
recording_middleware: RecordingMiddleware,
nested_middleware: RecordingMiddleware,
):
mcp_server.mount(nested_mcp_server, prefix="nested")
async with Client(mcp_server) as client:
await client.call_tool("add", {"a": 1, "b": 2})
assert recording_middleware.assert_called(times=3)
assert recording_middleware.assert_called(method="tools/call", times=3)
assert recording_middleware.assert_called(hook="on_message", times=1)
assert recording_middleware.assert_called(hook="on_request", times=1)
assert recording_middleware.assert_called(hook="on_call_tool", times=1)
assert nested_middleware.assert_called(times=0)
async def test_call_tool_on_nested_server(
self,
mcp_server: FastMCP,
nested_mcp_server: FastMCP,
recording_middleware: RecordingMiddleware,
nested_middleware: RecordingMiddleware,
):
mcp_server.mount(nested_mcp_server, prefix="nested")
async with Client(mcp_server) as client:
await client.call_tool("nested_add", {"a": 1, "b": 2})
assert recording_middleware.assert_called(times=3)
assert recording_middleware.assert_called(method="tools/call", times=3)
assert recording_middleware.assert_called(hook="on_message", times=1)
assert recording_middleware.assert_called(hook="on_request", times=1)
assert recording_middleware.assert_called(hook="on_call_tool", times=1)
assert nested_middleware.assert_called(times=3)
assert nested_middleware.assert_called(method="tools/call", times=3)
assert nested_middleware.assert_called(hook="on_message", times=1)
assert nested_middleware.assert_called(hook="on_request", times=1)
assert nested_middleware.assert_called(hook="on_call_tool", times=1)
async def test_read_resource_on_parent_server(
self,
mcp_server: FastMCP,
nested_mcp_server: FastMCP,
recording_middleware: RecordingMiddleware,
nested_middleware: RecordingMiddleware,
):
mcp_server.mount(nested_mcp_server, prefix="nested")
async with Client(mcp_server) as client:
await client.read_resource("resource://test")
assert recording_middleware.assert_called(times=3)
assert recording_middleware.assert_called(method="resources/read", times=3)
assert recording_middleware.assert_called(hook="on_message", times=1)
assert recording_middleware.assert_called(hook="on_request", times=1)
assert recording_middleware.assert_called(hook="on_read_resource", times=1)
assert nested_middleware.assert_called(times=0)
async def test_read_resource_on_nested_server(
self,
mcp_server: FastMCP,
nested_mcp_server: FastMCP,
recording_middleware: RecordingMiddleware,
nested_middleware: RecordingMiddleware,
):
mcp_server.mount(nested_mcp_server, prefix="nested")
async with Client(mcp_server) as client:
await client.read_resource("resource://nested/test")
assert recording_middleware.assert_called(times=3)
assert recording_middleware.assert_called(method="resources/read", times=3)
assert recording_middleware.assert_called(hook="on_message", times=1)
assert recording_middleware.assert_called(hook="on_request", times=1)
assert recording_middleware.assert_called(hook="on_read_resource", times=1)
assert nested_middleware.assert_called(times=3)
assert nested_middleware.assert_called(method="resources/read", times=3)
assert nested_middleware.assert_called(hook="on_message", times=1)
assert nested_middleware.assert_called(hook="on_request", times=1)
assert nested_middleware.assert_called(hook="on_read_resource", times=1)
async def test_get_prompt_on_parent_server(
self,
mcp_server: FastMCP,
nested_mcp_server: FastMCP,
recording_middleware: RecordingMiddleware,
nested_middleware: RecordingMiddleware,
):
mcp_server.mount(nested_mcp_server, prefix="nested")
async with Client(mcp_server) as client:
await client.get_prompt("test_prompt", {"x": "test"})
assert recording_middleware.assert_called(times=3)
assert recording_middleware.assert_called(method="prompts/get", times=3)
assert recording_middleware.assert_called(hook="on_message", times=1)
assert recording_middleware.assert_called(hook="on_request", times=1)
assert recording_middleware.assert_called(hook="on_get_prompt", times=1)
assert nested_middleware.assert_called(times=0)
async def test_get_prompt_on_nested_server(
self,
mcp_server: FastMCP,
nested_mcp_server: FastMCP,
recording_middleware: RecordingMiddleware,
nested_middleware: RecordingMiddleware,
):
mcp_server.mount(nested_mcp_server, prefix="nested")
async with Client(mcp_server) as client:
await client.get_prompt("nested_test_prompt", {"x": "test"})
assert recording_middleware.assert_called(times=3)
assert recording_middleware.assert_called(method="prompts/get", times=3)
assert recording_middleware.assert_called(hook="on_message", times=1)
assert recording_middleware.assert_called(hook="on_request", times=1)
assert recording_middleware.assert_called(hook="on_get_prompt", times=1)
assert nested_middleware.assert_called(times=3)
assert nested_middleware.assert_called(method="prompts/get", times=3)
assert nested_middleware.assert_called(hook="on_message", times=1)
assert nested_middleware.assert_called(hook="on_request", times=1)
assert nested_middleware.assert_called(hook="on_get_prompt", times=1)
async def test_list_tools_on_nested_server(
self,
mcp_server: FastMCP,
nested_mcp_server: FastMCP,
recording_middleware: RecordingMiddleware,
nested_middleware: RecordingMiddleware,
):
mcp_server.mount(nested_mcp_server, prefix="nested")
async with Client(mcp_server) as client:
await client.list_tools()
assert recording_middleware.assert_called(times=3)
assert recording_middleware.assert_called(method="tools/list", times=3)
assert recording_middleware.assert_called(hook="on_message", times=1)
assert recording_middleware.assert_called(hook="on_request", times=1)
assert recording_middleware.assert_called(hook="on_list_tools", times=1)
assert nested_middleware.assert_called(times=3)
assert nested_middleware.assert_called(method="tools/list", times=3)
assert nested_middleware.assert_called(hook="on_message", times=1)
assert nested_middleware.assert_called(hook="on_request", times=1)
assert nested_middleware.assert_called(hook="on_list_tools", times=1)
async def test_list_resources_on_nested_server(
self,
mcp_server: FastMCP,
nested_mcp_server: FastMCP,
recording_middleware: RecordingMiddleware,
nested_middleware: RecordingMiddleware,
):
mcp_server.mount(nested_mcp_server, prefix="nested")
async with Client(mcp_server) as client:
await client.list_resources()
assert recording_middleware.assert_called(times=3)
assert recording_middleware.assert_called(method="resources/list", times=3)
assert recording_middleware.assert_called(hook="on_message", times=1)
assert recording_middleware.assert_called(hook="on_request", times=1)
assert recording_middleware.assert_called(hook="on_list_resources", times=1)
assert nested_middleware.assert_called(times=3)
assert nested_middleware.assert_called(method="resources/list", times=3)
assert nested_middleware.assert_called(hook="on_message", times=1)
assert nested_middleware.assert_called(hook="on_request", times=1)
assert nested_middleware.assert_called(hook="on_list_resources", times=1)
async def test_list_resource_templates_on_nested_server(
self,
mcp_server: FastMCP,
nested_mcp_server: FastMCP,
recording_middleware: RecordingMiddleware,
nested_middleware: RecordingMiddleware,
):
mcp_server.mount(nested_mcp_server, prefix="nested")
async with Client(mcp_server) as client:
await client.list_resource_templates()
assert recording_middleware.assert_called(times=3)
assert recording_middleware.assert_called(
method="resources/templates/list", times=3
)
assert recording_middleware.assert_called(hook="on_message", times=1)
assert recording_middleware.assert_called(hook="on_request", times=1)
assert recording_middleware.assert_called(
hook="on_list_resource_templates", times=1
)
assert nested_middleware.assert_called(times=3)
assert nested_middleware.assert_called(
method="resources/templates/list", times=3
)
assert nested_middleware.assert_called(hook="on_message", times=1)
assert nested_middleware.assert_called(hook="on_request", times=1)
assert nested_middleware.assert_called(
hook="on_list_resource_templates", times=1
)