From a42c0c40b04ef932e5b0f4a5ad2a7890841fda45 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Wed, 18 Jun 2025 21:41:28 -0400 Subject: [PATCH] Add middleware for all current handlers --- src/fastmcp/server/middleware.py | 2 +- src/fastmcp/server/server.py | 296 +++++++++++++-------- tests/server/middleware/test_middleware.py | 262 +++++++++++++++++- 3 files changed, 442 insertions(+), 118 deletions(-) diff --git a/src/fastmcp/server/middleware.py b/src/fastmcp/server/middleware.py index 3f73bfe07..2363d41be 100644 --- a/src/fastmcp/server/middleware.py +++ b/src/fastmcp/server/middleware.py @@ -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) diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index 86298d5e1..8eabe815d 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -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] diff --git a/tests/server/middleware/test_middleware.py b/tests/server/middleware/test_middleware.py index 3a2bf7f2c..67f15138c 100644 --- a/tests/server/middleware/test_middleware.py +++ b/tests/server/middleware/test_middleware.py @@ -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 + )