mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-23 14:04:18 +02:00
Preserve discovery middleware contracts
🤖 Generated with OpenAI Codex
This commit is contained in:
parent
17817d5472
commit
82f76f7482
3 changed files with 66 additions and 6 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue