Run FastMCP middleware for every inbound message (#4553)

* Make the SDK seam the root of FastMCP middleware dispatch (D3)

Notifications, cancellations, and malformed/unroutable messages now reach
on_message/on_request/on_notification at the SDK seam. Component methods keep
their interior dispatch (typed hooks, tool-exception visibility) unchanged; the
seam covers only messages the interior never dispatches, so each hook fires once.

* Document the middleware seam coverage and suspend semantics (D3)

* Align seam docs and ask-visibility test with the result-cycle MRTR model

An InputRequiredResult is the full result of a complete request->response
cycle, not a suspension: component hooks observe an asking round's
InputRequiredToolResult as an ordinary return value.

* Replace 'seam' language with plain dispatch terminology

* Keep the raw middleware __call__ signature; forward middleware message edits

* Cover fires-once across an MRTR continuation round

* Align cherry-picked coverage test with renamed recorder

* Rewrite only the message, never the dispatch destination
This commit is contained in:
Jeremiah Lowin 2026-07-19 20:29:08 -04:00 committed by GitHub
commit 3213776b25
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 740 additions and 51 deletions

View file

@ -124,6 +124,12 @@ Server-side middleware is a new first-class SDK concept: `Server.middleware` is
*Verify:* `fastmcp_slim/fastmcp/server/low_level.py` (`FastMCPServerMiddleware`). *Verify:* `fastmcp_slim/fastmcp/server/low_level.py` (`FastMCPServerMiddleware`).
### Middleware observes every inbound message — New (coverage)
FastMCP's `Middleware` chain used to begin *inside* the per-method handlers, so `on_message`/`on_request`/`on_notification` only fired for messages that reached a tool/resource/prompt handler. Notifications, cancellations, and malformed or unroutable requests were invisible to middleware. `FastMCPServerMiddleware` — FastMCP's entry in the SDK's own middleware list — is now the dispatch root: it runs the `on_message`/`on_request`/`on_notification` pass for every message the interior handlers do not dispatch (all notifications including `notifications/cancelled`, `ping`, `logging/setLevel`, unknown methods, and component requests that fail validation before the handler runs). The component methods keep their interior dispatch unchanged, so `on_call_tool` and friends still receive the typed component result and a tool exception still propagates through `on_message`/`on_request` exactly where the built-in error/logging/timing middleware expect it — each hook fires exactly once per message. Multi-round (SEP-2322) calls compose cleanly with this: each round is a complete request→response cycle through the full chain, and an asking round's `call_next` returns the ask as an ordinary `InputRequiredToolResult` value (see the MRTR entry below). All thirteen built-in middleware pass their suites unmodified. See [What middleware sees](/servers/middleware#what-middleware-sees).
*Verify:* `fastmcp_slim/fastmcp/server/low_level.py` (`FastMCPServerMiddleware` root dispatch, `_INTERIOR_METHODS`), `fastmcp_slim/fastmcp/server/middleware/middleware.py` (`MiddlewarePhase`, `mark_interior_dispatched`), `fastmcp_slim/fastmcp/server/server.py` (`_dispatch_component_middleware`), `tests/server/middleware/test_message_visibility.py`.
### Per-session state re-homed to the connection — Absorbed ### Per-session state re-homed to the connection — Absorbed
Because `ServerSession` is now per-request, per-session state can no longer live on the session object. The minimum logging level is re-homed to a FastMCP-side map keyed by session id (via `connection.session_id`), and `client_supports_extension` becomes a free function reading `session.client_params.capabilities`. Because `ServerSession` is now per-request, per-session state can no longer live on the session object. The minimum logging level is re-homed to a FastMCP-side map keyed by session id (via `connection.session_id`), and `client_supports_extension` becomes a free function reading `session.client_params.capabilities`.

View file

@ -98,6 +98,22 @@ Rather than processing every message identically, FastMCP provides specialized h
When a client calls a tool, the middleware chain processes `on_message` first, then `on_request`, then `on_call_tool`. This hierarchy lets you target exactly the right scope—use `on_message` for logging everything, `on_request` for authentication, and `on_call_tool` for tool-specific behavior. When a client calls a tool, the middleware chain processes `on_message` first, then `on_request`, then `on_call_tool`. This hierarchy lets you target exactly the right scope—use `on_message` for logging everything, `on_request` for authentication, and `on_call_tool` for tool-specific behavior.
### What middleware sees
<VersionBadge version="4.0.0" />
Dispatch begins in the SDK's middleware layer — the single point every inbound message passes through. As a result, `on_message`, `on_request`, and `on_notification` observe **every** message a client sends, including the ones that never reach a tool, resource, or prompt handler:
- **Notifications** such as `notifications/cancelled`, `notifications/initialized`, and `notifications/progress` reach `on_message` and `on_notification`.
- **Cancellations** are observed as a `notifications/cancelled` message. The connection applies the cancellation itself and then hands the notification to your middleware.
- **Malformed or unroutable requests**—an unknown method, or a `tools/call` whose params fail validation before the tool runs—reach `on_message` and `on_request` as a raised error propagating through `call_next`, so logging and error-handling middleware record them.
The operation hooks (`on_call_tool`, `on_list_tools`, and the rest) fire exactly once per request, and their `call_next` still returns the typed component result—a `ToolResult`, a `list[Tool]`, and so on—so a tool exception propagates through `on_call_tool`, `on_request`, and `on_message` exactly where error, logging, and timing middleware expect it.
#### Multi-round tool calls
A guard tool asks the client for input by returning an `InputRequiredResult` (see [Elicitation on the modern protocol](/servers/elicitation#elicitation-on-the-modern-protocol)). Each round of a multi-round call is a complete request→response cycle that runs the **full middleware chain**: `on_call_tool` fires once per round, and on an asking round `call_next` returns the ask as that round's ordinary result value—an `InputRequiredToolResult`, a `ToolResult` subclass. Nothing is raised and nothing is held open, so default middleware completes normally on every round (logging logs the ask, timing times it, error handling does not fire—an ask is a legitimate result, not an error). Middleware that needs to treat an ask differently identifies it with an `isinstance(result, InputRequiredToolResult)` check; see [Middleware and multi-round calls](/servers/elicitation#middleware) for a worked example.
### Hook Signature ### Hook Signature
Every hook follows the same pattern: Every hook follows the same pattern:

View file

@ -3,6 +3,7 @@ from __future__ import annotations
import weakref import weakref
from collections.abc import Iterator, Mapping from collections.abc import Iterator, Mapping
from contextlib import contextmanager from contextlib import contextmanager
from dataclasses import replace
from typing import TYPE_CHECKING, Any, cast from typing import TYPE_CHECKING, Any, cast
import mcp_types import mcp_types
@ -37,6 +38,75 @@ if TYPE_CHECKING:
logger = get_logger(__name__) 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
# reaching it. Every other message is dispatched here at the root.
_INTERIOR_METHODS = frozenset(
{
"tools/call",
"tools/list",
"resources/read",
"resources/list",
"resources/templates/list",
"prompts/get",
"prompts/list",
}
)
def _raw_message(ctx: ServerRequestContext) -> Any:
"""The message payload handed to the root dispatch's ``on_message``/``on_request`` pass.
The raw inbound params mapping is used verbatim rather than a validated,
typed request model. This is deliberate: the outer pass must observe *every*
message, including malformed or unroutable ones, and reconstructing a typed
model would raise on exactly those messages and hide them from the hooks.
The method and request/notification kind are carried on the
``MiddlewareContext`` itself, so observation middleware still has everything
it needs.
"""
params = ctx.params
if isinstance(params, Mapping):
return dict(params)
return {} if params is None else params
def _forward_ctx(
ctx: ServerRequestContext, mw_ctx: Any, original: Any
) -> ServerRequestContext:
"""Fold middleware edits to the *message* back into the SDK context.
The outer pass hands middleware a *copy* of the raw params (see
``_raw_message``), so a hook that follows the documented inspect/modify
contract mutating ``context.message`` or passing ``context.copy(message=...)``
to ``call_next`` would otherwise have its edits silently dropped when the
bridge dispatched the original context. Rewriting through
``dataclasses.replace`` is how the SDK documents altering what the handler
sees. An untouched message forwards the original context unchanged.
``ctx.method`` is deliberately *not* rewritable here. Dispatch has already
branched on the method to decide that this message has no interior handler,
so redirecting it now say, turning a ``ping`` into a ``tools/list`` would
hand it to a component handler that runs the FastMCP chain a second time,
firing ``on_message`` and raw ``__call__`` overrides twice for one message
and duplicating whatever side effects (rate limiting, authorization,
logging) they carry. Rewriting the method is not part of the documented
middleware contract; only the message is.
"""
message = mw_ctx.message
if isinstance(message, Mapping):
params: Mapping[str, Any] | None = dict(message)
# `_raw_message` renders absent params as `{}`; keep that distinction so
# an untouched notification still dispatches with `params=None`.
if ctx.params is None and message == original and not message:
params = None
else:
params = ctx.params
if params == ctx.params:
return ctx
return replace(ctx, params=params)
def client_supports_extension(session: ServerSession, extension_id: str) -> bool: def client_supports_extension(session: ServerSession, extension_id: str) -> bool:
"""Check whether the connected client supports a given MCP extension. """Check whether the connected client supports a given MCP extension.
@ -68,15 +138,40 @@ def client_supports_extension(session: ServerSession, extension_id: str) -> bool
class FastMCPServerMiddleware: class FastMCPServerMiddleware:
"""SDK v2 server middleware that routes ``initialize`` through FastMCP middleware. """Root dispatch for the FastMCP middleware chain, in the SDK's middleware layer.
v2 no longer lets FastMCP subclass ``ServerSession`` (the runner constructs v2 no longer lets FastMCP subclass ``ServerSession`` (the runner constructs
it per request), so the old ``MiddlewareServerSession._received_request`` it per request), so the old ``MiddlewareServerSession._received_request``
override is replaced by a ``ServerMiddleware``. This middleware binds the override is replaced by a ``ServerMiddleware`` an ordinary entry in the
FastMCP request-context ContextVar for the whole chain (covering SDK's own middleware list. Sitting at the root of dispatch, this
``initialize``, where no handler adapter runs) and routes the initialize is the single entry point through which *every* inbound message flows
request through the FastMCP middleware chain so ``on_initialize`` hooks fire requests, notifications, cancellations, ``initialize``, and even malformed or
and can observe the ``InitializeResult`` or veto with ``MCPError``. unroutable messages the SDK can still hand us. It binds the FastMCP
request-context ContextVar and re-applies the app-scoped ``SharedContext`` for
the whole chain, then runs the FastMCP ``Middleware`` chain so
``on_message`` / ``on_request`` / ``on_notification`` observe the message.
Dispatch shapes:
- ``initialize`` runs the *whole* FastMCP chain here (``on_message`` ->
``on_request`` -> ``on_initialize``) because there is no interior handler
adapter for it: the SDK builds the ``InitializeResult`` directly, so this is
the only place ``on_initialize`` can observe it or veto with ``MCPError``.
- The component methods (``tools/call``, ``tools/list``, ``resources/read``,
...) still run their FastMCP chain *interior*, in the handler adapter, where
``on_call_tool`` receives the typed component result and a tool exception
propagates through ``on_message``/``on_request`` exactly where the built-in
error/logging/timing middleware expect it. The root dispatch does not re-run the
chain for these it only steps in when such a request fails *before* the
interior runs (malformed params, routing), so ``on_message`` still observes
the failure.
- Every other message all notifications (including ``notifications/cancelled``
and ``notifications/initialized``), ``ping``, ``logging/setLevel``, and any
unroutable/non-component request has no interior FastMCP dispatch, so the
root dispatch runs the ``"outer"`` pass (``on_message`` plus
``on_request``/``on_notification``) here, wrapping the real SDK dispatch.
This closes the long-standing gap where these messages were invisible to
FastMCP middleware.
""" """
def __init__(self, fastmcp: FastMCP): def __init__(self, fastmcp: FastMCP):
@ -93,13 +188,88 @@ class FastMCPServerMiddleware:
bind_request_context(ctx), bind_request_context(ctx),
self._seam_span(fastmcp, ctx), self._seam_span(fastmcp, ctx),
): ):
# Only initialize requests (request_id present) go through FastMCP if fastmcp is None:
# middleware here; every other request already binds the context in return await call_next(ctx)
# its own adapter, so we just pass through.
if ctx.method == "initialize" and ctx.request_id is not None: if ctx.method == "initialize" and ctx.request_id is not None:
if fastmcp is not None: return await self._run_initialize_mw(fastmcp, ctx, call_next)
return await self._run_initialize_mw(fastmcp, ctx, call_next) if ctx.request_id is not None and ctx.method in _INTERIOR_METHODS:
return await self._dispatch_component(fastmcp, ctx, call_next)
return await self._run_outer_mw(fastmcp, ctx, call_next, _raise=None)
async def _dispatch_component(
self,
fastmcp: FastMCP,
ctx: ServerRequestContext,
call_next: CallNext,
) -> HandlerResult:
"""Delegate a component request to the interior chain, covering early failures.
The interior handler adapter runs the FastMCP chain itself and records
``_interior_dispatched``. If the request instead fails before reaching it
(malformed params, method routing), the flag stays False and no hook fired
so the root dispatch runs the ``"outer"`` pass to observe the failure, re-raising
the original error inside it so ``on_message``/``on_request`` see it.
"""
from fastmcp.server.middleware.middleware import _interior_dispatched
token = _interior_dispatched.set(False)
try:
return await call_next(ctx) return await call_next(ctx)
except (MCPError, ValidationError) as exc:
if _interior_dispatched.get():
raise
return await self._run_outer_mw(fastmcp, ctx, call_next, _raise=exc)
finally:
_interior_dispatched.reset(token)
async def _run_outer_mw(
self,
fastmcp: FastMCP,
ctx: ServerRequestContext,
call_next: CallNext,
*,
_raise: BaseException | None,
) -> HandlerResult:
"""Run the method-agnostic (``on_message``/``on_request``) hook pass.
``call_next`` bridges to the real SDK dispatch (request-state boundary,
params validation, the notification handler), so these hooks observe the
actual wire outcome: a notification returns ``None``, an unroutable request
raises through ``call_next``. Message edits are folded back in through
``_forward_ctx``.
When ``_raise`` is set the operation already failed before the interior
ran, and the bridge re-raises it rather than dispatching. This pass is
the *observation* path for that failure, not a retry: re-dispatching a
corrected component request would run its handler, which runs the FastMCP
chain interior, firing ``on_message`` and the raw ``__call__`` override a
second time for one message. A hook cannot repair a malformed
``tools/call`` from here it sees the failure, and the failure stands.
"""
from fastmcp.server.context import Context
from fastmcp.server.middleware.middleware import MiddlewareContext
is_notification = ctx.request_id is None
original_message = _raw_message(ctx)
async def root_call_next(_mw_ctx: MiddlewareContext) -> HandlerResult:
if _raise is not None:
raise _raise
return await call_next(_forward_ctx(ctx, _mw_ctx, original_message))
async with Context(fastmcp=fastmcp, session=ctx.session) as fastmcp_ctx:
mw_context = MiddlewareContext(
message=original_message,
source="client",
type="notification" if is_notification else "request",
method=ctx.method,
fastmcp_context=fastmcp_ctx,
)
return await fastmcp._run_middleware(
mw_context,
cast("FastMCPCallNext[Any, Any]", root_call_next),
phase="outer",
)
@contextmanager @contextmanager
def _seam_span( def _seam_span(

View file

@ -2,6 +2,7 @@ from __future__ import annotations
import logging import logging
from collections.abc import Awaitable, Callable, Sequence from collections.abc import Awaitable, Callable, Sequence
from contextvars import ContextVar
from dataclasses import dataclass, field, replace from dataclasses import dataclass, field, replace
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import ( from typing import (
@ -32,6 +33,53 @@ __all__ = [
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
MiddlewarePhase = Literal["all", "outer", "typed"]
"""Which slice of a middleware's hooks to run in a single dispatch pass.
- ``"all"`` runs the whole chain in one pass (``on_message`` -> ``on_request`` /
``on_notification`` -> the typed per-method hook). This is what the interior
component methods (``call_tool``, ``list_tools``, ...) run for the methods they
serve, and what the ``initialize`` request runs at the dispatch root.
- ``"outer"`` runs only ``on_message`` and ``on_request``/``on_notification``.
The root dispatch (in the SDK's middleware layer) runs this pass for the messages the interior never
dispatches (notifications, cancellations, unroutable/non-component requests,
and pre-handler failures), so ``on_message`` observes *every* inbound message
without double-firing for the component methods the interior already covers.
- ``"typed"`` runs only the per-method hook. Reserved for a future full split;
no current dispatch path uses it.
"""
_interior_dispatched: ContextVar[bool] = ContextVar(
"fastmcp_interior_dispatched", default=False
)
"""Set to True by an interior component dispatch when it runs its middleware chain.
The root dispatch reads this to tell whether the FastMCP middleware
chain already fired *inside* the wire request (so ``on_message``/``on_request``
were observed there including any tool exception, exactly where the built-in
error/logging/timing middleware expect them). It is only consulted for the
component methods: if such a request fails *before* the interior runs (malformed
params, routing), the flag stays False and the root dispatch observes the failure itself.
"""
def mark_interior_dispatched() -> None:
"""Record that an interior component middleware chain ran for this message."""
_interior_dispatched.set(True)
_dispatch_phase: ContextVar[MiddlewarePhase] = ContextVar(
"fastmcp_dispatch_phase", default="all"
)
"""The dispatch phase for the middleware chain currently running.
Set by ``FastMCP._run_middleware`` around each chain execution and read by
``Middleware.__call__``, so the phase never appears in the middleware call
signature user middleware overriding the documented
``__call__(context, call_next)`` keeps working unchanged.
"""
T = TypeVar("T", default=Any) T = TypeVar("T", default=Any)
R = TypeVar("R", covariant=True, default=Any) R = TypeVar("R", covariant=True, default=Any)
@ -93,47 +141,61 @@ class Middleware:
context: MiddlewareContext[T], context: MiddlewareContext[T],
call_next: CallNext[T, Any], call_next: CallNext[T, Any],
) -> Any: ) -> Any:
"""Main entry point that orchestrates the pipeline.""" """Main entry point that orchestrates the pipeline.
The dispatch phase which slice of the hooks runs (see
``MiddlewarePhase``) is read from ``_dispatch_phase`` rather than
passed as an argument, so middleware that overrides this method with the
documented ``(context, call_next)`` signature keeps working unchanged.
Such an override runs once per message regardless of phase, which
matches its pre-existing behavior.
"""
handler_chain = await self._dispatch_handler( handler_chain = await self._dispatch_handler(
context, context,
call_next=call_next, call_next=call_next,
phase=_dispatch_phase.get(),
) )
return await handler_chain(context) return await handler_chain(context)
async def _dispatch_handler( async def _dispatch_handler(
self, context: MiddlewareContext[Any], call_next: CallNext[Any, Any] self,
context: MiddlewareContext[Any],
call_next: CallNext[Any, Any],
phase: MiddlewarePhase = "all",
) -> CallNext[Any, Any]: ) -> CallNext[Any, Any]:
"""Builds a chain of handlers for a given message.""" """Builds a chain of handlers for a given message and dispatch phase."""
handler = call_next handler = call_next
match context.method: if phase in ("all", "typed"):
case "initialize": match context.method:
handler = make_handler_wrapper(self.on_initialize, handler) case "initialize":
case "tools/call": handler = make_handler_wrapper(self.on_initialize, handler)
handler = make_handler_wrapper(self.on_call_tool, handler) case "tools/call":
case "resources/read": handler = make_handler_wrapper(self.on_call_tool, handler)
handler = make_handler_wrapper(self.on_read_resource, handler) case "resources/read":
case "prompts/get": handler = make_handler_wrapper(self.on_read_resource, handler)
handler = make_handler_wrapper(self.on_get_prompt, handler) case "prompts/get":
case "tools/list": handler = make_handler_wrapper(self.on_get_prompt, handler)
handler = make_handler_wrapper(self.on_list_tools, handler) case "tools/list":
case "resources/list": handler = make_handler_wrapper(self.on_list_tools, handler)
handler = make_handler_wrapper(self.on_list_resources, handler) case "resources/list":
case "resources/templates/list": handler = make_handler_wrapper(self.on_list_resources, handler)
handler = make_handler_wrapper( case "resources/templates/list":
self.on_list_resource_templates, handler = make_handler_wrapper(
handler, self.on_list_resource_templates,
) handler,
case "prompts/list": )
handler = make_handler_wrapper(self.on_list_prompts, handler) case "prompts/list":
handler = make_handler_wrapper(self.on_list_prompts, handler)
match context.type: if phase in ("all", "outer"):
case "request": match context.type:
handler = make_handler_wrapper(self.on_request, handler) case "request":
case "notification": handler = make_handler_wrapper(self.on_request, handler)
handler = make_handler_wrapper(self.on_notification, handler) case "notification":
handler = make_handler_wrapper(self.on_notification, handler)
handler = make_handler_wrapper(self.on_message, handler) handler = make_handler_wrapper(self.on_message, handler)
return handler return handler

View file

@ -71,6 +71,11 @@ from fastmcp.server.caching import build_cache_hints
from fastmcp.server.lifespan import Lifespan from fastmcp.server.lifespan import Lifespan
from fastmcp.server.low_level import LowLevelServer from fastmcp.server.low_level import LowLevelServer
from fastmcp.server.middleware import CallNext, Middleware, MiddlewareContext from fastmcp.server.middleware import CallNext, Middleware, MiddlewareContext
from fastmcp.server.middleware.middleware import (
MiddlewarePhase,
_dispatch_phase,
mark_interior_dispatched,
)
from fastmcp.server.mixins import LifespanMixin, MCPOperationsMixin, TransportMixin from fastmcp.server.mixins import LifespanMixin, MCPOperationsMixin, TransportMixin
from fastmcp.server.providers import LocalProvider, Provider from fastmcp.server.providers import LocalProvider, Provider
from fastmcp.server.providers.aggregate import AggregateProvider from fastmcp.server.providers.aggregate import AggregateProvider
@ -571,8 +576,18 @@ class FastMCP(
self, self,
context: MiddlewareContext[Any], context: MiddlewareContext[Any],
call_next: CallNext[Any, Any], call_next: CallNext[Any, Any],
*,
phase: MiddlewarePhase = "all",
) -> Any: ) -> Any:
"""Builds and executes the middleware chain.""" """Builds and executes the middleware chain for a single dispatch phase.
``phase`` selects whether a pass runs only the method-agnostic hooks
(``"outer"``, at the root dispatch) or only the typed per-method hook
(``"typed"``, interior); it defaults to ``"all"`` for the direct
programmatic path. It is conveyed through the ``_dispatch_phase``
ContextVar rather than the middleware call signature, so user middleware
overriding the documented ``__call__(context, call_next)`` is unaffected.
"""
chain = call_next chain = call_next
for mw in reversed(self.middleware): for mw in reversed(self.middleware):
next_chain: CallNext[Any, Any] = chain next_chain: CallNext[Any, Any] = chain
@ -585,7 +600,30 @@ class FastMCP(
return await mw(context, call_next) return await mw(context, call_next)
chain = cast(CallNext[Any, Any], wrapped) chain = cast(CallNext[Any, Any], wrapped)
return await chain(context) token = _dispatch_phase.set(phase)
try:
return await chain(context)
finally:
_dispatch_phase.reset(token)
async def _dispatch_component_middleware(
self,
context: MiddlewareContext[Any],
call_next: CallNext[Any, Any],
) -> Any:
"""Run the interior middleware chain for a component operation.
This is the dispatch site for the component methods (``tools/call``,
``tools/list``, ``resources/read``, ...). It runs the whole FastMCP chain
(``on_message`` -> ``on_request`` -> the typed per-method hook) in one
pass, so error-observing middleware see a tool exception propagate through
``on_message``/``on_request`` exactly as they always have. It also records
(via ``mark_interior_dispatched``) that the chain fired for this wire
message, so the root dispatch knows not to observe it a second
time.
"""
mark_interior_dispatched()
return await self._run_middleware(context, call_next, phase="all")
def add_middleware(self, middleware: Middleware) -> None: def add_middleware(self, middleware: Middleware) -> None:
self.middleware.append(middleware) self.middleware.append(middleware)
@ -695,7 +733,7 @@ class FastMCP(
method="tools/list", method="tools/list",
fastmcp_context=ctx, fastmcp_context=ctx,
) )
return await self._run_middleware( return await self._dispatch_component_middleware(
context=mw_context, context=mw_context,
call_next=lambda context: self.list_tools(run_middleware=False), call_next=lambda context: self.list_tools(run_middleware=False),
) )
@ -831,7 +869,7 @@ class FastMCP(
method="resources/list", method="resources/list",
fastmcp_context=ctx, fastmcp_context=ctx,
) )
return await self._run_middleware( return await self._dispatch_component_middleware(
context=mw_context, context=mw_context,
call_next=lambda context: self.list_resources(run_middleware=False), call_next=lambda context: self.list_resources(run_middleware=False),
) )
@ -966,7 +1004,7 @@ class FastMCP(
method="resources/templates/list", method="resources/templates/list",
fastmcp_context=ctx, fastmcp_context=ctx,
) )
return await self._run_middleware( return await self._dispatch_component_middleware(
context=mw_context, context=mw_context,
call_next=lambda context: self.list_resource_templates( call_next=lambda context: self.list_resource_templates(
run_middleware=False run_middleware=False
@ -1100,7 +1138,7 @@ class FastMCP(
method="prompts/list", method="prompts/list",
fastmcp_context=ctx, fastmcp_context=ctx,
) )
return await self._run_middleware( return await self._dispatch_component_middleware(
context=mw_context, context=mw_context,
call_next=lambda context: self.list_prompts(run_middleware=False), call_next=lambda context: self.list_prompts(run_middleware=False),
) )
@ -1302,7 +1340,7 @@ class FastMCP(
method="tools/call", method="tools/call",
fastmcp_context=ctx, fastmcp_context=ctx,
) )
return await self._run_middleware( return await self._dispatch_component_middleware(
context=mw_context, context=mw_context,
call_next=lambda context: self.call_tool( call_next=lambda context: self.call_tool(
context.message.name, context.message.name,
@ -1473,7 +1511,7 @@ class FastMCP(
method="resources/read", method="resources/read",
fastmcp_context=ctx, fastmcp_context=ctx,
) )
return await self._run_middleware( return await self._dispatch_component_middleware(
context=mw_context, context=mw_context,
call_next=lambda context: self.read_resource( call_next=lambda context: self.read_resource(
str(context.message.uri), str(context.message.uri),
@ -1674,7 +1712,7 @@ class FastMCP(
method="prompts/get", method="prompts/get",
fastmcp_context=ctx, fastmcp_context=ctx,
) )
return await self._run_middleware( return await self._dispatch_component_middleware(
context=mw_context, context=mw_context,
call_next=lambda context: self.render_prompt( call_next=lambda context: self.render_prompt(
context.message.name, context.message.name,

View file

@ -0,0 +1,397 @@
"""Middleware message visibility (v4 D3, the middleware hybrid rebase).
Dispatch begins in the SDK's middleware layer, so ``on_message``/``on_request``/
``on_notification`` observe *every* inbound message including the ones that
never reach a FastMCP handler (notifications, cancellations, and
malformed/unroutable requests) and were therefore invisible to FastMCP
middleware before. The typed per-method hooks keep firing exactly once, interior,
where ``call_next`` yields the typed component result.
"""
from typing import Any
import mcp_types
import pytest
from mcp.shared.exceptions import MCPError
from mcp_types import ElicitRequest, ElicitRequestFormParams, InputRequiredResult
from fastmcp import Client, FastMCP
from fastmcp.server.context import Context
from fastmcp.server.middleware import CallNext, Middleware, MiddlewareContext
from fastmcp.tools.base import InputRequiredToolResult
class HookRecorder(Middleware):
"""Records ``(hook, method)`` before delegating, so a hook is captured even
when ``call_next`` raises (a pre-handler failure)."""
def __init__(self) -> None:
self.records: list[tuple[str, str | None]] = []
async def on_message(self, context: MiddlewareContext, call_next: CallNext) -> Any:
self.records.append(("on_message", context.method))
return await call_next(context)
async def on_request(self, context: MiddlewareContext, call_next: CallNext) -> Any:
self.records.append(("on_request", context.method))
return await call_next(context)
async def on_notification(
self, context: MiddlewareContext, call_next: CallNext
) -> Any:
self.records.append(("on_notification", context.method))
return await call_next(context)
async def on_call_tool(
self, context: MiddlewareContext, call_next: CallNext
) -> Any:
self.records.append(("on_call_tool", context.method))
return await call_next(context)
async def on_list_tools(
self, context: MiddlewareContext, call_next: CallNext
) -> Any:
self.records.append(("on_list_tools", context.method))
return await call_next(context)
def _adder() -> FastMCP:
server = FastMCP("AdderServer")
@server.tool
def add(a: int, b: int) -> int:
return a + b
return server
class TestNotificationVisibility:
async def test_client_cancelled_notification_reaches_on_message(self):
"""A ``notifications/cancelled`` from the client is observed by
``on_message`` and ``on_notification`` it never reaches a FastMCP
handler, so before the rebase it was invisible to middleware."""
server = _adder()
recorder = HookRecorder()
server.add_middleware(recorder)
async with Client(server) as client:
await client.session.send_notification(
mcp_types.CancelledNotification(
params=mcp_types.CancelledNotificationParams(
request_id="never-issued"
)
)
)
# Round-trip on the same connection so the notification is dispatched
# before we assert (in-order delivery).
await client.call_tool("add", {"a": 1, "b": 2})
assert ("on_message", "notifications/cancelled") in recorder.records
assert ("on_notification", "notifications/cancelled") in recorder.records
async def test_client_progress_notification_reaches_on_message(self):
"""A generic client notification is observed by ``on_message``."""
server = _adder()
recorder = HookRecorder()
server.add_middleware(recorder)
async with Client(server) as client:
await client.session.send_notification(
mcp_types.ProgressNotification(
params=mcp_types.ProgressNotificationParams(
progress_token="tok", progress=1.0
)
)
)
await client.call_tool("add", {"a": 1, "b": 2})
assert ("on_message", "notifications/progress") in recorder.records
class TestUnroutableAndMalformed:
async def test_unroutable_method_observed_by_on_message(self):
"""An unknown method fails routing before any handler; the root dispatch still
runs ``on_message``/``on_request`` around the failure."""
server = _adder()
recorder = HookRecorder()
server.add_middleware(recorder)
async with Client(server) as client:
with pytest.raises(MCPError):
await client.session._dispatcher.send_raw_request(
"does/not/exist", {}, {}
)
assert ("on_message", "does/not/exist") in recorder.records
assert ("on_request", "does/not/exist") in recorder.records
async def test_malformed_component_params_observed_by_on_message(self):
"""A ``tools/call`` with malformed params fails validation before the
interior handler runs, so no typed hook fires but the root dispatch observes the
failure through ``on_message``, and ``on_call_tool`` does not fire."""
server = _adder()
recorder = HookRecorder()
server.add_middleware(recorder)
async with Client(server) as client:
with pytest.raises(MCPError):
await client.session._dispatcher.send_raw_request(
"tools/call", {"not_a_valid": "param"}, {}
)
assert ("on_message", "tools/call") in recorder.records
assert ("on_call_tool", "tools/call") not in recorder.records
class TestSingleFire:
async def test_each_hook_fires_once_per_component_call(self):
"""One ``tools/call`` fires ``on_message`` once and ``on_call_tool`` once —
the interior dispatch is the single entry for component methods; the root dispatch
does not double-run it."""
server = _adder()
recorder = HookRecorder()
server.add_middleware(recorder)
async with Client(server) as client:
await client.call_tool("add", {"a": 1, "b": 2})
on_message = [r for r in recorder.records if r == ("on_message", "tools/call")]
on_call_tool = [
r for r in recorder.records if r == ("on_call_tool", "tools/call")
]
assert len(on_message) == 1
assert len(on_call_tool) == 1
class TestRawMiddlewareCompatibility:
"""Middleware may override ``__call__(context, call_next)`` — the documented
raw signature. The dispatch phase travels out-of-band, so that contract is
unchanged and such middleware keeps working."""
async def test_raw_call_override_still_works(self):
seen: list[str | None] = []
class RawMiddleware(Middleware):
async def __call__(self, context, call_next):
seen.append(context.method)
return await call_next(context)
server = _adder()
server.add_middleware(RawMiddleware())
async with Client(server) as client:
result = await client.call_tool("add", {"a": 1, "b": 2})
await client.session.send_notification(
mcp_types.ProgressNotification(
params=mcp_types.ProgressNotificationParams(
progress_token="tok", progress=1.0
)
)
)
await client.call_tool("add", {"a": 1, "b": 2})
assert result.data == 3
# It observes both a component call and a message the root dispatch owns.
assert "tools/call" in seen
assert "notifications/progress" in seen
class TestMessageModification:
"""The root dispatch hands middleware a copy of the raw params, so edits made
through the documented inspect/modify contract must be folded back into the
SDK context before the real dispatch runs."""
async def test_modified_message_reaches_sdk_dispatch(self):
"""A ``logging/setLevel`` carrying an invalid level fails params
validation inside ``call_next``. Middleware that rewrites the message to
a valid level makes the request succeed which only happens if the edit
is actually forwarded."""
class RewriteLevel(Middleware):
async def on_message(self, context, call_next):
if context.method == "logging/setLevel":
context.message["level"] = "debug"
return await call_next(context)
server = _adder()
server.add_middleware(RewriteLevel())
async with Client(server) as client:
await client.session._dispatcher.send_raw_request(
"logging/setLevel", {"level": "not-a-valid-level"}, {}
)
async def test_unmodified_message_dispatches_unchanged(self):
"""An observation-only hook leaves dispatch untouched."""
server = _adder()
recorder = HookRecorder()
server.add_middleware(recorder)
async with Client(server) as client:
await client.session._dispatcher.send_raw_request(
"logging/setLevel", {"level": "debug"}, {}
)
assert ("on_message", "logging/setLevel") in recorder.records
async def test_method_rewrite_does_not_redirect_dispatch(self):
"""Only the message is rewritable. Dispatch has already branched on the
method to decide this message has no interior handler, so honoring a
rewrite into a component method would hand it to a handler that runs the
chain again firing the generic hooks twice for one message. The
rewrite is ignored and the invariant holds."""
class RewriteMethod(Middleware):
async def on_message(self, context, call_next):
if context.method == "ping":
return await call_next(context.copy(method="tools/list"))
return await call_next(context)
server = _adder()
recorder = HookRecorder()
# Recorder outermost, so it observes the message as it arrived; the
# rewriter runs inside it.
server.add_middleware(recorder)
server.add_middleware(RewriteMethod())
async with Client(server) as client:
await client.session._dispatcher.send_raw_request("ping", {}, {})
# Had the rewrite redirected dispatch, the component handler would have
# run the chain again — a second on_message, plus an on_list_tools for a
# request that was never a tools/list.
assert [r for r in recorder.records if r == ("on_message", "ping")] == [
("on_message", "ping")
]
assert not [r for r in recorder.records if r == ("on_message", "tools/list")]
assert not [r for r in recorder.records if r[0] == "on_list_tools"]
async def test_failed_component_request_is_observed_not_retried(self):
"""A component request that dies in validation reaches the hooks as a
failure. A hook cannot repair it from here: re-dispatching would run the
handler and fire the generic hooks a second time, so the failure stands
and ``on_message`` sees it exactly once."""
class RepairAttempt(Middleware):
async def on_message(self, context, call_next):
if context.method == "tools/call":
context.message["name"] = "add"
context.message["arguments"] = {"a": 1, "b": 2}
return await call_next(context)
server = _adder()
recorder = HookRecorder()
server.add_middleware(RepairAttempt())
server.add_middleware(recorder)
async with Client(server) as client:
with pytest.raises(MCPError):
await client.session._dispatcher.send_raw_request(
"tools/call", {"not_a_valid": "param"}, {}
)
calls = [r for r in recorder.records if r == ("on_message", "tools/call")]
assert len(calls) == 1
assert ("on_call_tool", "tools/call") not in recorder.records
def _guard_server() -> FastMCP:
server = FastMCP("Guard")
@server.tool
async def guard(ctx: Context) -> str | InputRequiredResult:
if ctx.input_responses is None:
request = ElicitRequest(
method="elicitation/create",
params=ElicitRequestFormParams(
message="Your name?",
requested_schema={
"type": "object",
"properties": {"name": {"type": "string"}},
"required": ["name"],
},
),
)
return InputRequiredResult(
result_type="input_required",
input_requests={"name": request},
request_state=None,
)
return "done"
return server
class TestAskVisibility:
async def test_ask_is_the_observed_result_of_a_guard_leg(self):
"""Each MRTR leg is a complete request→response cycle: a guard tool's ask
is the full, legitimate result of that leg. A component hook's
``call_next`` returns it as an ordinary value an
``InputRequiredToolResult`` (a ``ToolResult`` subclass) so the hook
completes normally and can identify the ask by ``isinstance``."""
class AskProbe(Middleware):
def __init__(self) -> None:
self.entered = 0
self.results: list[Any] = []
async def on_call_tool(
self, context: MiddlewareContext, call_next: CallNext
) -> Any:
self.entered += 1
result = await call_next(context)
self.results.append(result)
return result
server = _guard_server()
probe = AskProbe()
server.add_middleware(probe)
async with Client(server, mode="auto") as client:
result = await client.session.call_tool(
"guard", {}, allow_input_required=True
)
assert isinstance(result, InputRequiredResult)
assert probe.entered == 1
# The hook completed and observed the ask as the leg's result value.
assert len(probe.results) == 1
assert isinstance(probe.results[0], InputRequiredToolResult)
async def test_hooks_fire_once_per_round_across_a_continuation(self):
"""The fires-once invariant holds across a continuation — the one place
root dispatch and MRTR genuinely meet. Each round is its own complete
requestresponse cycle, so answering the ask runs the chain a second
time in full rather than double-firing on either round."""
server = _guard_server()
recorder = HookRecorder()
server.add_middleware(recorder)
async with Client(server, mode="auto") as client:
ask = await client.session.call_tool("guard", {}, allow_input_required=True)
assert isinstance(ask, InputRequiredResult)
answered = await client.session.call_tool(
"guard",
{},
input_responses={
"name": {"action": "accept", "content": {"name": "Ada"}}
},
request_state=ask.request_state,
allow_input_required=True,
)
assert isinstance(answered, mcp_types.CallToolResult)
# Two rounds — the ask and the answer — and exactly one chain per round.
on_message = [r for r in recorder.records if r == ("on_message", "tools/call")]
on_call_tool = [
r for r in recorder.records if r == ("on_call_tool", "tools/call")
]
assert len(on_message) == 2
assert len(on_call_tool) == 2
class TestSchedulingProbe:
async def test_trivial_noop(self):
"""Temporary probe: does merely adding a 7th test destabilize the run?"""
assert True