Port HTTP session manager, event store, and transport run paths to SDK v2

This commit is contained in:
Jeremiah Lowin 2026-07-05 19:15:03 -04:00
commit 1a9b1c435f
No known key found for this signature in database
4 changed files with 65 additions and 63 deletions

View file

@ -16,12 +16,17 @@ from key_value.aio.stores.memory import MemoryStore
from mcp.server.streamable_http import EventCallback, EventId, EventMessage, StreamId
from mcp.server.streamable_http import EventStore as SDKEventStore
from mcp_types import JSONRPCMessage
from pydantic import TypeAdapter
from fastmcp.utilities.logging import get_logger
from fastmcp.utilities.types import FastMCPBaseModel
logger = get_logger(__name__)
# In the v2 SDK `JSONRPCMessage` is a bare union (no `.model_validate`); use a
# TypeAdapter to validate a stored dict back into the correct member.
_jsonrpc_message_adapter: TypeAdapter[JSONRPCMessage] = TypeAdapter(JSONRPCMessage)
class EventEntry(FastMCPBaseModel):
"""Stored event entry."""
@ -223,7 +228,7 @@ class EventStore(SDKEventStore):
for event_id in event_ids[start_idx:]:
event = await self._event_store.get(key=event_id)
if event and event.message:
msg = JSONRPCMessage.model_validate(event.message)
msg = _jsonrpc_message_adapter.validate_python(event.message)
await send_callback(EventMessage(msg, event.event_id))
return stream_id

View file

