From 04f67da520f1f92ac1d48f88a8a1d664866bb080 Mon Sep 17 00:00:00 2001 From: Marcelo Trylesinski Date: Fri, 18 Sep 2026 11:52:31 +0200 Subject: [PATCH] Buffer replay before network delivery and preserve the live handoff --- docs/run/legacy-clients.md | 8 + src/mcp/server/streamable_http.py | 82 +++- tests/server/test_streamable_http_router.py | 424 +++++++++++++++++++- 3 files changed, 496 insertions(+), 18 deletions(-) diff --git a/docs/run/legacy-clients.md b/docs/run/legacy-clients.md index b19364bc09..591059c800 100644 --- a/docs/run/legacy-clients.md +++ b/docs/run/legacy-clients.md @@ -66,6 +66,14 @@ On one worker that is invisible. On two, it is the whole problem: a request that events to a client reconnecting to the *same* session), not a session store. It never makes a session reachable from another process. +!!! note "Replay is buffered before network delivery" + With `event_store=`, the SDK collects replayed events before sending them, so a slow replay + reader does not hold the event-store lock. The buffer spills to a temporary file above + 1 MiB instead of retaining the whole history in memory. Historical events, any new resumption + cursor, and live events are sent in that order. The lock does not serialize incoming POSTs + or reserve JSON-RPC request IDs. Live-stream backpressure can still delay other messages in + the same session. + !!! note "Request cleanup preserves newer streams" Closing an HTTP request releases its own streams, including during cancellation. If you reuse a JSON-RPC request ID after the previous request completes, cleanup from an diff --git a/src/mcp/server/streamable_http.py b/src/mcp/server/streamable_http.py index 5f86df316e..f03b07978f 100644 --- a/src/mcp/server/streamable_http.py +++ b/src/mcp/server/streamable_http.py @@ -6,6 +6,7 @@ responses, with streaming support for long-running operations. """ +import json import logging import math import re @@ -210,11 +211,13 @@ def __init__( self.mcp_session_id = mcp_session_id self.is_json_response_enabled = is_json_response_enabled self._event_store = event_store + self._event_store_lock = anyio.Lock() self._security = TransportSecurityMiddleware(security_settings) self._retry_interval = retry_interval self._request_streams: dict[RequestId, _RequestStreams] = {} self._sse_stream_writers: dict[RequestId, MemoryObjectSendStream[SSEEvent]] = {} self._terminated = False + self._connected = False self._idle_timeout = idle_timeout self._requests_in_flight = 0 self.idle_scope: anyio.CancelScope | None = None @@ -915,25 +918,43 @@ async def _replay_events(self, last_event_id: str, request: Request, send: Send) async def replay_sender(): stream_id: StreamId | None = None request_streams: _RequestStreams | None = None + priming_event: SSEEvent | None = None + replay_buffer = anyio.SpooledTemporaryFile(max_size=1024 * 1024) try: async with sse_stream_writer: async def send_event(event_message: EventMessage) -> None: - await sse_stream_writer.send(self._create_event_data(event_message)) - - stream_id = await event_store.replay_events_after(last_event_id, send_event) - if stream_id and stream_id not in self._request_streams: # pragma: no branch - request_streams = anyio.create_memory_object_stream[EventMessage]( - REQUEST_STREAM_BUFFER_SIZE + event_data = self._create_event_data(event_message) + # Whole records prevent ping frames from splitting a replayed SSE event. + await replay_buffer.write( + (json.dumps(event_data, ensure_ascii=False) + "\n").encode("utf-8") ) - self._sse_stream_writers[stream_id] = sse_stream_writer - self._request_streams[stream_id] = request_streams - priming_event = await self._mint_priming_event(stream_id, replay_protocol_version) - if priming_event is not None: - await sse_stream_writer.send(priming_event) - async with request_streams[1] as msg_reader: - async for event_message in msg_reader: - await sse_stream_writer.send(self._create_event_data(event_message)) + + async with self._event_store_lock: + if self._terminated or not self._connected: + return + stream_id = await event_store.replay_events_after(last_event_id, send_event) + if self._terminated or not self._connected: + return + if stream_id and stream_id not in self._request_streams: + request_streams = anyio.create_memory_object_stream[EventMessage]( + REQUEST_STREAM_BUFFER_SIZE + ) + self._sse_stream_writers[stream_id] = sse_stream_writer + self._request_streams[stream_id] = request_streams + priming_event = await self._mint_priming_event(stream_id, replay_protocol_version) + + await replay_buffer.seek(0) + while event_data := await replay_buffer.readline(): + await sse_stream_writer.send(json.loads(event_data)) + await replay_buffer.aclose() + if request_streams is None or self._terminated or not self._connected: + return + if priming_event is not None: + await sse_stream_writer.send(priming_event) + async with request_streams[1] as msg_reader: + async for event_message in msg_reader: + await sse_stream_writer.send(self._create_event_data(event_message)) except anyio.ClosedResourceError: # pragma: lax no cover # Expected when close_sse_stream() is called logger.debug("Replay SSE stream closed by close_sse_stream()") @@ -945,6 +966,8 @@ async def send_event(event_message: EventMessage) -> None: if request_streams is not None: assert stream_id is not None self._clean_up_memory_streams(stream_id, request_streams) + with anyio.CancelScope(shield=True): + await replay_buffer.aclose() # Create and start EventSourceResponse response = EventSourceResponse( @@ -999,10 +1022,12 @@ async def connect( self._write_stream_reader = write_stream_reader self._write_stream = write_stream + acquire_scope: anyio.CancelScope | None = None # Start a task group for message routing async with anyio.create_task_group() as tg: # Create a message router that distributes messages to request streams async def message_router(): + nonlocal acquire_scope try: async for session_message in write_stream_reader: # pragma: no branch # Determine which request stream(s) should receive this message @@ -1038,10 +1063,28 @@ async def message_router(): # messages will be replayed on the re-connect event_id = None if self._event_store: - event_id = await self._event_store.store_event(request_stream_id, message) - logger.debug(f"Stored {event_id} from {request_stream_id}") + with anyio.CancelScope() as lock_scope: + acquire_scope = lock_scope + try: + try: + self._event_store_lock.acquire_nowait() + except anyio.WouldBlock: + if not self._connected: + return + await self._event_store_lock.acquire() + finally: + acquire_scope = None + if lock_scope.cancelled_caught: + return + try: + event_id = await self._event_store.store_event(request_stream_id, message) + logger.debug(f"Stored {event_id} from {request_stream_id}") + target = self._request_streams.get(request_stream_id) + finally: + self._event_store_lock.release() + else: + target = self._request_streams.get(request_stream_id) - target = self._request_streams.get(request_stream_id) if target is not None: try: # Send both the message and the event ID @@ -1066,8 +1109,10 @@ async def message_router(): tg.start_soon(message_router) try: + self._connected = True yield read_stream, write_stream finally: + self._connected = False for stream_id, streams in list(self._request_streams.items()): self._clean_up_memory_streams(stream_id, streams) @@ -1080,3 +1125,6 @@ async def message_router(): except Exception as e: # pragma: no cover # During cleanup, we catch all exceptions since streams might be in various states logger.debug(f"Error closing streams: {e}") + # Cancel only lock acquisition, not a store write that has already started. + if acquire_scope is not None: + acquire_scope.cancel() diff --git a/tests/server/test_streamable_http_router.py b/tests/server/test_streamable_http_router.py index bd10a2b647..be433223cb 100644 --- a/tests/server/test_streamable_http_router.py +++ b/tests/server/test_streamable_http_router.py @@ -5,7 +5,7 @@ import anyio import httpx2 import pytest -from mcp_types import JSONRPCMessage, JSONRPCRequest, JSONRPCResponse, jsonrpc_message_adapter +from mcp_types import JSONRPCMessage, JSONRPCNotification, JSONRPCRequest, JSONRPCResponse, jsonrpc_message_adapter from starlette.types import Message, Scope from mcp.server.streamable_http import ( @@ -167,6 +167,428 @@ async def test_terminated_transport_answers_404() -> None: assert post.sent[0]["status"] == 404 +@pytest.mark.anyio +async def test_a_response_produced_during_replay_reaches_the_resumed_stream() -> None: + """SDK-defined: a response arriving after the replay snapshot cannot fall between replay and live delivery. + + A gated public EventStore and raw ASGI peer expose the handoff without depending on HTTP client scheduling. + """ + snapshot_taken = anyio.Event() + release_replay = anyio.Event() + response_delivered = anyio.Event() + disconnect = anyio.Event() + progress = JSONRPCNotification( + jsonrpc="2.0", method="notifications/progress", params={"progressToken": "call", "progress": 0.5, "total": 1} + ) + response = JSONRPCResponse(jsonrpc="2.0", id="request", result={"value": "resumed"}) + wire: list[bytes] = [] + + class SnapshotStore(EventStore): + def __init__(self) -> None: + self.events: list[tuple[StreamId, JSONRPCMessage | None]] = [] + + async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId: + self.events.append((stream_id, message)) + return str(len(self.events)) + + async def replay_events_after(self, last_event_id: EventId, send_callback: EventCallback) -> StreamId | None: + cursor = int(last_event_id) + stream_id, _ = self.events[cursor - 1] + snapshot = tuple(self.events[cursor:]) + snapshot_taken.set() + await release_replay.wait() + for index, (_, message) in enumerate(snapshot, cursor + 1): + assert message is not None + await send_callback(EventMessage(message, str(index))) + return stream_id + + store = SnapshotStore() + last_event_id = await store.store_event("request", None) + await store.store_event("request", progress) + transport = StreamableHTTPServerTransport(mcp_session_id=None, event_store=store) + scope: Scope = { + "type": "http", + "method": "GET", + "path": "/mcp", + "query_string": b"", + "headers": [ + (b"accept", b"text/event-stream"), + (b"mcp-protocol-version", b"2025-11-25"), + (b"last-event-id", last_event_id.encode()), + ], + } + + async def receive() -> Message: + await disconnect.wait() + return {"type": "http.disconnect"} + + async def send(message: Message) -> None: + if message["type"] == "http.response.start": + assert message["status"] == 200 + else: + body = message.get("body", b"") + wire.append(body) + if response.model_dump_json(by_alias=True, exclude_unset=True).encode() in body: + response_delivered.set() + + with anyio.fail_after(5): + async with transport.connect() as (_, write_stream), anyio.create_task_group() as tg: + tg.start_soon(transport.handle_request, scope, receive, send) + await snapshot_taken.wait() + await write_stream.send(SessionMessage(response)) + await anyio.wait_all_tasks_blocked() + release_replay.set() + await response_delivered.wait() + disconnect.set() + + received = httpx2.Response( + 200, + headers={"content-type": "text/event-stream"}, + content=b"".join(wire), + request=httpx2.Request("GET", "http://localhost/mcp"), + ) + assert [ + jsonrpc_message_adapter.validate_json(event.data) for event in httpx2.EventSource(received) if event.data + ] == [ + progress, + response, + ] + + +@pytest.mark.anyio +@pytest.mark.parametrize("terminate", [False, True], ids=["context-exit", "terminated"]) +@pytest.mark.parametrize("wait_for_router", [False, True], ids=["after-send", "queued"]) +async def test_transport_shutdown_cancels_a_router_waiting_for_replay(terminate: bool, wait_for_router: bool) -> None: + """SDK-defined: shutdown releases a blocked router, and its late replay cannot open a dead stream.""" + replay_started = anyio.Event() + release_replay = anyio.Event() + replay_finished = anyio.Event() + sent: list[Message] = [] + + class BlockingStore(EventStore): + async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId: + return "cursor" + + async def replay_events_after(self, last_event_id: EventId, send_callback: EventCallback) -> StreamId | None: + replay_started.set() + await release_replay.wait() + return "request" + + store = BlockingStore() + cursor = await store.store_event("request", None) + transport = StreamableHTTPServerTransport(mcp_session_id=None, event_store=store) + scope: Scope = { + "type": "http", + "method": "GET", + "path": "/mcp", + "query_string": b"", + "headers": [ + (b"accept", b"text/event-stream"), + (b"last-event-id", cursor.encode()), + (b"mcp-protocol-version", b"2025-11-25"), + ], + } + + async def receive() -> Message: + await anyio.sleep_forever() + raise NotImplementedError + + async def send(message: Message) -> None: + sent.append(message) + + async def replay() -> None: + await transport.handle_request(scope, receive, send) + replay_finished.set() + + with anyio.fail_after(5): + async with anyio.create_task_group() as requests: + async with transport.connect() as (_, write_stream): + requests.start_soon(replay) + await replay_started.wait() + await anyio.wait_all_tasks_blocked() + await write_stream.send(SessionMessage(JSONRPCResponse(jsonrpc="2.0", id="request", result={}))) + if wait_for_router: + await anyio.wait_all_tasks_blocked() + if terminate: + await transport.terminate() + release_replay.set() + await replay_finished.wait() + + assert transport.is_terminated is terminate + assert sent[0]["type"] == "http.response.start" + assert sent[0]["status"] == 200 + assert b"".join(message.get("body", b"") for message in sent) == b"" + + +@pytest.mark.anyio +@pytest.mark.parametrize("window", ["lock", "hook"]) +async def test_termination_during_replay_setup_ends_the_response_without_priming(window: str) -> None: + """SDK-defined: termination wins over a replay waiting for the lock or its event-store hook.""" + storing = anyio.Event() + replay_started = anyio.Event() + headers_sent = anyio.Event() + release = anyio.Event() + replay_finished = anyio.Event() + sent: list[Message] = [] + response = JSONRPCResponse(jsonrpc="2.0", id="request", result={}) + + class GatedStore(EventStore): + async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId: + if window == "lock": + storing.set() + await release.wait() + return "stored" + + async def replay_events_after(self, last_event_id: EventId, send_callback: EventCallback) -> StreamId | None: + replay_started.set() + await release.wait() + return "request" + + transport = StreamableHTTPServerTransport(None, event_store=GatedStore()) + scope: Scope = { + "type": "http", + "method": "GET", + "path": "/mcp", + "query_string": b"", + "headers": [ + (b"accept", b"text/event-stream"), + (b"last-event-id", b"cursor"), + (b"mcp-protocol-version", b"2025-11-25"), + ], + } + + async def receive() -> Message: + await anyio.sleep_forever() + raise NotImplementedError + + async def send(message: Message) -> None: + sent.append(message) + if message["type"] == "http.response.start": + assert message["status"] == 200 + headers_sent.set() + + async def replay() -> None: + await transport.handle_request(scope, receive, send) + replay_finished.set() + + with anyio.fail_after(5): + async with transport.connect() as (_, write_stream), anyio.create_task_group() as requests: + if window == "lock": + await write_stream.send(SessionMessage(response)) + await storing.wait() + requests.start_soon(replay) + await headers_sent.wait() + if window == "hook": + await replay_started.wait() + await write_stream.send(SessionMessage(response)) + await anyio.wait_all_tasks_blocked() + await transport.terminate() + release.set() + await replay_finished.wait() + await anyio.wait_all_tasks_blocked() + + assert replay_started.is_set() is (window == "hook") + assert b"".join(message.get("body", b"") for message in sent) == b"" + + +@pytest.mark.anyio +@pytest.mark.parametrize("history_size", [0, 16, 1024 * 1024 + 1], ids=["priming", "memory", "spilled"]) +@pytest.mark.parametrize("disconnect_early", [False, True], ids=["drain", "disconnect"]) +async def test_a_blocked_replay_does_not_prevent_a_sibling_response(history_size: int, disconnect_early: bool) -> None: + """SDK-defined: a replay's network backpressure cannot hold the event-store lock. + + The raw ASGI peer blocks response headers, making the network stall deterministic without HTTPX2. + """ + replay_started = anyio.Event() + headers_sent = anyio.Event() + release_headers = anyio.Event() + priming_received = anyio.Event() + tail_received = anyio.Event() + disconnect = anyio.Event() + post_finished = anyio.Event() + history = JSONRPCNotification( + jsonrpc="2.0", + method="notifications/progress", + params={"progressToken": "p", "progress": 0.5, "message": "x" * history_size}, + ) + tail = JSONRPCResponse(jsonrpc="2.0", id="replay", result={"done": True}) + sibling = JSONRPCResponse(jsonrpc="2.0", id="sibling", result={"ok": True}) + chunks: list[bytes] = [] + + class ReplayStore(EventStore): + def __init__(self) -> None: + self.events: list[JSONRPCMessage | None] = [] + + async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId: + self.events.append(message) + return str(len(self.events)) + + async def replay_events_after(self, last_event_id: EventId, send_callback: EventCallback) -> StreamId | None: + assert last_event_id == "cursor" + replay_started.set() + if history_size: + await send_callback(EventMessage(history, "historical")) + return "replay" + + transport = StreamableHTTPServerTransport(None, is_json_response_enabled=True, event_store=ReplayStore()) + scope: Scope = { + "type": "http", + "method": "GET", + "path": "/mcp", + "query_string": b"", + "headers": [ + (b"accept", b"text/event-stream"), + (b"last-event-id", b"cursor"), + (b"mcp-protocol-version", b"2025-11-25"), + ], + } + + async def receive() -> Message: + await disconnect.wait() + return {"type": "http.disconnect"} + + async def send(message: Message) -> None: + if message["type"] == "http.response.start": + assert message["status"] == 200 + headers_sent.set() + await release_headers.wait() + elif body := message.get("body", b""): + assert body.endswith(b"\r\n\r\n") + chunks.append(body) + if body.endswith(b"data: \r\n\r\n"): + priming_received.set() + if tail.model_dump_json(by_alias=True, exclude_unset=True).encode() in body: + tail_received.set() + + post = _AsgiPost( + b'{"jsonrpc":"2.0","id":"sibling","method":"tools/list"}', + [(b"accept", b"application/json"), (b"content-type", b"application/json")], + ) + + async def run_post() -> None: + await transport.handle_request(post.scope, post.receive, post.send) + post_finished.set() + + with anyio.fail_after(5): + async with transport.connect() as (read_stream, write_stream): + async with anyio.create_task_group() as requests: + requests.start_soon(transport.handle_request, scope, receive, send) + await headers_sent.wait() + await replay_started.wait() + requests.start_soon(run_post) + incoming = await read_stream.receive() + assert isinstance(incoming, SessionMessage) + assert incoming.message == JSONRPCRequest(jsonrpc="2.0", id="sibling", method="tools/list") + await write_stream.send(SessionMessage(sibling)) + await post_finished.wait() + assert ( + jsonrpc_message_adapter.validate_json(b"".join(message.get("body", b"") for message in post.sent)) + == sibling + ) + if disconnect_early: + disconnect.set() + else: + release_headers.set() + await priming_received.wait() + await write_stream.send(SessionMessage(tail)) + await tail_received.wait() + transport.close_sse_stream(tail.id) + await transport.terminate() + + gc.collect() + if disconnect_early: + assert chunks == [] + return + received = httpx2.Response( + 200, + headers={"content-type": "text/event-stream"}, + content=b"".join(chunks), + request=httpx2.Request("GET", "http://localhost/mcp"), + ) + events = list(httpx2.EventSource(received)) + assert [jsonrpc_message_adapter.validate_json(event.data) for event in events if event.data] == ( + [history, tail] if history_size else [tail] + ) + assert events[-2].data == "" + assert events[-1].id != events[-2].id + + +@pytest.mark.anyio +@pytest.mark.parametrize("wait_for_store", [False, True], ids=["after-send", "during-store"]) +async def test_normal_transport_exit_waits_for_an_active_event_store_write(wait_for_store: bool) -> None: + """SDK-defined: normal session teardown finishes an accepted event-store write instead of cancelling it.""" + storing = anyio.Event() + release = anyio.Event() + exited = anyio.Event() + response = JSONRPCResponse(jsonrpc="2.0", id="last", result={"done": True}) + stored: list[JSONRPCMessage | None] = [] + + class GatedStore(EventStore): + async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId: + storing.set() + await release.wait() + stored.append(message) + return "committed" + + async def replay_events_after(self, last_event_id: EventId, send_callback: EventCallback) -> StreamId | None: + raise NotImplementedError + + transport = StreamableHTTPServerTransport(None, event_store=GatedStore()) + + async def run_session() -> None: + async with transport.connect() as (_, write_stream): + await write_stream.send(SessionMessage(response)) + if wait_for_store: + await storing.wait() + exited.set() + + with anyio.fail_after(5): + async with anyio.create_task_group() as tasks: + tasks.start_soon(run_session) + await storing.wait() + await anyio.wait_all_tasks_blocked() + try: + assert not exited.is_set() + finally: + release.set() + await exited.wait() + assert stored == [response] + + +@pytest.mark.anyio +async def test_caller_cancellation_still_interrupts_an_event_store_write() -> None: + """SDK-defined: draining a store on normal exit does not shield it against caller cancellation.""" + storing = anyio.Event() + stopped = anyio.Event() + + class GatedStore(EventStore): + async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId: + try: + with anyio.fail_after(5): + storing.set() + await anyio.sleep_forever() + finally: + stopped.set() + raise NotImplementedError + + async def replay_events_after(self, last_event_id: EventId, send_callback: EventCallback) -> StreamId | None: + raise NotImplementedError + + transport = StreamableHTTPServerTransport(None, event_store=GatedStore()) + + async def run_session() -> None: + async with transport.connect() as (_, write_stream): + await write_stream.send(SessionMessage(JSONRPCResponse(jsonrpc="2.0", id="last", result={}))) + await anyio.sleep_forever() + + with anyio.fail_after(5): + async with anyio.create_task_group() as tasks: + tasks.start_soon(run_session) + await storing.wait() + tasks.cancel_scope.cancel() + assert stopped.is_set() + + @pytest.mark.anyio @pytest.mark.parametrize("json_response", [False, True], ids=["sse", "json"]) @pytest.mark.parametrize("pause_priming", [False, True], ids=["live", "priming"])