From 07655e662079ca7b213332b033b455d90df3a00e Mon Sep 17 00:00:00 2001
From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com>
Date: Fri, 25 Jul 2025 15:46:48 -0700
Subject: [PATCH] Fix typing, add tests for tool call middleware (#1269)
---
docs/servers/middleware.mdx | 34 ++++++++++++++++++++
src/fastmcp/server/middleware/middleware.py | 26 ++-------------
tests/server/middleware/test_middleware.py | 35 ++++++++++++++++++++-
3 files changed, 71 insertions(+), 24 deletions(-)
diff --git a/docs/servers/middleware.mdx b/docs/servers/middleware.mdx
index 6ab894aae..3251a5bea 100644
--- a/docs/servers/middleware.mdx
+++ b/docs/servers/middleware.mdx
@@ -202,6 +202,40 @@ class ListingFilterMiddleware(Middleware):
This filtering happens before the components are converted to MCP format and returned to the client, so the tags (which are FastMCP-specific) are naturally stripped in the final response.
+
+When filtering components in listing operations, ensure you also prevent execution of filtered components in the corresponding execution hooks (`on_call_tool`, `on_read_resource`, `on_get_prompt`) to maintain consistency.
+
+
+### Tool Call Modification
+
+For execution operations like tool calls, you can modify arguments before execution or transform results afterward:
+
+```python
+from fastmcp.server.middleware import Middleware, MiddlewareContext
+
+class ToolCallMiddleware(Middleware):
+ async def on_call_tool(self, context: MiddlewareContext, call_next):
+ # Modify arguments before execution
+ if context.message.name == "calculate":
+ # Ensure positive inputs
+ if context.message.arguments.get("value", 0) < 0:
+ context.message.arguments["value"] = abs(context.message.arguments["value"])
+
+ result = await call_next(context)
+
+ # Transform result after execution
+ if context.message.name == "get_data":
+ # Add metadata to result
+ if result.structured_content:
+ result.structured_content["processed_at"] = "2024-01-01T00:00:00Z"
+
+ return result
+```
+
+
+For more complex tool rewriting scenarios, consider using [Tool Transformation](/patterns/tool-transformation) patterns which provide a more structured approach to creating modified tool variants.
+
+
### Anatomy of a Hook
Every middleware hook follows the same pattern. Let's examine the `on_message` hook to understand the structure:
diff --git a/src/fastmcp/server/middleware/middleware.py b/src/fastmcp/server/middleware/middleware.py
index fa2e93455..9df826665 100644
--- a/src/fastmcp/server/middleware/middleware.py
+++ b/src/fastmcp/server/middleware/middleware.py
@@ -20,7 +20,7 @@ import mcp.types as mt
from fastmcp.prompts.prompt import Prompt
from fastmcp.resources.resource import Resource
from fastmcp.resources.template import ResourceTemplate
-from fastmcp.tools.tool import Tool
+from fastmcp.tools.tool import Tool, ToolResult
if TYPE_CHECKING:
from fastmcp.server.context import Context
@@ -43,26 +43,6 @@ class CallNext(Protocol[T, R]):
def __call__(self, context: MiddlewareContext[T]) -> Awaitable[R]: ...
-ServerResultT = TypeVar(
- "ServerResultT",
- bound=mt.EmptyResult
- | mt.InitializeResult
- | mt.CompleteResult
- | mt.GetPromptResult
- | mt.ListPromptsResult
- | mt.ListResourcesResult
- | mt.ListResourceTemplatesResult
- | mt.ReadResourceResult
- | mt.CallToolResult
- | mt.ListToolsResult,
-)
-
-
-@runtime_checkable
-class ServerResultProtocol(Protocol[ServerResultT]):
- root: ServerResultT
-
-
@dataclass(kw_only=True, frozen=True)
class MiddlewareContext(Generic[T]):
"""
@@ -167,8 +147,8 @@ class Middleware:
async def on_call_tool(
self,
context: MiddlewareContext[mt.CallToolRequestParams],
- call_next: CallNext[mt.CallToolRequestParams, mt.CallToolResult],
- ) -> mt.CallToolResult:
+ call_next: CallNext[mt.CallToolRequestParams, ToolResult],
+ ) -> ToolResult:
return await call_next(context)
async def on_read_resource(
diff --git a/tests/server/middleware/test_middleware.py b/tests/server/middleware/test_middleware.py
index 60e83ef8e..c19261c12 100644
--- a/tests/server/middleware/test_middleware.py
+++ b/tests/server/middleware/test_middleware.py
@@ -7,7 +7,8 @@ import pytest
from fastmcp import Client, FastMCP
from fastmcp.server.context import Context
-from fastmcp.server.middleware import Middleware, MiddlewareContext
+from fastmcp.server.middleware import CallNext, Middleware, MiddlewareContext
+from fastmcp.tools.tool import ToolResult
@dataclass
@@ -411,6 +412,38 @@ class TestMiddlewareHooks:
assert len(prompts) == 1
assert prompts[0].name == "public_prompt"
+ async def test_call_tool_middleware(self):
+ server = FastMCP()
+
+ @server.tool
+ def add(a: int, b: int) -> int:
+ return a + b
+
+ class CallToolMiddleware(Middleware):
+ async def on_call_tool(
+ self,
+ context: MiddlewareContext[mcp.types.CallToolRequestParams],
+ call_next: CallNext[mcp.types.CallToolRequestParams, ToolResult],
+ ):
+ # modify argument
+ if context.message.name == "add":
+ context.message.arguments["a"] += 100 # type: ignore
+
+ result = await call_next(context)
+
+ # modify result
+ if context.message.name == "add":
+ result.structured_content["result"] += 5 # type: ignore
+
+ return result
+
+ server.add_middleware(CallToolMiddleware())
+
+ async with Client(server) as client:
+ result = await client.call_tool("add", {"a": 1, "b": 2})
+
+ assert result.structured_content["result"] == 108 # type: ignore
+
class TestNestedMiddlewareHooks:
@pytest.fixture