Preserve discovery middleware contracts

🤖 Generated with OpenAI Codex
This commit is contained in:
Jake Kaplan 2026-08-05 21:48:23 -04:00
commit 82f76f7482
3 changed files with 66 additions and 6 deletions

View file

@ -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):

View file

@ -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

View file

@ -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):