From 82f76f74820b0861dca3f351f6fa509378aa4ca1 Mon Sep 17 00:00:00 2001 From: Jake Kaplan Date: Wed, 5 Aug 2026 21:48:23 -0400 Subject: [PATCH] Preserve discovery middleware contracts MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 🤖 Generated with OpenAI Codex --- fastmcp_slim/fastmcp/server/low_level.py | 19 ++++++---- .../middleware/test_discovery_middleware.py | 36 +++++++++++++++++++ .../middleware/test_message_visibility.py | 17 +++++++++ 3 files changed, 66 insertions(+), 6 deletions(-) diff --git a/fastmcp_slim/fastmcp/server/low_level.py b/fastmcp_slim/fastmcp/server/low_level.py index 3d2ad9bf9..16415fea5 100644 --- a/fastmcp_slim/fastmcp/server/low_level.py +++ b/fastmcp_slim/fastmcp/server/low_level.py @@ -38,7 +38,6 @@ if TYPE_CHECKING: logger = get_logger(__name__) - # The request methods that FastMCP serves through a handler adapter, each of which # runs the FastMCP middleware chain interior (see MCPOperationsMixin). The root # dispatch leaves these to the interior dispatch and only observes them if they fail before @@ -332,15 +331,23 @@ class FastMCPServerMiddleware: from fastmcp.server.context import Context from fastmcp.server.middleware.middleware import MiddlewareContext - params = ctx.params if isinstance(ctx.params, dict) else {} - discover_message = mcp_types.DiscoverRequest.model_validate( - {"method": "server/discover", "params": params}, by_name=False - ) + try: + discover_message = mcp_types.DiscoverRequest.model_validate( + {"method": "server/discover", "params": ctx.params}, by_name=False + ) + except ValidationError as exc: + return await self._run_outer_mw(fastmcp, ctx, call_next, _raise=exc) async def call_original_handler( _mw_ctx: MiddlewareContext, ) -> mcp_types.DiscoverResult: - raw = await call_next(ctx) + message = _mw_ctx.message + params = ( + message.params.model_dump(by_alias=True, mode="json", exclude_none=True) + if message.params is not None + else None + ) + raw = await call_next(replace(ctx, params=params)) if isinstance(raw, mcp_types.DiscoverResult): return raw if isinstance(raw, Mapping): diff --git a/tests/server/middleware/test_discovery_middleware.py b/tests/server/middleware/test_discovery_middleware.py index d7af479c7..4cd9ab617 100644 --- a/tests/server/middleware/test_discovery_middleware.py +++ b/tests/server/middleware/test_discovery_middleware.py @@ -29,3 +29,39 @@ async def test_on_discover_receives_and_transforms_typed_result(): assert isinstance(middleware.request, mcp_types.DiscoverRequest) assert isinstance(middleware.result, mcp_types.DiscoverResult) + + +async def test_on_discover_forwards_modified_params(): + modified = False + server = FastMCP("modified-discovery") + default_handler = server._mcp_server._handle_discover + + async def capture_params(ctx, params): + nonlocal modified + assert params is not None + assert params.meta is not None + modified = params.meta["com.example/modified"] is True + return await default_handler(ctx, params) + + server._mcp_server.add_request_handler( + "server/discover", mcp_types.RequestParams, capture_params + ) + + class ModifyParams(Middleware): + async def on_discover(self, context, call_next): + assert context.message.params is not None + assert context.message.params.meta is not None + context.message.params = mcp_types.RequestParams( + meta={ + **context.message.params.meta, + "com.example/modified": True, + } + ) + return await call_next(context) + + server.add_middleware(ModifyParams()) + + async with Client(server, mode="auto"): + pass + + assert modified diff --git a/tests/server/middleware/test_message_visibility.py b/tests/server/middleware/test_message_visibility.py index 8145c2b06..5c2465637 100644 --- a/tests/server/middleware/test_message_visibility.py +++ b/tests/server/middleware/test_message_visibility.py @@ -159,6 +159,23 @@ class TestUnroutableAndMalformed: assert ("on_message", "tools/call") in recorder.records assert ("on_call_tool", "tools/call") not in recorder.records + async def test_malformed_discover_params_observed_by_generic_hooks(self): + server = _adder() + recorder = HookRecorder() + server.add_middleware(recorder) + + async with Client(server) as client: + recorder.records.clear() + with pytest.raises(MCPError): + await _raw_request( + client, + "server/discover", + {"_meta": {"progressToken": []}}, + ) + + assert ("on_message", "server/discover") in recorder.records + assert ("on_request", "server/discover") in recorder.records + class TestSingleFire: async def test_each_hook_fires_once_per_component_call(self):