From 1a9b1c435f794e5b9204f867f6b23fe8f5df554d Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Sun, 5 Jul 2026 19:15:03 -0400 Subject: [PATCH] Port HTTP session manager, event store, and transport run paths to SDK v2 --- fastmcp_slim/fastmcp/server/event_store.py | 7 +- fastmcp_slim/fastmcp/server/http.py | 9 +++ tests/server/http/test_stale_access_token.py | 36 ++++++---- tests/server/test_event_store.py | 76 +++++++------------- 4 files changed, 65 insertions(+), 63 deletions(-) diff --git a/fastmcp_slim/fastmcp/server/event_store.py b/fastmcp_slim/fastmcp/server/event_store.py index 7b988d61b..86897aac7 100644 --- a/fastmcp_slim/fastmcp/server/event_store.py +++ b/fastmcp_slim/fastmcp/server/event_store.py @@ -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 diff --git a/fastmcp_slim/fastmcp/server/http.py b/fastmcp_slim/fastmcp/server/http.py index c659fd85a..b96bcdd5e 100644 --- a/fastmcp_slim/fastmcp/server/http.py +++ b/fastmcp_slim/fastmcp/server/http.py @@ -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(), diff --git a/tests/server/http/test_stale_access_token.py b/tests/server/http/test_stale_access_token.py index 507b4deb6..c2cf19180 100644 --- a/tests/server/http/test_stale_access_token.py +++ b/tests/server/http/test_stale_access_token.py @@ -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) diff --git a/tests/server/test_event_store.py b/tests/server/test_event_store.py index b70120ca5..edb00b5e8 100644 --- a/tests/server/test_event_store.py +++ b/tests/server/test_event_store.py @@ -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"