@ -49,6 +49,7 @@ class FastMCPStreamableHTTPSessionManager(StreamableHTTPSessionManager):
stateless: bool = False,
security_settings: TransportSecuritySettings | None = None,
retry_interval: int | None = None,
session_idle_timeout: float | None = None,
) -> None:
self._shared_event_store: EventStore | None = None
super().__init__(
@ -58,6 +59,7 @@ class FastMCPStreamableHTTPSessionManager(StreamableHTTPSessionManager):
stateless=stateless,
security_settings=security_settings,
retry_interval=retry_interval,
session_idle_timeout=session_idle_timeout,
)
@property
@ -597,6 +599,13 @@ def create_streamable_http_app(
retry_interval=retry_interval,
json_response=json_response,
stateless=stateless_http,
# FastMCP owns DNS-rebinding protection via HostOriginGuardMiddleware,
# which is more expressive and already the documented surface. Always
# disable the SDK's own protection so the two layers don't
# double-block with confusing errors from two allowlists.
security_settings=TransportSecuritySettings(
enable_dns_rebinding_protection=False
),
)
async with (
server._lifespan_manager(),

View file

@ -11,11 +11,27 @@ from unittest.mock import MagicMock
from mcp.server.auth.middleware.auth_context import auth_context_var
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser
from mcp.server.context import ServerRequestContext as RequestContext
from starlette.requests import Request
from fastmcp.server.auth import AccessToken
from fastmcp.server.dependencies import get_access_token, request_ctx
from fastmcp.server.dependencies import (
FastMCPRequestContext,
fastmcp_request_ctx,
get_access_token,
)
def _make_ctx(request: Request | None) -> FastMCPRequestContext:
return FastMCPRequestContext(
session=MagicMock(),
request_id="0",
meta=None,
request=request,
protocol_version="2025-06-18",
close_sse_stream=None,
lifespan_context=MagicMock(),
_srctx=MagicMock(meta=None),
)
class TestStaleAccessToken:
@ -62,15 +78,11 @@ class TestStaleAccessToken:
}
mock_request = Request(scope)
# Create a mock RequestContext with the request
mock_request_context = MagicMock(spec=RequestContext)
mock_request_context.request = mock_request
# Set up the context vars:
# - auth_context_var has STALE token
# - request_ctx has request with FRESH token
# - fastmcp_request_ctx has request with FRESH token
auth_token = auth_context_var.set(stale_user)
request_token = request_ctx.set(mock_request_context)
request_token = fastmcp_request_ctx.set(_make_ctx(mock_request))
try:
# Call get_access_token - should return FRESH token
@ -86,7 +98,7 @@ class TestStaleAccessToken:
finally:
# Clean up context vars
auth_context_var.reset(auth_token)
request_ctx.reset(request_token)
fastmcp_request_ctx.reset(request_token)
def test_get_access_token_falls_back_to_context_var_when_no_request(self):
"""
@ -141,11 +153,9 @@ class TestStaleAccessToken:
"user": UnauthenticatedUser(),
}
mock_request = Request(scope)
mock_request_context = MagicMock(spec=RequestContext)
mock_request_context.request = mock_request
auth_token = auth_context_var.set(user)
request_token = request_ctx.set(mock_request_context)
request_token = fastmcp_request_ctx.set(_make_ctx(mock_request))
try:
result = get_access_token()
@ -155,4 +165,4 @@ class TestStaleAccessToken:
assert result.token == "context-var-token"
finally:
auth_context_var.reset(auth_token)
request_ctx.reset(request_token)
fastmcp_request_ctx.reset(request_token)

View file

@ -2,7 +2,7 @@
import pytest
from mcp.server.streamable_http import EventMessage
from mcp_types import JSONRPCMessage, JSONRPCRequest
from mcp_types import JSONRPCRequest
from fastmcp.server.event_store import (
EventEntry,
@ -50,7 +50,7 @@ class TestEventStore:
@pytest.fixture
def sample_message(self):
return JSONRPCMessage(root=JSONRPCRequest(jsonrpc="2.0", method="test", id=1))
return JSONRPCRequest(jsonrpc="2.0", method="test", id=1)
async def test_store_event_returns_event_id(self, event_store, sample_message):
event_id = await event_store.store_event("stream-1", sample_message)
@ -89,7 +89,7 @@ class TestEventStore:
assert stream_id == "stream-1"
assert len(replayed_events) == 1
assert replayed_events[0].event_id == second_event_id
replayed_message = replayed_events[0].message.root
replayed_message = replayed_events[0].message
assert isinstance(replayed_message, JSONRPCRequest)
assert replayed_message.method == "test"
@ -99,9 +99,7 @@ class TestEventStore:
priming_id = await event_store.store_event("stream-1", None)
# Store a real event
real_message = JSONRPCMessage(
root=JSONRPCRequest(jsonrpc="2.0", method="test", id=1)
)
real_message = JSONRPCRequest(jsonrpc="2.0", method="test", id=1)
await event_store.store_event("stream-1", real_message)
# Replay after priming event
@ -130,9 +128,7 @@ class TestEventStore:
# Store more events than the limit
event_ids = []
for i in range(7):
msg = JSONRPCMessage(
root=JSONRPCRequest(jsonrpc="2.0", method=f"test-{i}", id=i)
)
msg = JSONRPCRequest(jsonrpc="2.0", method=f"test-{i}", id=i)
event_id = await event_store.store_event("stream-1", msg)
event_ids.append(event_id)
@ -152,12 +148,8 @@ class TestEventStore:
async def test_multiple_streams_are_isolated(self, event_store):
"""Events from different streams should not interfere with each other."""
msg1 = JSONRPCMessage(
root=JSONRPCRequest(jsonrpc="2.0", method="stream1-test", id=1)
)
msg2 = JSONRPCMessage(
root=JSONRPCRequest(jsonrpc="2.0", method="stream2-test", id=2)
)
msg1 = JSONRPCRequest(jsonrpc="2.0", method="stream1-test", id=1)
msg2 = JSONRPCRequest(jsonrpc="2.0", method="stream2-test", id=2)
stream1_event = await event_store.store_event("stream-1", msg1)
await event_store.store_event("stream-1", msg1)
@ -191,15 +183,9 @@ class TestEventStore:
):
session_a_store = SessionScopedEventStore(event_store, "session-a")
session_b_store = SessionScopedEventStore(event_store, "session-b")
msg_a1 = JSONRPCMessage(
root=JSONRPCRequest(jsonrpc="2.0", method="session-a-1", id=1)
)
msg_a2 = JSONRPCMessage(
root=JSONRPCRequest(jsonrpc="2.0", method="session-a-2", id=2)
)
msg_b = JSONRPCMessage(
root=JSONRPCRequest(jsonrpc="2.0", method="session-b", id=3)
)
msg_a1 = JSONRPCRequest(jsonrpc="2.0", method="session-a-1", id=1)
msg_a2 = JSONRPCRequest(jsonrpc="2.0", method="session-a-2", id=2)
msg_b = JSONRPCRequest(jsonrpc="2.0", method="session-b", id=3)
session_a_event = await session_a_store.store_event(stream_id, msg_a1)
await session_b_store.store_event(stream_id, msg_b)
@ -216,7 +202,7 @@ class TestEventStore:
assert replayed_stream_id == stream_id
assert [event.event_id for event in replayed_events] == [session_a_second_event]
replayed_message = replayed_events[0].message.root
replayed_message = replayed_events[0].message
assert isinstance(replayed_message, JSONRPCRequest)
assert replayed_message.method == "session-a-2"
@ -225,12 +211,8 @@ class TestEventStore:
):
session_a_store = SessionScopedEventStore(event_store, "session-a")
session_b_store = SessionScopedEventStore(event_store, "session-b")
msg_b1 = JSONRPCMessage(
root=JSONRPCRequest(jsonrpc="2.0", method="session-b-1", id=1)
)
msg_b2 = JSONRPCMessage(
root=JSONRPCRequest(jsonrpc="2.0", method="session-b-2", id=2)
)
msg_b1 = JSONRPCRequest(jsonrpc="2.0", method="session-b-1", id=1)
msg_b2 = JSONRPCRequest(jsonrpc="2.0", method="session-b-2", id=2)
foreign_event_id = await session_b_store.store_event("_GET_stream", msg_b1)
await session_b_store.store_event("_GET_stream", msg_b2)
@ -262,7 +244,7 @@ class TestEventStore:
async def test_default_storage_is_memory(self):
"""Test that EventStore defaults to in-memory storage."""
event_store = EventStore()
msg = JSONRPCMessage(root=JSONRPCRequest(jsonrpc="2.0", method="test", id=1))
msg = JSONRPCRequest(jsonrpc="2.0", method="test", id=1)
event_id = await event_store.store_event("stream-1", msg)
assert event_id is not None
@ -285,26 +267,22 @@ class TestEventStoreIntegration:
event_store = EventStore()
# Create a realistic JSON-RPC request wrapped in JSONRPCMessage
original_msg = JSONRPCMessage(
root=JSONRPCRequest(
jsonrpc="2.0",
method="tools/call",
id="request-123",
params={"name": "my_tool", "arguments": {"x": 1, "y": 2}},
)
original_msg = JSONRPCRequest(
jsonrpc="2.0",
method="tools/call",
id="request-123",
params={"name": "my_tool", "arguments": {"x": 1, "y": 2}},
)
# Store it
event_id = await event_store.store_event("stream-1", original_msg)
# Store another event so we have something to replay
second_msg = JSONRPCMessage(
root=JSONRPCRequest(
jsonrpc="2.0",
method="tools/call",
id="request-456",
params={"name": "my_tool", "arguments": {"x": 3, "y": 4}},
)
second_msg = JSONRPCRequest(
jsonrpc="2.0",
method="tools/call",
id="request-456",
params={"name": "my_tool", "arguments": {"x": 3, "y": 4}},
)
await event_store.store_event("stream-1", second_msg)
@ -318,6 +296,6 @@ class TestEventStoreIntegration:
assert len(replayed) == 1
assert replayed[0].event_id is not None
assert isinstance(replayed[0].message.root, JSONRPCRequest)
assert replayed[0].message.root.method == "tools/call"
assert replayed[0].message.root.id == "request-456"
assert isinstance(replayed[0].message, JSONRPCRequest)
assert replayed[0].message.method == "tools/call"
assert replayed[0].message.id == "request-456"