mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-22 05:24:18 +02:00
[codex] Fix event store replay isolation (#4402)
This commit is contained in:
parent
5de15e0c21
commit
feaae683a4
3 changed files with 173 additions and 3 deletions
|
|
@ -37,6 +37,58 @@ class StreamEventList(FastMCPBaseModel):
|
|||
event_ids: list[str]
|
||||
|
||||
|
||||
class SessionScopedEventStore(SDKEventStore):
|
||||
"""EventStore adapter that isolates stream IDs to one transport session."""
|
||||
|
||||
def __init__(self, event_store: SDKEventStore, session_id: str):
|
||||
self._event_store = event_store
|
||||
self._stream_prefix = f"{len(session_id)}:{session_id}:"
|
||||
|
||||
def _scope_stream_id(self, stream_id: StreamId) -> StreamId:
|
||||
return f"{self._stream_prefix}{stream_id}"
|
||||
|
||||
def _unscope_stream_id(self, stream_id: StreamId) -> StreamId | None:
|
||||
if not stream_id.startswith(self._stream_prefix):
|
||||
return None
|
||||
return stream_id[len(self._stream_prefix) :]
|
||||
|
||||
async def store_event(
|
||||
self, stream_id: StreamId, message: JSONRPCMessage | None
|
||||
) -> EventId:
|
||||
return await self._event_store.store_event(
|
||||
self._scope_stream_id(stream_id), message
|
||||
)
|
||||
|
||||
async def replay_events_after(
|
||||
self,
|
||||
last_event_id: EventId,
|
||||
send_callback: EventCallback,
|
||||
) -> StreamId | None:
|
||||
replayed_events: list[EventMessage] = []
|
||||
|
||||
async def buffer_event(event: EventMessage) -> None:
|
||||
replayed_events.append(event)
|
||||
|
||||
scoped_stream_id = await self._event_store.replay_events_after(
|
||||
last_event_id, buffer_event
|
||||
)
|
||||
if scoped_stream_id is None:
|
||||
return None
|
||||
|
||||
stream_id = self._unscope_stream_id(scoped_stream_id)
|
||||
if stream_id is None:
|
||||
logger.warning(
|
||||
"Event ID %s does not belong to this session-scoped event store",
|
||||
last_event_id,
|
||||
)
|
||||
return None
|
||||
|
||||
for event in replayed_events:
|
||||
await send_callback(event)
|
||||
|
||||
return stream_id
|
||||
|
||||
|
||||
class EventStore(SDKEventStore):
|
||||
"""EventStore implementation backed by AsyncKeyValue.
|
||||
|
||||
|
|
|
|||
|
|
@ -3,7 +3,8 @@ from __future__ import annotations
|
|||
from collections.abc import AsyncGenerator, Callable, Generator
|
||||
from contextlib import asynccontextmanager, contextmanager
|
||||
from contextvars import ContextVar
|
||||
from typing import TYPE_CHECKING
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from uuid import uuid4
|
||||
|
||||
from mcp.server.auth.routes import build_resource_metadata_url
|
||||
from mcp.server.lowlevel.server import LifespanResultT
|
||||
|
|
@ -12,6 +13,7 @@ from mcp.server.streamable_http import (
|
|||
EventStore,
|
||||
)
|
||||
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
|
||||
from mcp.server.transport_security import TransportSecuritySettings
|
||||
from starlette.applications import Starlette
|
||||
from starlette.middleware import Middleware
|
||||
from starlette.requests import Request
|
||||
|
|
@ -21,6 +23,7 @@ from starlette.types import Lifespan, Receive, Scope, Send
|
|||
|
||||
from fastmcp.server.auth import AuthProvider
|
||||
from fastmcp.server.auth.middleware import RequireAuthMiddleware
|
||||
from fastmcp.server.event_store import SessionScopedEventStore
|
||||
from fastmcp.utilities.logging import get_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -29,6 +32,41 @@ if TYPE_CHECKING:
|
|||
logger = get_logger(__name__)
|
||||
|
||||
|
||||
class FastMCPStreamableHTTPSessionManager(StreamableHTTPSessionManager):
|
||||
"""Session manager that scopes resumability storage per transport session."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
app: Any,
|
||||
event_store: EventStore | None = None,
|
||||
json_response: bool = False,
|
||||
stateless: bool = False,
|
||||
security_settings: TransportSecuritySettings | None = None,
|
||||
retry_interval: int | None = None,
|
||||
) -> None:
|
||||
self._shared_event_store: EventStore | None = None
|
||||
super().__init__(
|
||||
app=app,
|
||||
event_store=event_store,
|
||||
json_response=json_response,
|
||||
stateless=stateless,
|
||||
security_settings=security_settings,
|
||||
retry_interval=retry_interval,
|
||||
)
|
||||
|
||||
@property
|
||||
def event_store(self) -> EventStore | None:
|
||||
if self._shared_event_store is None:
|
||||
return None
|
||||
# The SDK reads `self.event_store` once when constructing each transport.
|
||||
# A fresh adapter gives that transport a private stream namespace.
|
||||
return SessionScopedEventStore(self._shared_event_store, session_id=uuid4().hex)
|
||||
|
||||
@event_store.setter
|
||||
def event_store(self, event_store: EventStore | None) -> None:
|
||||
self._shared_event_store = event_store
|
||||
|
||||
|
||||
class StreamableHTTPASGIApp:
|
||||
"""ASGI application wrapper for Streamable HTTP server transport."""
|
||||
|
||||
|
|
@ -359,7 +397,7 @@ def create_streamable_http_app(
|
|||
# Create a lifespan manager to start and stop the session manager
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: Starlette) -> AsyncGenerator[None, None]:
|
||||
streamable_http_app.session_manager = StreamableHTTPSessionManager(
|
||||
streamable_http_app.session_manager = FastMCPStreamableHTTPSessionManager(
|
||||
app=server._mcp_server,
|
||||
event_store=event_store,
|
||||
retry_interval=retry_interval,
|
||||
|
|
|
|||
|
|
@ -4,7 +4,13 @@ import pytest
|
|||
from mcp.server.streamable_http import EventMessage
|
||||
from mcp.types import JSONRPCMessage, JSONRPCRequest
|
||||
|
||||
from fastmcp.server.event_store import EventEntry, EventStore, StreamEventList
|
||||
from fastmcp.server.event_store import (
|
||||
EventEntry,
|
||||
EventStore,
|
||||
SessionScopedEventStore,
|
||||
StreamEventList,
|
||||
)
|
||||
from fastmcp.server.http import FastMCPStreamableHTTPSessionManager
|
||||
|
||||
|
||||
class TestEventEntry:
|
||||
|
|
@ -179,6 +185,80 @@ class TestEventStore:
|
|||
assert stream_id == "stream-2"
|
||||
assert len(stream2_replayed) == 1
|
||||
|
||||
@pytest.mark.parametrize("stream_id", ["_GET_stream", "1"])
|
||||
async def test_session_scoped_stores_isolate_overlapping_stream_ids(
|
||||
self, event_store, stream_id
|
||||
):
|
||||
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)
|
||||
)
|
||||
|
||||
session_a_event = await session_a_store.store_event(stream_id, msg_a1)
|
||||
await session_b_store.store_event(stream_id, msg_b)
|
||||
session_a_second_event = await session_a_store.store_event(stream_id, msg_a2)
|
||||
|
||||
replayed_events: list[EventMessage] = []
|
||||
|
||||
async def callback(event: EventMessage):
|
||||
replayed_events.append(event)
|
||||
|
||||
replayed_stream_id = await session_a_store.replay_events_after(
|
||||
session_a_event, callback
|
||||
)
|
||||
|
||||
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
|
||||
assert isinstance(replayed_message, JSONRPCRequest)
|
||||
assert replayed_message.method == "session-a-2"
|
||||
|
||||
async def test_session_scoped_replay_rejects_foreign_last_event_id(
|
||||
self, event_store
|
||||
):
|
||||
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)
|
||||
)
|
||||
|
||||
foreign_event_id = await session_b_store.store_event("_GET_stream", msg_b1)
|
||||
await session_b_store.store_event("_GET_stream", msg_b2)
|
||||
|
||||
replayed_events: list[EventMessage] = []
|
||||
|
||||
async def callback(event: EventMessage):
|
||||
replayed_events.append(event)
|
||||
|
||||
replayed_stream_id = await session_a_store.replay_events_after(
|
||||
foreign_event_id, callback
|
||||
)
|
||||
|
||||
assert replayed_stream_id is None
|
||||
assert replayed_events == []
|
||||
|
||||
def test_session_manager_returns_scoped_event_stores(self, event_store):
|
||||
session_manager = FastMCPStreamableHTTPSessionManager(
|
||||
app=object(), event_store=event_store
|
||||
)
|
||||
|
||||
first_transport_store = session_manager.event_store
|
||||
second_transport_store = session_manager.event_store
|
||||
|
||||
assert isinstance(first_transport_store, SessionScopedEventStore)
|
||||
assert isinstance(second_transport_store, SessionScopedEventStore)
|
||||
assert first_transport_store is not second_transport_store
|
||||
|
||||
async def test_default_storage_is_memory(self):
|
||||
"""Test that EventStore defaults to in-memory storage."""
|
||||
event_store = EventStore()
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue