Bridge initialize InitializeResult through middleware; re-home session state to Connection

This commit is contained in:
Jeremiah Lowin 2026-07-05 20:42:28 -04:00
commit 92b0022c2d
No known key found for this signature in database
2 changed files with 74 additions and 19 deletions

View file

@ -672,23 +672,42 @@ class Context:
"This typically means you're outside a request context."
)
# Check for cached session ID
session_id = getattr(session, "_fastmcp_state_prefix", None)
if session_id is not None:
return session_id
# In SDK v2 the ServerSession is constructed fresh per request, so the
# stable per-client identity lives on the underlying Connection, which
# persists for the whole client session. Cache the state prefix on the
# connection (its `session_id` for HTTP, its `state` dict otherwise) so
# session-scoped state survives across tool calls.
connection = getattr(session, "_connection", None)
# For HTTP, try to get from header
if request_ctx is not None:
# Check for a cached prefix on the stable connection (or the session, as
# a fallback for on_initialize where only a raw session is available).
if connection is not None:
cached = connection.state.get("_fastmcp_state_prefix")
if cached is not None:
return cached
session_cached = getattr(session, "_fastmcp_state_prefix", None)
if session_cached is not None:
return session_cached
# For HTTP, prefer the connection's negotiated session id, then the
# incoming request header.
session_id: str | None = None
if connection is not None:
session_id = connection.session_id
if session_id is None and request_ctx is not None:
request = request_ctx.request
if request:
session_id = request.headers.get("mcp-session-id")
# For STDIO/SSE/in-memory, generate a UUID
# For STDIO/SSE/in-memory, generate a UUID.
if session_id is None:
session_id = str(uuid4())
# Cache on session for consistency
session._fastmcp_state_prefix = session_id # type: ignore[attr-defined] # ty:ignore[unresolved-attribute]
# Cache on the stable connection (falling back to the session).
if connection is not None:
connection.state["_fastmcp_state_prefix"] = session_id
else:
session._fastmcp_state_prefix = session_id # type: ignore[attr-defined] # ty:ignore[unresolved-attribute]
return session_id
@property

View file

@ -1,7 +1,7 @@
from __future__ import annotations
import weakref
from collections.abc import Iterator
from collections.abc import Iterator, Mapping
from contextlib import contextmanager
from typing import TYPE_CHECKING, Any, cast
@ -22,6 +22,8 @@ from mcp.server.lowlevel.server import (
from mcp.server.models import InitializationOptions
from mcp.server.session import ServerSession
from mcp.server.stdio import stdio_server as stdio_server
from mcp.shared.exceptions import MCPError
from pydantic import ValidationError
from fastmcp.apps.config import UI_EXTENSION_ID
from fastmcp.utilities.logging import get_logger
@ -121,15 +123,33 @@ class FastMCPServerMiddleware:
init_message = mcp_types.InitializeRequest.model_validate(
{"method": "initialize", "params": params}, by_name=False
)
except Exception:
except ValidationError:
init_message = None
# Track the initialize result produced by the SDK chain so a FastMCP
# middleware that raises `MCPError` *after* `call_next` can be
# logged-and-swallowed (the result is already committed) rather than
# producing a duplicate error response — preserving the pre-v2 contract.
captured_result: mcp_types.InitializeResult | None = None
call_next_completed = False
async def call_original_handler(
_mw_ctx: MiddlewareContext,
) -> HandlerResult:
) -> mcp_types.InitializeResult | None:
# call_next(ctx) runs the rest of the SDK chain, which for
# initialize returns the serialized InitializeResult dict.
return await call_next(ctx)
# initialize returns the serialized InitializeResult dict. FastMCP
# middleware `on_initialize` hooks expect a typed InitializeResult,
# so deserialize before handing control back up the FastMCP chain.
# The runner's `_dump_result` re-serializes whatever we return, so a
# returned model round-trips cleanly.
nonlocal captured_result, call_next_completed
raw = await call_next(ctx)
if isinstance(raw, mcp_types.InitializeResult):
captured_result = raw
elif isinstance(raw, Mapping):
captured_result = mcp_types.InitializeResult.model_validate(dict(raw))
call_next_completed = True
return captured_result if raw is not None else None
async with Context(fastmcp=fastmcp, session=ctx.session) as fastmcp_ctx:
mw_context = MiddlewareContext(
@ -139,10 +159,26 @@ class FastMCPServerMiddleware:
method="initialize",
fastmcp_context=fastmcp_ctx,
)
return await fastmcp._run_middleware(
mw_context,
cast("FastMCPCallNext[Any, Any]", call_original_handler),
)
try:
return await fastmcp._run_middleware(
mw_context,
cast("FastMCPCallNext[Any, Any]", call_original_handler),
)
except MCPError:
# A middleware raised after the initialize response was already
# produced: log and return the committed result instead of
# re-raising to avoid responding to initialize twice. If the
# error was raised before `call_next` succeeded, re-raise so the
# dispatcher turns it into the wire error.
if not call_next_completed:
raise
logger.warning(
"MCPError raised by FastMCP middleware after the initialize "
"response was produced; logging and not re-raising to avoid a "
"duplicate response.",
exc_info=True,
)
return captured_result
class LowLevelServer(_Server[LifespanResultT]):