From a1b2b89911757ebb252c0b68c535f7905009213a Mon Sep 17 00:00:00 2001 From: Marcelo Trylesinski Date: Wed, 16 Sep 2026 12:26:36 +0200 Subject: [PATCH 1/7] Restore Trio support and test both AnyIO backends --- docs/client/subscriptions.md | 2 + docs/get-started/testing.md | 14 +- docs/run/legacy-clients.md | 6 + src/mcp/client/sse.py | 86 +++++------ src/mcp/client/streamable_http.py | 60 ++++---- src/mcp/client/subscriptions.py | 65 +++++---- src/mcp/server/sse.py | 10 +- src/mcp/server/streamable_http.py | 77 +++++----- src/mcp/shared/_httpx_utils.py | 13 +- tests/client/test_stdio.py | 9 +- tests/client/test_subscriptions.py | 98 +++++++++++++ tests/conftest.py | 14 +- tests/docs_src/test_asgi.py | 4 + tests/docs_src/test_context.py | 2 + tests/docs_src/test_deploy.py | 3 + tests/docs_src/test_identity_assertion.py | 2 + tests/docs_src/test_legacy_clients.py | 2 + tests/docs_src/test_subscriptions.py | 87 ++++++----- tests/docs_src/test_troubleshooting.py | 4 + .../transports/test_hosting_resume.py | 9 +- ...est_1363_race_condition_streamable_http.py | 107 +++----------- tests/server/test_sse_security.py | 21 +-- tests/server/test_streamable_http_router.py | 138 +++++++++++++++++- tests/shared/test_sse.py | 8 +- 24 files changed, 546 insertions(+), 295 deletions(-) diff --git a/docs/client/subscriptions.md b/docs/client/subscriptions.md index bf5b0a36d8..dddb52abf2 100644 --- a/docs/client/subscriptions.md +++ b/docs/client/subscriptions.md @@ -59,6 +59,8 @@ Requests run freely beside an open stream, from the watcher task or any other, o To stop watching, leave the block: there is no `unsubscribe` call. Cancelling the task that owns the block does that for you, and the SDK cancels the listen request the way the transport expects: over streamable HTTP, by closing that request's stream. A watcher that runs for the life of your app never returns on its own, so cancel it, or its task group's scope, at shutdown. +Exit waits up to five seconds for the SDK's listen task to finish, even if the caller is cancelled. With an in-memory server, this lets cooperative handler cleanup finish before you open another subscription. The limit prevents an uncooperative handler from blocking subscription exit indefinitely. It is not an acknowledgment that a remote server has finished its cleanup. + ## Streams end A stream ends in one of two ways, both ordinary control flow. A graceful server close ends the `async for`; an abrupt drop raises `SubscriptionLost`. diff --git a/docs/get-started/testing.md b/docs/get-started/testing.md index 98c671738f..b9211f04d1 100644 --- a/docs/get-started/testing.md +++ b/docs/get-started/testing.md @@ -12,18 +12,18 @@ Let's assume you have a simple server with a single tool: --8<-- "docs_src/testing/tutorial001.py" ``` -To run the test below you'll need two extra (development) dependencies: +Install the development dependencies to run the test on both async backends: === "uv" ```bash - uv add --dev pytest inline-snapshot + uv add --dev pytest inline-snapshot trio ``` === "pip" ```bash - pip install pytest inline-snapshot + pip install pytest inline-snapshot trio ``` !!! info @@ -45,9 +45,9 @@ from mcp.types import CallToolResult, TextContent from server import mcp -@pytest.fixture -def anyio_backend(): # (1)! - return "asyncio" +@pytest.fixture(params=["asyncio", "trio"]) +def anyio_backend(request: pytest.FixtureRequest) -> str: # (1)! + return request.param @pytest.fixture @@ -69,7 +69,7 @@ async def test_call_add_tool(client: Client): ) ``` -1. If you are using `trio`, return `"trio"` instead. See the [anyio documentation](https://anyio.readthedocs.io/en/stable/testing.html#specifying-the-backends-to-run-on) for the details. +1. Each test runs once with `asyncio` and once with `trio`. Testing both catches backend-specific assumptions and scheduling races. If your application requires one backend, keep only that name in `params`. See the [anyio documentation](https://anyio.readthedocs.io/en/stable/testing.html#specifying-the-backends-to-run-on) for details. 2. The fixture yields a connected client. Every test that takes `client` gets a fresh in-memory connection to the same server. There you go! You can now extend your tests to cover more scenarios. diff --git a/docs/run/legacy-clients.md b/docs/run/legacy-clients.md index 15809fd608..70aec64b3e 100644 --- a/docs/run/legacy-clients.md +++ b/docs/run/legacy-clients.md @@ -66,6 +66,12 @@ 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 "Slow replay delays messages in the same session" + With `event_store=`, the SDK pauses new message delivery within the session until replay + finishes and the live stream is registered. This prevents responses from being lost between + replay and live delivery. A slow replay can delay other requests in that session, but not + other sessions. + ## Session lifetime and limits A legacy session does not live forever, and one process does not hold an unlimited number of diff --git a/src/mcp/client/sse.py b/src/mcp/client/sse.py index f72ff273a9..77c2ea58af 100644 --- a/src/mcp/client/sse.py +++ b/src/mcp/client/sse.py @@ -16,6 +16,7 @@ McpHttpClientFactory, create_mcp_http_client, request_within_origin, + sse_events, sse_within_origin, ) from mcp.shared.message import SessionMessage @@ -76,48 +77,49 @@ async def sse_client( async def sse_reader(task_status: TaskStatus[str] = anyio.TASK_STATUS_IGNORED): try: - async for sse in event_source: # pragma: no branch - logger.debug(f"Received SSE event: {sse.event}") - match sse.event: - case "endpoint": - endpoint_url = urljoin(url, sse.data) - logger.debug(f"Received endpoint URL: {endpoint_url}") - - url_parsed = urlparse(url) - endpoint_parsed = urlparse(endpoint_url) - if ( # pragma: no cover - url_parsed.netloc != endpoint_parsed.netloc - or url_parsed.scheme != endpoint_parsed.scheme - ): - error_msg = ( # pragma: no cover - f"Endpoint origin does not match connection origin: {endpoint_url}" - ) - logger.error(error_msg) # pragma: no cover - raise ValueError(error_msg) # pragma: no cover - - if on_session_created: - session_id = _extract_session_id_from_endpoint(endpoint_url) - if session_id: - on_session_created(session_id) - - task_status.started(endpoint_url) - - case "message": - # Skip empty data (keep-alive pings) - if not sse.data: - continue - try: - message = types.jsonrpc_message_adapter.validate_json(sse.data, by_name=False) - logger.debug(f"Received server message: {message}") - except Exception as exc: # pragma: no cover - logger.exception("Error parsing server message") # pragma: no cover - await read_stream_writer.send(exc) # pragma: no cover - continue # pragma: no cover - - session_message = SessionMessage(message) - await read_stream_writer.send(session_message) - case _: # pragma: no cover - logger.warning(f"Unknown SSE event: {sse.event}") # pragma: no cover + async with sse_events(event_source) as events: + async for sse in events: # pragma: no branch + logger.debug(f"Received SSE event: {sse.event}") + match sse.event: + case "endpoint": + endpoint_url = urljoin(url, sse.data) + logger.debug(f"Received endpoint URL: {endpoint_url}") + + url_parsed = urlparse(url) + endpoint_parsed = urlparse(endpoint_url) + if ( # pragma: no cover + url_parsed.netloc != endpoint_parsed.netloc + or url_parsed.scheme != endpoint_parsed.scheme + ): + error_msg = ( # pragma: no cover + f"Endpoint origin does not match connection origin: {endpoint_url}" + ) + logger.error(error_msg) # pragma: no cover + raise ValueError(error_msg) # pragma: no cover + + if on_session_created: + session_id = _extract_session_id_from_endpoint(endpoint_url) + if session_id: + on_session_created(session_id) + + task_status.started(endpoint_url) + + case "message": + # Skip empty data (keep-alive pings) + if not sse.data: + continue + try: + message = types.jsonrpc_message_adapter.validate_json(sse.data, by_name=False) + logger.debug(f"Received server message: {message}") + except Exception as exc: # pragma: no cover + logger.exception("Error parsing server message") # pragma: no cover + await read_stream_writer.send(exc) # pragma: no cover + continue # pragma: no cover + + session_message = SessionMessage(message) + await read_stream_writer.send(session_message) + case _: # pragma: no cover + logger.warning(f"Unknown SSE event: {sse.event}") # pragma: no cover except SSEError as sse_exc: # pragma: lax no cover logger.exception("Encountered SSE exception") raise sse_exc diff --git a/src/mcp/client/streamable_http.py b/src/mcp/client/streamable_http.py index 82de50fd05..e6c08d9a0f 100644 --- a/src/mcp/client/streamable_http.py +++ b/src/mcp/client/streamable_http.py @@ -37,6 +37,7 @@ create_mcp_http_client, redirect_location, request_within_origin, + sse_events, sse_within_origin, stream_within_origin, ) @@ -231,7 +232,10 @@ async def handle_get_stream(self, client: httpx2.AsyncClient, read_stream_writer if last_event_id: headers[LAST_EVENT_ID] = last_event_id - async with sse_within_origin(client, self.url, headers=headers) as event_source: + async with ( + sse_within_origin(client, self.url, headers=headers) as event_source, + sse_events(event_source) as events, + ): if (redirect := _unfollowed_redirect(event_source.response)) is not None: # The same GET would be redirected again, so retrying cannot help. logger.warning(f"GET stream not opened: {redirect}") @@ -239,7 +243,7 @@ async def handle_get_stream(self, client: httpx2.AsyncClient, read_stream_writer event_source.response.raise_for_status() logger.debug("GET SSE connection established") - async for sse in event_source: + async for sse in events: # Track last event ID for reconnection if sse.id: last_event_id = sse.id @@ -278,7 +282,10 @@ async def _handle_resumption_request(self, ctx: RequestContext) -> None: if isinstance(ctx.session_message.message, JSONRPCRequest): # pragma: no branch original_request_id = ctx.session_message.message.id - async with sse_within_origin(ctx.client, self.url, headers=headers) as event_source: + async with ( + sse_within_origin(ctx.client, self.url, headers=headers) as event_source, + sse_events(event_source) as events, + ): if (redirect := _unfollowed_redirect(event_source.response)) is not None: logger.warning(redirect) assert original_request_id is not None @@ -289,7 +296,7 @@ async def _handle_resumption_request(self, ctx: RequestContext) -> None: event_source.response.raise_for_status() logger.debug("Resumption GET SSE connection established") - async for sse in event_source: # pragma: no branch + async for sse in events: # pragma: no branch is_complete = await self._handle_sse_event( sse, ctx.read_stream_writer, @@ -464,27 +471,27 @@ async def _handle_sse_response( original_request_id = ctx.session_message.message.id try: - event_source = EventSource(response) - async for sse in event_source: # pragma: no branch - # Track last event ID for potential reconnection - if sse.id: - last_event_id = sse.id + async with sse_events(EventSource(response)) as events: + async for sse in events: # pragma: no branch + # Track last event ID for potential reconnection + if sse.id: + last_event_id = sse.id - # Track retry interval from server - if sse.retry is not None: - retry_interval_ms = sse.retry + # Track retry interval from server + if sse.retry is not None: + retry_interval_ms = sse.retry - is_complete = await self._handle_sse_event( - sse, - ctx.read_stream_writer, - original_request_id=original_request_id, - resumption_callback=(ctx.metadata.on_resumption_token_update if ctx.metadata else None), - ) - # If the SSE event indicates completion, like returning response/error - # break the loop - if is_complete: - await response.aclose() - return # Normal completion, no reconnect needed + is_complete = await self._handle_sse_event( + sse, + ctx.read_stream_writer, + original_request_id=original_request_id, + resumption_callback=(ctx.metadata.on_resumption_token_update if ctx.metadata else None), + ) + # If the SSE event indicates completion, like returning response/error + # break the loop + if is_complete: + await response.aclose() + return # Normal completion, no reconnect needed except Exception: logger.debug("SSE stream ended", exc_info=True) # pragma: lax no cover @@ -542,7 +549,10 @@ async def _handle_reconnection( headers[LAST_EVENT_ID] = last_event_id try: - async with sse_within_origin(ctx.client, self.url, headers=headers) as event_source: + async with ( + sse_within_origin(ctx.client, self.url, headers=headers) as event_source, + sse_events(event_source) as events, + ): event_source.response.raise_for_status() logger.info("Reconnected to SSE stream") @@ -550,7 +560,7 @@ async def _handle_reconnection( reconnect_last_event_id: str = last_event_id reconnect_retry_ms = retry_interval_ms - async for sse in event_source: + async for sse in events: if sse.id: # pragma: no branch reconnect_last_event_id = sse.id if sse.retry is not None: diff --git a/src/mcp/client/subscriptions.py b/src/mcp/client/subscriptions.py index 27283909be..f8a1af262c 100644 --- a/src/mcp/client/subscriptions.py +++ b/src/mcp/client/subscriptions.py @@ -242,41 +242,50 @@ async def listen( opts: CallOptions = {"request_id": request_id} session._stamp(data, opts) # pyright: ignore[reportPrivateUsage] driver_scope = anyio.CancelScope() + driver_done = anyio.Event() async def drive() -> None: # Deliberately no result timeout: the response arrives when the stream ends. - with driver_scope: - try: - await session._dispatcher.send_raw_request( # pyright: ignore[reportPrivateUsage] - data["method"], data.get("params"), opts - ) - except MCPError as error: - route.settle("lost", error=error) - return - except ValueError as error: - # A raw request id collided with our minted listen id: fail this subscription - # and release the route in this same slice, so it cannot consume the raw caller's ack. - session._unregister_listen_route(request_id) # pyright: ignore[reportPrivateUsage] - route.settle("lost", error=MCPError(types.INTERNAL_ERROR, str(error))) - return - # A result, whatever its body, is the spec's graceful close; with no prior ack - # it opens the subscription already closed. - route.set_acked(types.SubscriptionFilter()) - route.settle("graceful") + try: + with driver_scope: + try: + await session._dispatcher.send_raw_request( # pyright: ignore[reportPrivateUsage] + data["method"], data.get("params"), opts + ) + except MCPError as error: + route.settle("lost", error=error) + return + except ValueError as error: + # A raw request id collided with our minted listen id: fail this subscription + # and release the route in this same slice, so it cannot consume the raw caller's ack. + session._unregister_listen_route(request_id) # pyright: ignore[reportPrivateUsage] + route.settle("lost", error=MCPError(types.INTERNAL_ERROR, str(error))) + return + # A result, whatever its body, is the spec's graceful close; with no prior ack + # it opens the subscription already closed. + route.set_acked(types.SubscriptionFilter()) + route.settle("graceful") + finally: + driver_done.set() # Register the demux route before the request is written so the ack cannot race it. route = session._register_listen_route(request_id) # pyright: ignore[reportPrivateUsage] try: task_group.start_soon(drive) - with anyio.fail_after(session._session_read_timeout_seconds): # pyright: ignore[reportPrivateUsage] - await route.acked.wait() - if route.honored is None: - # Only reachable on failure paths: a graceful no-ack result acked an empty filter in drive(). - if route.error is not None: - raise route.error - raise SubscriptionLost(f"subscription {request_id!r} ended before it was acknowledged") - yield Subscription(route, request_id, route.honored, on_event) + try: + with anyio.fail_after(session._session_read_timeout_seconds): # pyright: ignore[reportPrivateUsage] + await route.acked.wait() + if route.honored is None: + # Only reachable on failure paths: a graceful no-ack result acked an empty filter in drive(). + if route.error is not None: + raise route.error + raise SubscriptionLost(f"subscription {request_id!r} ended before it was acknowledged") + yield Subscription(route, request_id, route.honored, on_event) + finally: + route.settle("local") + driver_scope.cancel() + # Direct handlers unwind in the driver; remote cancellation has no acknowledgment. + with anyio.move_on_after(5, shield=True): + await driver_done.wait() finally: - route.settle("local") - driver_scope.cancel() session._unregister_listen_route(request_id) # pyright: ignore[reportPrivateUsage] diff --git a/src/mcp/server/sse.py b/src/mcp/server/sse.py index d71ef25004..77e45ec47d 100644 --- a/src/mcp/server/sse.py +++ b/src/mcp/server/sse.py @@ -192,7 +192,15 @@ async def sse_writer(): ) try: - async with anyio.create_task_group() as tg: + async with ( + read_stream_writer, + read_stream, + write_stream, + write_stream_reader, + sse_stream_writer, + sse_stream_reader, + anyio.create_task_group() as tg, + ): async def response_wrapper(scope: Scope, receive: Receive, send: Send): """The EventSourceResponse returning signals a client close / disconnect. diff --git a/src/mcp/server/streamable_http.py b/src/mcp/server/streamable_http.py index 416dd9e2b4..c320ee3d04 100644 --- a/src/mcp/server/streamable_http.py +++ b/src/mcp/server/streamable_http.py @@ -209,6 +209,7 @@ 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[ @@ -375,8 +376,9 @@ async def _run_sse_writer( logger.exception("Error in SSE writer") finally: logger.debug("Closing SSE writer") - self._sse_stream_writers.pop(request_id, None) - await self._clean_up_memory_streams(request_id) + if self._sse_stream_writers.get(request_id) is sse_stream_writer: + self._sse_stream_writers.pop(request_id, None) + await self._clean_up_memory_streams(request_id) def _create_error_response( self, @@ -953,6 +955,7 @@ async def _replay_events(self, last_event_id: str, request: Request, send: Send) sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[SSEEvent](0) async def replay_sender(): + stream_id: StreamId | None = None try: async with sse_stream_writer: # Define an async callback for sending events @@ -960,42 +963,32 @@ async def send_event(event_message: EventMessage) -> None: event_data = self._create_event_data(event_message) await sse_stream_writer.send(event_data) - # Replay past events and get the stream ID - stream_id = await event_store.replay_events_after(last_event_id, send_event) - - # If stream ID not in mapping, create it - if stream_id and stream_id not in self._request_streams: # pragma: no branch - try: - # Register SSE writer so close_sse_stream() can close it - self._sse_stream_writers[stream_id] = sse_stream_writer - - # Prime the resumed connection so the client sees the stream - # is re-registered. The replay→live-tail ordering window here - # is pre-existing and tracked separately. - 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) - - # Create new request streams for this connection - self._request_streams[stream_id] = anyio.create_memory_object_stream[EventMessage]( - REQUEST_STREAM_BUFFER_SIZE - ) - msg_reader = self._request_streams[stream_id][1] - - # Forward messages to SSE - async with msg_reader: - async for event_message in msg_reader: - event_data = self._create_event_data(event_message) - - await sse_stream_writer.send(event_data) - finally: - self._sse_stream_writers.pop(stream_id, None) - await self._clean_up_memory_streams(stream_id) + # ponytail: replay stalls this session's router; per-stream cursors would narrow the lock. + async with self._event_store_lock: + stream_id = await event_store.replay_events_after(last_event_id, send_event) + if not stream_id or stream_id in self._request_streams: + return + self._sse_stream_writers[stream_id] = sse_stream_writer + self._request_streams[stream_id] = anyio.create_memory_object_stream[EventMessage]( + REQUEST_STREAM_BUFFER_SIZE + ) + msg_reader = self._request_streams[stream_id][1] + 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 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()") except Exception: # pragma: lax no cover logger.exception("Error in replay sender") + finally: + if stream_id is not None and self._sse_stream_writers.get(stream_id) is sse_stream_writer: + self._sse_stream_writers.pop(stream_id, None) + await self._clean_up_memory_streams(stream_id) # Create and start EventSourceResponse response = EventSourceResponse( @@ -1089,16 +1082,20 @@ 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}") + async with self._event_store_lock: + 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) + else: + target = self._request_streams.get(request_stream_id) - if request_stream_id in self._request_streams: + if target is not None: try: # Send both the message and the event ID - await self._request_streams[request_stream_id][0].send(EventMessage(message, event_id)) + await target[0].send(EventMessage(message, event_id)) except (anyio.BrokenResourceError, anyio.ClosedResourceError): # pragma: no cover - # Stream might be closed, remove from registry - self._request_streams.pop(request_stream_id, None) + if self._request_streams.get(request_stream_id) is target: + self._request_streams.pop(request_stream_id, None) else: logger.debug( f"""Request stream {request_stream_id} not found @@ -1133,3 +1130,5 @@ 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}") + if self._event_store is not None: + tg.cancel_scope.cancel() diff --git a/src/mcp/shared/_httpx_utils.py b/src/mcp/shared/_httpx_utils.py index 940b9f08cc..6e5486a1ea 100644 --- a/src/mcp/shared/_httpx_utils.py +++ b/src/mcp/shared/_httpx_utils.py @@ -1,7 +1,7 @@ """Utilities for creating and using httpx2 AsyncClient instances in the MCP transports.""" from abc import ABC, abstractmethod -from collections.abc import AsyncGenerator +from collections.abc import AsyncGenerator, AsyncIterator from contextlib import asynccontextmanager from typing import Any, Protocol @@ -165,6 +165,17 @@ async def sse_within_origin( yield httpx2.EventSource(response) +@asynccontextmanager +async def sse_events(source: httpx2.EventSource) -> AsyncGenerator[AsyncIterator[httpx2.ServerSentEvent]]: + """Close the SSE iterator when its consumer stops before the response ends.""" + events = source.__aiter__() + try: + yield events + finally: + assert isinstance(events, AsyncGenerator) + await events.aclose() + + def redirect_location(response: httpx2.Response) -> httpx2.URL | None: """Where `response` redirects to, for use in a message: without userinfo, query or fragment, which can carry state that does not belong in an error or a log line. None if not a redirect.""" diff --git a/tests/client/test_stdio.py b/tests/client/test_stdio.py index 91f829ff98..bba6bcd01a 100644 --- a/tests/client/test_stdio.py +++ b/tests/client/test_stdio.py @@ -858,9 +858,16 @@ async def test_invalid_utf8_flushed_by_a_dying_server_does_not_break_shutdown( abort the drain or surface a UnicodeDecodeError out of the context manager. """ ping = JSONRPCRequest(jsonrpc="2.0", id=1, method="ping") - process = FakeProcess(on_stdin_close=lambda: process.exit(0)) + process = FakeProcess() terminated = install_fake_process(monkeypatch, process) + def exit_when_flushed() -> None: + if process.stdin_closed.is_set() and process.pending_stdout_chunks() == 0: + process.exit(0) + + process.on_stdin_close = exit_when_flushed + process.on_stdout_receive = exit_when_flushed + with anyio.fail_after(5): async with stdio_client(FAKE_PARAMS): # Park the reader delivering a message nobody receives, then queue diff --git a/tests/client/test_subscriptions.py b/tests/client/test_subscriptions.py index c9877c4ab0..d862561850 100644 --- a/tests/client/test_subscriptions.py +++ b/tests/client/test_subscriptions.py @@ -10,6 +10,7 @@ import mcp_types as types import pytest from mcp_types import SubscriptionFilter +from trio.testing import MockClock import mcp.client.subscriptions as subscriptions_module from mcp import Client, MCPError @@ -38,6 +39,11 @@ pytestmark = pytest.mark.anyio +@pytest.fixture(autouse=True) +def _module_runner_lease() -> None: + """Opt out of the shared runner because the cleanup timeout test parametrizes `anyio_backend`.""" + + def _bus_server(bus: InMemorySubscriptionBus, *, max_subscriptions: int | None = None) -> Server[Any]: """A lowlevel server whose only feature is serving listen streams from `bus`.""" handler = ( @@ -226,6 +232,96 @@ async def test_exiting_the_context_frees_the_server_slot(): assert second.subscription_id != first.subscription_id +async def test_a_cancelled_task_can_close_a_subscription_opened_by_another_task() -> None: + """SDK-defined: cross-task exit joins direct handler cleanup even when the closing task is cancelled.""" + handler = ListenHandler(InMemorySubscriptionBus()) + cleanup_started = anyio.Event() + release_cleanup = anyio.Event() + cleanup_finished = anyio.Event() + closed = anyio.Event() + + async def slow_cleanup( + ctx: ServerRequestContext, params: types.SubscriptionsListenRequestParams + ) -> types.SubscriptionsListenResult: + assert params.notifications.tools_list_changed is True + try: + return await handler(ctx, params) + finally: + with anyio.fail_after(5, shield=True): + cleanup_started.set() + await release_cleanup.wait() + cleanup_finished.set() + + server = Server("subs", on_subscriptions_listen=slow_cleanup) + with anyio.fail_after(5): + async with Client(server) as client, anyio.create_task_group() as tg: + subscription = client.listen(tools_list_changed=True) + await subscription.__aenter__() + + async def close() -> None: + with anyio.CancelScope() as scope: + scope.cancel() + try: + await anyio.Event().wait() + except anyio.get_cancelled_exc_class() as exc: + await subscription.__aexit__(type(exc), exc, exc.__traceback__) + assert cleanup_finished.is_set() + closed.set() + raise + + tg.start_soon(close) + try: + await cleanup_started.wait() + await anyio.wait_all_tasks_blocked() + assert not closed.is_set() + finally: + release_cleanup.set() + await closed.wait() + + +@pytest.mark.parametrize( + "anyio_backend", + [pytest.param(("trio", {"clock": MockClock(autojump_threshold=0)}), id="trio-mockclock")], +) +async def test_exiting_a_subscription_bounds_uncooperative_direct_handler_cleanup() -> None: + """SDK-defined: exit stops waiting after five seconds if a direct handler shields its cleanup.""" + handler = ListenHandler(InMemorySubscriptionBus()) + cleanup_started = anyio.Event() + release_cleanup = anyio.Event() + cleanup_finished = anyio.Event() + + async def shielded_cleanup( + ctx: ServerRequestContext, params: types.SubscriptionsListenRequestParams + ) -> types.SubscriptionsListenResult: + assert params.notifications.tools_list_changed is True + try: + return await handler(ctx, params) + finally: + # The watchdog must outlast the SDK's five-second cleanup cap. + with anyio.fail_after(10, shield=True): + cleanup_started.set() + await release_cleanup.wait() + cleanup_finished.set() + + server = Server("subs", on_subscriptions_listen=shielded_cleanup) + async with Client(server) as client: + subscription = client.listen(tools_list_changed=True) + with anyio.fail_after(5): + sub = await subscription.__aenter__() + try: + started = anyio.current_time() + await subscription.__aexit__(None, None, None) + assert anyio.current_time() - started == 5 # MockClock time, never wall-clock time. + assert cleanup_started.is_set() + assert not cleanup_finished.is_set() + with pytest.raises(StopAsyncIteration): + await anext(sub) + finally: + release_cleanup.set() + with anyio.fail_after(5): + await cleanup_finished.wait() + + async def test_concurrent_subscriptions_demux_independently(): """Two open subscriptions each receive only their own filter's events.""" bus = InMemorySubscriptionBus() @@ -599,6 +695,8 @@ async def test_client_listen_installs_the_cache_eviction_barrier_exactly_when_a_ with anyio.fail_after(5): async with uncached_client.listen(tools_list_changed=True) as sub: # pragma: no branch assert sub._on_event is None # pyright: ignore[reportPrivateUsage] + await bus.publish(ToolsListChanged()) + assert await anext(sub) == ToolsListChanged() async def test_the_cache_eviction_barrier_maps_events_and_contains_store_faults( diff --git a/tests/conftest.py b/tests/conftest.py index 9ade27e7f3..6813a6f320 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -18,14 +18,14 @@ import mcp.shared._otel # noqa: E402 -@pytest.fixture(scope="session") -def anyio_backend() -> str: - return "asyncio" +@pytest.fixture(scope="session", params=["asyncio", "trio"]) +def anyio_backend(request: pytest.FixtureRequest) -> str: + return request.param @pytest.fixture(scope="module", autouse=True) async def _module_runner_lease(anyio_backend: str) -> AsyncIterator[None]: - """Share one event loop across each module's tests instead of one per test. + """Share one event loop per module and backend instead of one per test. anyio's pytest plugin tears its runner down whenever the last lease is released, so with only function-scoped async fixtures every async test @@ -34,15 +34,15 @@ async def _module_runner_lease(anyio_backend: str) -> AsyncIterator[None]: run can transiently exhaust kernel socket buffers — surfacing in CI as `OSError: [WinError 10055]` raised from `asyncio.new_event_loop()` before an arbitrary test's body even starts. Holding a module-scoped lease caps - the churn at one loop per module per xdist worker. + the churn at one loop per module and backend per xdist worker. Modules that parametrize `anyio_backend` or call `trio.run(...)` directly must shadow this fixture with a sync no-op: a module-scoped lease cannot depend on the function-scoped parameter (pytest raises ScopeMismatch at setup), and the lease's live asyncio loop lingers over direct trio runs, whose signal handling collides with the loop's wakeup fd on Windows. The - lease also makes sniffio report asyncio to the module's sync tests, so a - sync test must not call `anyio.run()` itself. + lease also makes sniffio report the leased backend to the module's sync + tests, so a sync test must not call `anyio.run()` itself. """ yield diff --git a/tests/docs_src/test_asgi.py b/tests/docs_src/test_asgi.py index d241dc71f9..8653206642 100644 --- a/tests/docs_src/test_asgi.py +++ b/tests/docs_src/test_asgi.py @@ -1,6 +1,7 @@ """`docs/run/asgi.md`: every claim the page makes, proved against the real SDK.""" import inspect +from importlib import reload import httpx2 import pytest @@ -113,6 +114,7 @@ async def about(request: Request) -> Response: async def test_the_host_lifespan_enters_the_session_manager() -> None: """tutorial002: the host app's lifespan owns `session_manager.run()` and starts and stops cleanly.""" + reload(tutorial002) async with tutorial002.lifespan(tutorial002.app): async with Client(tutorial002.mcp) as client: result = await client.call_tool("add_note", {"text": "milk"}) @@ -130,6 +132,7 @@ async def test_two_servers_get_two_mounts() -> None: async def test_one_lifespan_starts_both_session_managers() -> None: """tutorial003: a single `AsyncExitStack` lifespan runs both managers; both servers answer.""" + reload(tutorial003) async with tutorial003.lifespan(tutorial003.app): async with Client(tutorial003.notes) as client: notes_result = await client.call_tool("add_note", {"text": "milk"}) @@ -215,6 +218,7 @@ async def test_the_default_app_is_localhost_only() -> None: async def test_the_documented_browser_origin_works_end_to_end() -> None: """tutorial005: the page's scenario for real. The public hostname, the browser origin, a realistic preflight naming the `Mcp-*` headers, then the actual request.""" + reload(tutorial005) transport = httpx2.ASGITransport(app=tutorial005.app) async with tutorial005.lifespan(tutorial005.app): async with httpx2.AsyncClient(transport=transport, base_url="https://mcp.example.com") as http: diff --git a/tests/docs_src/test_context.py b/tests/docs_src/test_context.py index 617d113b2b..2c41ab1505 100644 --- a/tests/docs_src/test_context.py +++ b/tests/docs_src/test_context.py @@ -1,6 +1,7 @@ """`docs/handlers/context.md`: every claim the page makes, proved against the real SDK.""" import re +from importlib import reload import pytest from inline_snapshot import snapshot @@ -62,6 +63,7 @@ async def test_a_context_only_tool_takes_no_arguments() -> None: async def test_register_a_tool_at_runtime_and_notify_the_client() -> None: """tutorial003: `mcp.add_tool` takes effect immediately and `send_tool_list_changed` reaches the client.""" + reload(tutorial003) messages: list[object] = [] async def collect(message: object) -> None: diff --git a/tests/docs_src/test_deploy.py b/tests/docs_src/test_deploy.py index 268c7ab74d..442c35be4a 100644 --- a/tests/docs_src/test_deploy.py +++ b/tests/docs_src/test_deploy.py @@ -1,5 +1,7 @@ """`docs/run/deploy.md`: every claim the page makes, proved against the real SDK.""" +from importlib import reload + import anyio import httpx2 import pytest @@ -51,6 +53,7 @@ async def test_the_default_app_rejects_a_real_hostname_before_mcp_runs() -> None async def test_the_allowlisted_app_serves_its_hostname_and_still_rejects_others() -> None: """tutorial001: `allowed_hosts=` opens exactly the hostname you named, and nothing else.""" + reload(tutorial001) transport = httpx2.ASGITransport(app=tutorial001.app) async with tutorial001.mcp.session_manager.run(): async with httpx2.AsyncClient(transport=transport, base_url="https://mcp.example.com") as http: diff --git a/tests/docs_src/test_identity_assertion.py b/tests/docs_src/test_identity_assertion.py index 3a15ef94ec..6e172edf67 100644 --- a/tests/docs_src/test_identity_assertion.py +++ b/tests/docs_src/test_identity_assertion.py @@ -1,6 +1,7 @@ """`docs/client/identity-assertion.md`: every claim the page makes, proved against the real SDK.""" import inspect +from importlib import reload from urllib.parse import parse_qsl import httpx2 @@ -144,6 +145,7 @@ async def test_the_metadata_advertises_the_grant_type_and_the_id_jag_profile() - async def test_the_whole_grant_is_one_token_request() -> None: """The `!!! check`: a 401, the well-known fetch, one `POST /token`, the retry; the subject reaches the tool.""" + reload(tutorial001) mcp = MCPServer( "Notes", token_verifier=ProviderTokenVerifier(tutorial002.provider), diff --git a/tests/docs_src/test_legacy_clients.py b/tests/docs_src/test_legacy_clients.py index 90daf1bd96..e5065ddea1 100644 --- a/tests/docs_src/test_legacy_clients.py +++ b/tests/docs_src/test_legacy_clients.py @@ -1,6 +1,7 @@ """`docs/run/legacy-clients.md`: every claim the page makes, proved against the real SDK.""" import inspect +from importlib import reload import httpx2 import pytest @@ -115,6 +116,7 @@ async def test_stateless_http_never_mints_a_session() -> None: async def test_stateless_http_kills_the_legacy_back_channel_and_only_the_legacy_one() -> None: """tutorial002: over the same `stateless_http=True` app, the modern client still gets its answer and the legacy client's call fails as the top-level `MCPError` the `!!! check` quotes.""" + reload(tutorial002) async with ( tutorial002.app.router.lifespan_context(tutorial002.app), httpx2.ASGITransport(tutorial002.app) as transport, diff --git a/tests/docs_src/test_subscriptions.py b/tests/docs_src/test_subscriptions.py index e2d9bf2c77..dd98eb12a8 100644 --- a/tests/docs_src/test_subscriptions.py +++ b/tests/docs_src/test_subscriptions.py @@ -1,6 +1,7 @@ """`docs/{handlers,client}/subscriptions.md`: every claim the two pages make, proved against the real SDK.""" -from collections.abc import Awaitable, Callable +from collections.abc import AsyncIterator +from contextlib import asynccontextmanager from typing import Any import anyio @@ -19,18 +20,14 @@ tutorial006, ) from mcp import Client +from mcp.client.subscriptions import Subscription from mcp.server.auth.middleware.auth_context import auth_context_var from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser from mcp.server.auth.provider import AccessToken -from mcp.server.context import ServerRequestContext from mcp.server.lowlevel import Server from mcp.server.subscriptions import SUBSCRIPTION_ID_META_KEY, ListenHandler, ToolsListChanged from mcp.shared.exceptions import MCPError -_ReadResource = Callable[ - [ServerRequestContext[Any], types.ReadResourceRequestParams], Awaitable[types.ReadResourceResult] -] - # See test_index.py for why this is a per-module mark and not a conftest hook. pytestmark = [pytest.mark.anyio, pytest.mark.filterwarnings("error::mcp.MCPDeprecationWarning")] @@ -68,24 +65,18 @@ async def wait_for(self, count: int) -> None: await self._arrival.wait() -class _Reads: - """Counts server-side resource reads so a test can await the Nth refetch.""" +class _Output: + """Counts completed prints so tests wait for client output, not server-side reads.""" def __init__(self) -> None: self.count = 0 self._bump = anyio.Event() - def counting(self, handler: _ReadResource) -> _ReadResource: - async def counted( - ctx: ServerRequestContext[Any], params: types.ReadResourceRequestParams - ) -> types.ReadResourceResult: - result = await handler(ctx, params) - self.count += 1 - self._bump.set() - self._bump = anyio.Event() - return result - - return counted + def __call__(self, *values: str | list[str]) -> None: + print(*values) + self.count += 1 + self._bump.set() + self._bump = anyio.Event() async def wait_for(self, count: int) -> None: with anyio.fail_after(5): @@ -209,18 +200,34 @@ async def listen() -> None: async def test_follow_board_prints_the_refetched_board_and_the_new_tool_list( - capsys: pytest.CaptureFixture[str], + capsys: pytest.CaptureFixture[str], monkeypatch: pytest.MonkeyPatch ) -> None: """tutorial003: each event drives a refetch - the board reprints, and a tools change reprints the tool names.""" + output = _Output() + monkeypatch.setattr(tutorial003, "print", output, raising=False) + listening = anyio.Event() async with Client(tutorial001.mcp) as client: + listen = client.listen + + @asynccontextmanager + async def listen_when_ready( + *, tools_list_changed: bool, resource_subscriptions: list[str] + ) -> AsyncIterator[Subscription]: + async with listen( + tools_list_changed=tools_list_changed, resource_subscriptions=resource_subscriptions + ) as sub: + listening.set() + yield sub + + monkeypatch.setattr(client, "listen", listen_when_ready) async with anyio.create_task_group() as tg: tg.start_soon(tutorial003.follow_board, client) - # Let the watcher park on its stream (ack complete) before publishing. - await anyio.wait_all_tasks_blocked() + with anyio.fail_after(5): + await listening.wait() await client.call_tool("complete_task", {"board": "sprint", "task": "design"}) - await anyio.wait_all_tasks_blocked() + await output.wait_for(1) await client.call_tool("enable_reports", {}) - await anyio.wait_all_tasks_blocked() + await output.wait_for(2) tg.cancel_scope.cancel() printed = capsys.readouterr().out @@ -243,13 +250,15 @@ def _assert_snapshot_then_current_board(printed: str) -> None: assert printed.strip().endswith(FINISHED_BOARD), printed +@pytest.mark.parametrize("anyio_backend", [pytest.param("asyncio", id="asyncio")]) async def test_the_asyncio_watcher_runs_beside_the_main_flow(capsys: pytest.CaptureFixture[str]) -> None: """tutorial004 (asyncio tab): run_sprint opens the subscription, snapshots the board, then a watcher task reprints it while the main flow keeps calling tools. The example connects over HTTP; the in-memory client here is the maintainer-side stand-in.""" async with Client(tutorial001.mcp) as client: - await tutorial004_asyncio.run_sprint(client) + with anyio.fail_after(5): + await tutorial004_asyncio.run_sprint(client) _assert_snapshot_then_current_board(capsys.readouterr().out) @@ -257,14 +266,16 @@ async def test_the_asyncio_watcher_runs_beside_the_main_flow(capsys: pytest.Capt async def test_the_trio_watcher_runs_beside_the_main_flow(capsys: pytest.CaptureFixture[str]) -> None: """tutorial004 (trio tab): the same shape as the asyncio tab, with a nursery owning the watcher.""" async with Client(tutorial001.mcp) as client: - await tutorial004_trio.run_sprint(client) + with anyio.fail_after(5): + await tutorial004_trio.run_sprint(client) _assert_snapshot_then_current_board(capsys.readouterr().out) async def test_the_anyio_watcher_runs_beside_the_main_flow(capsys: pytest.CaptureFixture[str]) -> None: """tutorial004 (anyio tab): the same shape again, with a task group owning the watcher.""" async with Client(tutorial001.mcp) as client: - await tutorial004_anyio.run_sprint(client) + with anyio.fail_after(5): + await tutorial004_anyio.run_sprint(client) _assert_snapshot_then_current_board(capsys.readouterr().out) @@ -272,16 +283,19 @@ async def test_the_anyio_watcher_runs_beside_the_main_flow(capsys: pytest.Captur "anyio_backend", [pytest.param(("trio", {"clock": MockClock(autojump_threshold=0)}), id="trio-mockclock")], ) -async def test_the_follower_re_listens_after_the_stream_ends(capsys: pytest.CaptureFixture[str]) -> None: +async def test_the_follower_re_listens_after_the_stream_ends( + capsys: pytest.CaptureFixture[str], monkeypatch: pytest.MonkeyPatch +) -> None: """tutorial005: a graceful server close ends one stream; the loop backs off, re-listens, and refetches. Runs on trio's autojumping MockClock so the loop's backoff sleep takes no wall-clock time. """ - reads = _Reads() + output = _Output() + monkeypatch.setattr(tutorial005, "print", output, raising=False) handler = ListenHandler(tutorial002.bus) server = Server( "sprint-board", - on_read_resource=reads.counting(tutorial002.read_resource), + on_read_resource=tutorial002.read_resource, on_list_tools=tutorial002.list_tools, on_call_tool=tutorial002.call_tool, on_subscriptions_listen=handler, @@ -290,17 +304,16 @@ async def test_the_follower_re_listens_after_the_stream_ends(capsys: pytest.Capt async with Client(server) as client: async with anyio.create_task_group() as tg: tg.start_soon(tutorial005.keep_following, client) - # First stream: the entry refetch reads the board, then an event reads it again. - await reads.wait_for(1) + # First stream: print the entry snapshot, then the board after the event. + await output.wait_for(1) await client.call_tool("complete_task", {"task": "design"}) - await reads.wait_for(2) + await output.wait_for(2) - # End that stream gracefully. The loop backs off (the mock clock jumps the - # sleep), re-listens, and refetches on entry: that is the third read. + # The mock clock jumps the backoff; the third print is the new stream's snapshot. handler.close() - await reads.wait_for(3) + await output.wait_for(3) await client.call_tool("complete_task", {"task": "build"}) - await reads.wait_for(4) + await output.wait_for(4) tg.cancel_scope.cancel() printed = capsys.readouterr().out diff --git a/tests/docs_src/test_troubleshooting.py b/tests/docs_src/test_troubleshooting.py index 1e1b5e15b8..ab69bb86b3 100644 --- a/tests/docs_src/test_troubleshooting.py +++ b/tests/docs_src/test_troubleshooting.py @@ -1,6 +1,7 @@ """`docs/troubleshooting.md`: every error string the page names, reproduced against the real SDK.""" import logging +from importlib import reload from typing import Any import httpx2 @@ -147,6 +148,7 @@ async def test_the_default_streamable_http_app_answers_a_real_hostname_with_421( caplog: pytest.LogCaptureFixture, ) -> None: """tutorial003: one 421, three spellings. The page presents all three as the same event.""" + reload(tutorial003) transport = httpx2.ASGITransport(app=tutorial003.app) async with tutorial003.mcp.session_manager.run(): # What curl (or the reverse proxy's access log) shows: the status and the plain-text body. @@ -169,6 +171,7 @@ async def test_the_default_streamable_http_app_answers_a_real_hostname_with_421( async def test_an_allowlisted_hostname_connects_and_calls_a_tool() -> None: """tutorial004: `transport_security=` names the deployed hostname, and the same client connects.""" + reload(tutorial004) transport = httpx2.ASGITransport(app=tutorial004.app) async with tutorial004.mcp.session_manager.run(): async with httpx2.AsyncClient(transport=transport) as http_client: @@ -265,6 +268,7 @@ async def test_a_legacy_ctx_elicit_without_a_callback_says_elicitation_not_suppo async def test_ctx_elicit_over_stateless_http_has_no_back_channel() -> None: """tutorial008: `stateless_http=True` leaves the server no channel to send `elicitation/create`.""" + reload(tutorial008) transport = httpx2.ASGITransport(app=tutorial008.app) async with tutorial008.mcp.session_manager.run(): async with httpx2.AsyncClient(transport=transport) as http_client: diff --git a/tests/interaction/transports/test_hosting_resume.py b/tests/interaction/transports/test_hosting_resume.py index 2699ab32f1..df917513db 100644 --- a/tests/interaction/transports/test_hosting_resume.py +++ b/tests/interaction/transports/test_hosting_resume.py @@ -10,6 +10,7 @@ """ import json +from collections.abc import AsyncGenerator import anyio import httpx2 @@ -70,9 +71,13 @@ def _tools_call(request_id: int, name: str, arguments: dict[str, object]) -> str async def _read_events(response: httpx2.Response, count: int) -> list[ServerSentEvent]: - """Read exactly `count` SSE events from a streaming response without closing it.""" + """Read exactly `count` SSE events and close the iterator.""" source = aiter(EventSource(response)) - return [await anext(source) for _ in range(count)] + try: + return [await anext(source) for _ in range(count)] + finally: + assert isinstance(source, AsyncGenerator) + await source.aclose() @requirement("hosting:resume:event-ids") diff --git a/tests/issues/test_1363_race_condition_streamable_http.py b/tests/issues/test_1363_race_condition_streamable_http.py index f98194b7b5..1b0dc885f1 100644 --- a/tests/issues/test_1363_race_condition_streamable_http.py +++ b/tests/issues/test_1363_race_condition_streamable_http.py @@ -16,12 +16,10 @@ """ import logging -import threading from collections.abc import AsyncGenerator from contextlib import asynccontextmanager import anyio -import anyio.to_thread import httpx2 import pytest from starlette.applications import Starlette @@ -62,42 +60,6 @@ async def lifespan(app: Starlette) -> AsyncGenerator[None, None]: return Starlette(routes=routes, lifespan=lifespan) -class ServerThread(threading.Thread): - """Thread that runs the ASGI application lifespan in a separate event loop.""" - - def __init__(self, app: Starlette): - super().__init__(daemon=True) - self.app = app - self._stop_event = threading.Event() - self._ready_event = threading.Event() - - def run(self) -> None: - """Run the lifespan in a new event loop.""" - - # Create a new event loop for this thread - async def run_lifespan(): - # Use the lifespan context (always present in our tests) - lifespan_context = getattr(self.app.router, "lifespan_context", None) - assert lifespan_context is not None # Tests always create apps with lifespan - async with lifespan_context(self.app): - # Only signal readiness once lifespan startup has completed, i.e. the - # session manager's task group exists and requests can be handled. - self._ready_event.set() - # Wait until stop is requested - while not self._stop_event.is_set(): - await anyio.sleep(0.1) - - anyio.run(run_lifespan) - - def wait_ready(self, timeout: float = 5.0) -> None: - """Block until the lifespan has started; call from a worker thread, not the event loop.""" - assert self._ready_event.wait(timeout), "server thread did not start its lifespan in time" - - def stop(self) -> None: - """Signal the thread to stop.""" - self._stop_event.set() - - def check_logs_for_race_condition_errors(caplog: pytest.LogCaptureFixture, test_name: str) -> None: """Check logs for ClosedResourceError and other race condition errors. @@ -128,7 +90,7 @@ def check_logs_for_race_condition_errors(caplog: pytest.LogCaptureFixture, test_ @pytest.mark.anyio -async def test_race_condition_invalid_accept_headers(caplog: pytest.LogCaptureFixture): +async def test_race_condition_invalid_accept_headers(caplog: pytest.LogCaptureFixture) -> None: """Test the race condition with invalid Accept headers. This test reproduces the exact scenario described in issue #1363: @@ -137,15 +99,8 @@ async def test_race_condition_invalid_accept_headers(caplog: pytest.LogCaptureFi - This should trigger the race condition where message_router encounters ClosedResourceError """ app = create_app() - server_thread = ServerThread(app) - server_thread.start() - - try: - # Wait for the server thread to enter the lifespan before sending requests - await anyio.to_thread.run_sync(server_thread.wait_ready) - - # Suppress WARNING logs (expected validation errors) and capture ERROR logs - with caplog.at_level(logging.ERROR): + with caplog.at_level(logging.ERROR), anyio.fail_after(5): + async with app.router.lifespan_context(app): # Test with missing text/event-stream in Accept header async with httpx2.AsyncClient( transport=httpx2.ASGITransport(app=app), base_url="http://testserver", timeout=5.0 @@ -191,32 +146,21 @@ async def test_race_condition_invalid_accept_headers(caplog: pytest.LogCaptureFi # Should get 406 Not Acceptable assert response.status_code == 406 - # Give background tasks time to complete - await anyio.sleep(0.2) + # Let the message routers finish before lifespan shutdown cancels them. + await anyio.wait_all_tasks_blocked() - finally: - server_thread.stop() - server_thread.join(timeout=5.0) - # Check logs for race condition errors - check_logs_for_race_condition_errors(caplog, "test_race_condition_invalid_accept_headers") + check_logs_for_race_condition_errors(caplog, "test_race_condition_invalid_accept_headers") @pytest.mark.anyio -async def test_race_condition_invalid_content_type(caplog: pytest.LogCaptureFixture): +async def test_race_condition_invalid_content_type(caplog: pytest.LogCaptureFixture) -> None: """Test the race condition with invalid Content-Type headers. This test reproduces the race condition scenario with Content-Type validation failure. """ app = create_app() - server_thread = ServerThread(app) - server_thread.start() - - try: - # Wait for the server thread to enter the lifespan before sending requests - await anyio.to_thread.run_sync(server_thread.wait_ready) - - # Suppress WARNING logs (expected validation errors) and capture ERROR logs - with caplog.at_level(logging.ERROR): + with caplog.at_level(logging.ERROR), anyio.fail_after(5): + async with app.router.lifespan_context(app): # Test with invalid Content-Type async with httpx2.AsyncClient( transport=httpx2.ASGITransport(app=app), base_url="http://testserver", timeout=5.0 @@ -231,32 +175,21 @@ async def test_race_condition_invalid_content_type(caplog: pytest.LogCaptureFixt ) assert response.status_code == 400 - # Give background tasks time to complete - await anyio.sleep(0.2) + # Let the message router finish before lifespan shutdown cancels it. + await anyio.wait_all_tasks_blocked() - finally: - server_thread.stop() - server_thread.join(timeout=5.0) - # Check logs for race condition errors - check_logs_for_race_condition_errors(caplog, "test_race_condition_invalid_content_type") + check_logs_for_race_condition_errors(caplog, "test_race_condition_invalid_content_type") @pytest.mark.anyio -async def test_race_condition_message_router_async_for(caplog: pytest.LogCaptureFixture): +async def test_race_condition_message_router_async_for(caplog: pytest.LogCaptureFixture) -> None: """Uses json_response=True to trigger the `if self.is_json_response_enabled` branch, which reproduces the ClosedResourceError when message_router is suspended in async for loop while transport cleanup closes streams concurrently. """ app = create_app(json_response=True) - server_thread = ServerThread(app) - server_thread.start() - - try: - # Wait for the server thread to enter the lifespan before sending requests - await anyio.to_thread.run_sync(server_thread.wait_ready) - - # Suppress WARNING logs (expected validation errors) and capture ERROR logs - with caplog.at_level(logging.ERROR): + with caplog.at_level(logging.ERROR), anyio.fail_after(5): + async with app.router.lifespan_context(app): # Use httpx2.ASGITransport to test the ASGI app directly async with httpx2.AsyncClient( transport=httpx2.ASGITransport(app=app), base_url="http://testserver", timeout=5.0 @@ -273,11 +206,7 @@ async def test_race_condition_message_router_async_for(caplog: pytest.LogCapture # Should get a successful response assert response.status_code in (200, 201) - # Give background tasks time to complete - await anyio.sleep(0.2) + # Let the message router finish before lifespan shutdown cancels it. + await anyio.wait_all_tasks_blocked() - finally: - server_thread.stop() - server_thread.join(timeout=5.0) - # Check logs for race condition errors in message router - check_logs_for_race_condition_errors(caplog, "test_race_condition_message_router_async_for") + check_logs_for_race_condition_errors(caplog, "test_race_condition_message_router_async_for") diff --git a/tests/server/test_sse_security.py b/tests/server/test_sse_security.py index 7e84428600..417e7d8d9f 100644 --- a/tests/server/test_sse_security.py +++ b/tests/server/test_sse_security.py @@ -9,8 +9,6 @@ import sse_starlette.sse from mcp_types import JSONRPCRequest, JSONRPCResponse from starlette.applications import Starlette -from starlette.requests import Request -from starlette.responses import Response from starlette.routing import Mount, Route from starlette.types import Message, Receive, Scope, Send @@ -45,20 +43,17 @@ def sse_security_client(security_settings: TransportSecuritySettings | None = No server = Server(SERVER_NAME) sse_transport = SseServerTransport("/messages/", security_settings) - async def handle_sse(request: Request) -> Response: - try: - async with sse_transport.connect_sse(request.scope, request.receive, request._send) as (read, write): - await server.run(read, write, server.create_initialization_options()) - except ValueError as e: - # Validation error was already handled inside connect_sse, which sent the rejection - # response itself; its non-empty body checkpoints, so the test reads the rejection - # status before the trailing Response() below sends a second response start. - logger.debug(f"SSE connection failed validation: {e}") - return Response() + class SSEApp: + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + try: + async with sse_transport.connect_sse(scope, receive, send) as (read, write): + await server.run(read, write, server.create_initialization_options()) + except ValueError: + logger.exception("SSE connection failed validation") app = Starlette( routes=[ - Route("/sse", endpoint=handle_sse), + Route("/sse", endpoint=SSEApp()), Mount("/messages/", app=sse_transport.handle_post_message), ] ) diff --git a/tests/server/test_streamable_http_router.py b/tests/server/test_streamable_http_router.py index 0c5796c1f5..7f93a89253 100644 --- a/tests/server/test_streamable_http_router.py +++ b/tests/server/test_streamable_http_router.py @@ -1,8 +1,9 @@ """Regression coverage for the StreamableHTTP per-session response router.""" import anyio +import httpx2 import pytest -from mcp_types import JSONRPCMessage, JSONRPCResponse +from mcp_types import JSONRPCMessage, JSONRPCNotification, JSONRPCResponse, jsonrpc_message_adapter from starlette.types import Message, Scope from mcp.server.streamable_http import ( @@ -157,3 +158,138 @@ async def test_terminated_transport_answers_404() -> None: assert post.sent[0]["type"] == "http.response.start" 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 +async def test_transport_shutdown_cancels_a_router_waiting_for_replay() -> None: + """SDK-defined: a stalled replay cannot keep the session's response router alive after termination.""" + replay_started = 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() + return await anyio.sleep_forever() + + 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())], + } + + async def receive() -> Message: + await anyio.sleep_forever() + raise NotImplementedError + + async def send(message: Message) -> None: + sent.append(message) + + with anyio.fail_after(5): + async with anyio.create_task_group() as requests: + async with transport.connect() as (_, write_stream): + requests.start_soon(transport.handle_request, scope, receive, send) + await replay_started.wait() + await write_stream.send(SessionMessage(JSONRPCResponse(jsonrpc="2.0", id="request", result={}))) + await anyio.wait_all_tasks_blocked() + await transport.terminate() + requests.cancel_scope.cancel() + + assert transport.is_terminated + assert sent[0]["type"] == "http.response.start" + assert sent[0]["status"] == 200 diff --git a/tests/shared/test_sse.py b/tests/shared/test_sse.py index 77d1b28a0a..ef47ca774b 100644 --- a/tests/shared/test_sse.py +++ b/tests/shared/test_sse.py @@ -120,8 +120,12 @@ async def test_raw_sse_connection() -> None: assert response.headers["content-type"] == "text/event-stream; charset=utf-8" lines = response.aiter_lines() - assert await anext(lines) == "event: endpoint" - assert (await anext(lines)).startswith("data: /messages/?session_id=") + try: + assert await anext(lines) == "event: endpoint" + assert (await anext(lines)).startswith("data: /messages/?session_id=") + finally: + assert isinstance(lines, AsyncGenerator) + await lines.aclose() @pytest.mark.anyio From 48484957209d64a636fb575a0bcd732759cb7e38 Mon Sep 17 00:00:00 2001 From: Marcelo Trylesinski Date: Wed, 16 Sep 2026 12:40:44 +0200 Subject: [PATCH 2/7] Ignore Trio's Linux pidfd encoding warning --- pyproject.toml | 2 ++ 1 file changed, 2 insertions(+) diff --git a/pyproject.toml b/pyproject.toml index b2f26da55f..0aa254b25d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -284,6 +284,8 @@ filterwarnings = [ # CI and scripts/test set PYTHONWARNDEFAULTENCODING=1, so "error" rejects any text I/O # of ours that omits encoding=; pytest-examples' own unguarded text I/O isn't ours. "ignore:'encoding' argument not specified:EncodingWarning:pytest_examples", + # Trio wraps Linux pidfds in text mode without encoding=; only fileno()/close() are used. + "ignore:'encoding' argument not specified:EncodingWarning:trio\\._subprocess$", ] [tool.markdown.lint] From 249279a6434193c38ac4e9402da55ec95602f688 Mon Sep 17 00:00:00 2001 From: Marcelo Trylesinski Date: Wed, 16 Sep 2026 13:12:23 +0200 Subject: [PATCH 3/7] Run race-condition log checks on every exit path --- ...est_1363_race_condition_streamable_http.py | 189 +++++++++--------- 1 file changed, 96 insertions(+), 93 deletions(-) diff --git a/tests/issues/test_1363_race_condition_streamable_http.py b/tests/issues/test_1363_race_condition_streamable_http.py index 1b0dc885f1..f290dc35b4 100644 --- a/tests/issues/test_1363_race_condition_streamable_http.py +++ b/tests/issues/test_1363_race_condition_streamable_http.py @@ -99,57 +99,58 @@ async def test_race_condition_invalid_accept_headers(caplog: pytest.LogCaptureFi - This should trigger the race condition where message_router encounters ClosedResourceError """ app = create_app() - with caplog.at_level(logging.ERROR), anyio.fail_after(5): - async with app.router.lifespan_context(app): - # Test with missing text/event-stream in Accept header - async with httpx2.AsyncClient( - transport=httpx2.ASGITransport(app=app), base_url="http://testserver", timeout=5.0 - ) as client: - response = await client.post( - "/", - json={"jsonrpc": "2.0", "method": "initialize", "id": 1, "params": {}}, - headers={ - "Accept": "application/json", # Missing text/event-stream - "Content-Type": "application/json", - }, - ) - # Should get 406 Not Acceptable due to missing text/event-stream - assert response.status_code == 406 - - # Test with missing application/json in Accept header - async with httpx2.AsyncClient( - transport=httpx2.ASGITransport(app=app), base_url="http://testserver", timeout=5.0 - ) as client: - response = await client.post( - "/", - json={"jsonrpc": "2.0", "method": "initialize", "id": 1, "params": {}}, - headers={ - "Accept": "text/event-stream", # Missing application/json - "Content-Type": "application/json", - }, - ) - # Should get 406 Not Acceptable due to missing application/json - assert response.status_code == 406 - - # Test with completely invalid Accept header - async with httpx2.AsyncClient( - transport=httpx2.ASGITransport(app=app), base_url="http://testserver", timeout=5.0 - ) as client: - response = await client.post( - "/", - json={"jsonrpc": "2.0", "method": "initialize", "id": 1, "params": {}}, - headers={ - "Accept": "text/plain", # Invalid Accept header - "Content-Type": "application/json", - }, - ) - # Should get 406 Not Acceptable - assert response.status_code == 406 - - # Let the message routers finish before lifespan shutdown cancels them. - await anyio.wait_all_tasks_blocked() - - check_logs_for_race_condition_errors(caplog, "test_race_condition_invalid_accept_headers") + try: + with caplog.at_level(logging.ERROR), anyio.fail_after(5): + async with app.router.lifespan_context(app): + # Test with missing text/event-stream in Accept header + async with httpx2.AsyncClient( + transport=httpx2.ASGITransport(app=app), base_url="http://testserver", timeout=5.0 + ) as client: + response = await client.post( + "/", + json={"jsonrpc": "2.0", "method": "initialize", "id": 1, "params": {}}, + headers={ + "Accept": "application/json", # Missing text/event-stream + "Content-Type": "application/json", + }, + ) + # Should get 406 Not Acceptable due to missing text/event-stream + assert response.status_code == 406 + + # Test with missing application/json in Accept header + async with httpx2.AsyncClient( + transport=httpx2.ASGITransport(app=app), base_url="http://testserver", timeout=5.0 + ) as client: + response = await client.post( + "/", + json={"jsonrpc": "2.0", "method": "initialize", "id": 1, "params": {}}, + headers={ + "Accept": "text/event-stream", # Missing application/json + "Content-Type": "application/json", + }, + ) + # Should get 406 Not Acceptable due to missing application/json + assert response.status_code == 406 + + # Test with completely invalid Accept header + async with httpx2.AsyncClient( + transport=httpx2.ASGITransport(app=app), base_url="http://testserver", timeout=5.0 + ) as client: + response = await client.post( + "/", + json={"jsonrpc": "2.0", "method": "initialize", "id": 1, "params": {}}, + headers={ + "Accept": "text/plain", # Invalid Accept header + "Content-Type": "application/json", + }, + ) + # Should get 406 Not Acceptable + assert response.status_code == 406 + + # Let the message routers finish before lifespan shutdown cancels them. + await anyio.wait_all_tasks_blocked() + finally: + check_logs_for_race_condition_errors(caplog, "test_race_condition_invalid_accept_headers") @pytest.mark.anyio @@ -159,26 +160,27 @@ async def test_race_condition_invalid_content_type(caplog: pytest.LogCaptureFixt This test reproduces the race condition scenario with Content-Type validation failure. """ app = create_app() - with caplog.at_level(logging.ERROR), anyio.fail_after(5): - async with app.router.lifespan_context(app): - # Test with invalid Content-Type - async with httpx2.AsyncClient( - transport=httpx2.ASGITransport(app=app), base_url="http://testserver", timeout=5.0 - ) as client: - response = await client.post( - "/", - json={"jsonrpc": "2.0", "method": "initialize", "id": 1, "params": {}}, - headers={ - "Accept": "application/json, text/event-stream", - "Content-Type": "text/plain", # Invalid Content-Type - }, - ) - assert response.status_code == 400 - - # Let the message router finish before lifespan shutdown cancels it. - await anyio.wait_all_tasks_blocked() - - check_logs_for_race_condition_errors(caplog, "test_race_condition_invalid_content_type") + try: + with caplog.at_level(logging.ERROR), anyio.fail_after(5): + async with app.router.lifespan_context(app): + # Test with invalid Content-Type + async with httpx2.AsyncClient( + transport=httpx2.ASGITransport(app=app), base_url="http://testserver", timeout=5.0 + ) as client: + response = await client.post( + "/", + json={"jsonrpc": "2.0", "method": "initialize", "id": 1, "params": {}}, + headers={ + "Accept": "application/json, text/event-stream", + "Content-Type": "text/plain", # Invalid Content-Type + }, + ) + assert response.status_code == 400 + + # Let the message router finish before lifespan shutdown cancels it. + await anyio.wait_all_tasks_blocked() + finally: + check_logs_for_race_condition_errors(caplog, "test_race_condition_invalid_content_type") @pytest.mark.anyio @@ -188,25 +190,26 @@ async def test_race_condition_message_router_async_for(caplog: pytest.LogCapture in async for loop while transport cleanup closes streams concurrently. """ app = create_app(json_response=True) - with caplog.at_level(logging.ERROR), anyio.fail_after(5): - async with app.router.lifespan_context(app): - # Use httpx2.ASGITransport to test the ASGI app directly - async with httpx2.AsyncClient( - transport=httpx2.ASGITransport(app=app), base_url="http://testserver", timeout=5.0 - ) as client: - # Send a valid initialize request - response = await client.post( - "/", - json={"jsonrpc": "2.0", "method": "initialize", "id": 1, "params": {}}, - headers={ - "Accept": "application/json, text/event-stream", - "Content-Type": "application/json", - }, - ) - # Should get a successful response - assert response.status_code in (200, 201) - - # Let the message router finish before lifespan shutdown cancels it. - await anyio.wait_all_tasks_blocked() - - check_logs_for_race_condition_errors(caplog, "test_race_condition_message_router_async_for") + try: + with caplog.at_level(logging.ERROR), anyio.fail_after(5): + async with app.router.lifespan_context(app): + # Use httpx2.ASGITransport to test the ASGI app directly + async with httpx2.AsyncClient( + transport=httpx2.ASGITransport(app=app), base_url="http://testserver", timeout=5.0 + ) as client: + # Send a valid initialize request + response = await client.post( + "/", + json={"jsonrpc": "2.0", "method": "initialize", "id": 1, "params": {}}, + headers={ + "Accept": "application/json, text/event-stream", + "Content-Type": "application/json", + }, + ) + # Should get a successful response + assert response.status_code in (200, 201) + + # Let the message router finish before lifespan shutdown cancels it. + await anyio.wait_all_tasks_blocked() + finally: + check_logs_for_race_condition_errors(caplog, "test_race_condition_message_router_async_for") From 9538140d3bd274802ed94e3ceb5c3f8a42740575 Mon Sep 17 00:00:00 2001 From: Marcelo Trylesinski Date: Wed, 16 Sep 2026 13:12:24 +0200 Subject: [PATCH 4/7] Preserve successor streams when closing reused request IDs --- docs/run/legacy-clients.md | 8 +- src/mcp/server/streamable_http.py | 213 ++++++++------------ src/mcp/shared/_httpx_utils.py | 2 +- tests/server/test_streamable_http_router.py | 124 +++++++++++- tests/shared/test_streamable_http.py | 60 +++--- 5 files changed, 238 insertions(+), 169 deletions(-) diff --git a/docs/run/legacy-clients.md b/docs/run/legacy-clients.md index 70aec64b3e..80773e7270 100644 --- a/docs/run/legacy-clients.md +++ b/docs/run/legacy-clients.md @@ -67,10 +67,10 @@ On one worker that is invisible. On two, it is the whole problem: a request that session reachable from another process. !!! note "Slow replay delays messages in the same session" - With `event_store=`, the SDK pauses new message delivery within the session until replay - finishes and the live stream is registered. This prevents responses from being lost between - replay and live delivery. A slow replay can delay other requests in that session, but not - other sessions. + With `event_store=`, the session's message router waits until replay finishes and the live + stream is registered. This prevents newly produced responses from being lost between replay + and live delivery. It does not serialize incoming POSTs or reserve JSON-RPC request IDs. + A slow replay can delay other requests in that session, but not other sessions. ## Session lifetime and limits diff --git a/src/mcp/server/streamable_http.py b/src/mcp/server/streamable_http.py index c320ee3d04..1f1547b129 100644 --- a/src/mcp/server/streamable_http.py +++ b/src/mcp/server/streamable_http.py @@ -15,7 +15,7 @@ from dataclasses import dataclass from functools import partial from http import HTTPStatus -from typing import Any, Final +from typing import Any, Final, TypeAlias import anyio import pydantic_core @@ -111,6 +111,7 @@ class EventMessage: EventCallback = Callable[[EventMessage], Awaitable[None]] +_RequestStreams: TypeAlias = tuple[MemoryObjectSendStream[EventMessage], MemoryObjectReceiveStream[EventMessage]] class EventStore(ABC): @@ -212,13 +213,7 @@ def __init__( self._event_store_lock = anyio.Lock() self._security = TransportSecurityMiddleware(security_settings) self._retry_interval = retry_interval - self._request_streams: dict[ - RequestId, - tuple[ - MemoryObjectSendStream[EventMessage], - MemoryObjectReceiveStream[EventMessage], - ], - ] = {} + self._request_streams: dict[RequestId, _RequestStreams] = {} self._sse_stream_writers: dict[RequestId, MemoryObjectSendStream[SSEEvent]] = {} self._terminated = False self._idle_timeout = idle_timeout @@ -356,12 +351,11 @@ async def _mint_priming_event(self, stream_id: StreamId, protocol_version: str) async def _run_sse_writer( self, - request_id: RequestId, sse_stream_writer: MemoryObjectSendStream[SSEEvent], request_stream_reader: MemoryObjectReceiveStream[EventMessage], priming_event: SSEEvent | None, ) -> None: - """Forward `_request_streams[request_id]` onto the SSE wire for one POST.""" + """Forward this POST's request stream onto the SSE wire.""" try: async with sse_stream_writer, request_stream_reader: if priming_event is not None: @@ -376,9 +370,6 @@ async def _run_sse_writer( logger.exception("Error in SSE writer") finally: logger.debug("Closing SSE writer") - if self._sse_stream_writers.get(request_id) is sse_stream_writer: - self._sse_stream_writers.pop(request_id, None) - await self._clean_up_memory_streams(request_id) def _create_error_response( self, @@ -459,19 +450,11 @@ async def _terminate_unanswered_request(self, request_id: RequestId) -> None: error = ErrorData(code=REQUEST_CANCELLED, message="Request cancelled") await self._write_stream.send(SessionMessage(JSONRPCError(jsonrpc="2.0", id=request_id, error=error))) - async def _clean_up_memory_streams(self, request_id: RequestId) -> None: - """Clean up memory streams for a given request ID.""" - if request_id in self._request_streams: # pragma: no branch - try: - # Close the request stream - await self._request_streams[request_id][0].aclose() - await self._request_streams[request_id][1].aclose() - except Exception: # pragma: no cover - # During cleanup, we catch all exceptions since streams might be in various states - logger.debug("Error closing memory streams - may already be closed") - finally: - # Remove the request stream from the mapping - self._request_streams.pop(request_id, None) + def _clean_up_memory_streams(self, request_id: RequestId, streams: _RequestStreams) -> None: + streams[0].close() + streams[1].close() + if self._request_streams.get(request_id) is streams: + del self._request_streams[request_id] async def handle_request(self, scope: Scope, receive: Receive, send: Send) -> None: """Application entry point that handles all HTTP requests.""" @@ -644,34 +627,28 @@ async def _handle_post_request(self, scope: Scope, request: Request, receive: Re request_id = str(message.id) if self.is_json_response_enabled: - self._request_streams[request_id] = anyio.create_memory_object_stream[EventMessage]( - REQUEST_STREAM_BUFFER_SIZE - ) - request_stream_reader = self._request_streams[request_id][1] - # Process the message - metadata = self._message_metadata( - request, on_request_unanswered=partial(self._terminate_unanswered_request, message.id) - ) - session_message = SessionMessage(message, metadata=metadata) - await writer.send(session_message) + request_streams = anyio.create_memory_object_stream[EventMessage](REQUEST_STREAM_BUFFER_SIZE) + self._request_streams[request_id] = request_streams try: - # `message_router` deposits only this request's own response - # here: anything else scoped to the request has no wire in - # JSON-response mode. - event_message = await request_stream_reader.receive() - except (anyio.EndOfStream, anyio.ClosedResourceError): - # The stream closed with no response: the session was - # terminated while this request was in flight. - logger.debug(f"Session terminated with request {request_id} in flight; no response to send") - response = self._create_error_response( - "Session terminated before the request completed", - HTTPStatus.INTERNAL_SERVER_ERROR, - INTERNAL_ERROR, + metadata = self._message_metadata( + request, on_request_unanswered=partial(self._terminate_unanswered_request, message.id) ) - else: - response = self._create_json_response(event_message.message) + session_message = SessionMessage(message, metadata=metadata) + await writer.send(session_message) + try: + # JSON mode routes only this request's response to its stream. + event_message = await request_streams[1].receive() + except (anyio.EndOfStream, anyio.ClosedResourceError): + logger.debug(f"Session terminated with request {request_id} in flight; no response to send") + response = self._create_error_response( + "Session terminated before the request completed", + HTTPStatus.INTERNAL_SERVER_ERROR, + INTERNAL_ERROR, + ) + else: + response = self._create_json_response(event_message.message) finally: - await self._clean_up_memory_streams(request_id) + self._clean_up_memory_streams(request_id, request_streams) await response(scope, receive, send) else: # Mint the priming event before any per-request state exists: @@ -681,40 +658,38 @@ async def _handle_post_request(self, scope: Scope, request: Request, receive: Re priming_event = await self._mint_priming_event(request_id, protocol_version) sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[SSEEvent](0) - self._sse_stream_writers[request_id] = sse_stream_writer - self._request_streams[request_id] = anyio.create_memory_object_stream[EventMessage]( - REQUEST_STREAM_BUFFER_SIZE - ) - request_stream_reader = self._request_streams[request_id][1] - - headers = { - "Cache-Control": "no-cache, no-transform", - "Connection": "keep-alive", - "Content-Type": CONTENT_TYPE_SSE, - **({MCP_SESSION_ID_HEADER: self.mcp_session_id} if self.mcp_session_id else {}), - } - response = EventSourceResponse( - content=sse_stream_reader, - data_sender_callable=partial( - self._run_sse_writer, request_id, sse_stream_writer, request_stream_reader, priming_event - ), - headers=headers, - ) - - # Start the SSE response (this will send headers immediately) + request_streams = anyio.create_memory_object_stream[EventMessage](REQUEST_STREAM_BUFFER_SIZE) try: - # First send the response to establish the SSE connection - async with anyio.create_task_group() as tg: - tg.start_soon(response, scope, receive, send) - # Then send the message to be processed by the server - session_message = self._create_session_message(message, request, request_id, protocol_version) - await writer.send(session_message) - except Exception: # pragma: lax no cover - logger.exception("SSE response error") - await sse_stream_writer.aclose() - await self._clean_up_memory_streams(request_id) + self._sse_stream_writers[request_id] = sse_stream_writer + self._request_streams[request_id] = request_streams + headers = { + "Cache-Control": "no-cache, no-transform", + "Connection": "keep-alive", + "Content-Type": CONTENT_TYPE_SSE, + **({MCP_SESSION_ID_HEADER: self.mcp_session_id} if self.mcp_session_id else {}), + } + response = EventSourceResponse( + content=sse_stream_reader, + data_sender_callable=partial( + self._run_sse_writer, sse_stream_writer, request_streams[1], priming_event + ), + headers=headers, + ) + try: + async with anyio.create_task_group() as tg: + tg.start_soon(response, scope, receive, send) + session_message = self._create_session_message( + message, request, request_id, protocol_version + ) + await writer.send(session_message) + except Exception: # pragma: lax no cover + logger.exception("SSE response error") finally: - await sse_stream_reader.aclose() + sse_stream_writer.close() + sse_stream_reader.close() + if self._sse_stream_writers.get(request_id) is sse_stream_writer: + self._sse_stream_writers.pop(request_id, None) + self._clean_up_memory_streams(request_id, request_streams) except Exception as err: logger.exception("Error handling POST request") @@ -775,19 +750,12 @@ async def _handle_get_request(self, request: Request, send: Send) -> None: await response(request.scope, request.receive, send) return - # Create SSE stream sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[SSEEvent](0) + request_streams = anyio.create_memory_object_stream[EventMessage](REQUEST_STREAM_BUFFER_SIZE) async def standalone_sse_writer(): try: - # Create a standalone message stream for server-initiated messages - - self._request_streams[GET_STREAM_KEY] = anyio.create_memory_object_stream[EventMessage]( - REQUEST_STREAM_BUFFER_SIZE - ) - standalone_stream_reader = self._request_streams[GET_STREAM_KEY][1] - - async with sse_stream_writer, standalone_stream_reader: + async with sse_stream_writer, request_streams[1] as standalone_stream_reader: # Process messages from the standalone stream async for event_message in standalone_stream_reader: # For the standalone stream, we handle: @@ -805,24 +773,22 @@ async def standalone_sse_writer(): logger.exception("Error in standalone SSE writer") # pragma: no cover finally: logger.debug("Closing standalone SSE writer") - await self._clean_up_memory_streams(GET_STREAM_KEY) - - # Create and start EventSourceResponse - response = EventSourceResponse( - content=sse_stream_reader, - data_sender_callable=standalone_sse_writer, - headers=headers, - ) try: - # This will send headers immediately and establish the SSE connection - await response(request.scope, request.receive, send) - except Exception: # pragma: lax no cover - logger.exception("Error in standalone SSE response") - await self._clean_up_memory_streams(GET_STREAM_KEY) + self._request_streams[GET_STREAM_KEY] = request_streams + response = EventSourceResponse( + content=sse_stream_reader, + data_sender_callable=standalone_sse_writer, + headers=headers, + ) + try: + await response(request.scope, request.receive, send) + except Exception: # pragma: lax no cover + logger.exception("Error in standalone SSE response") finally: - await sse_stream_writer.aclose() - await sse_stream_reader.aclose() + sse_stream_writer.close() + sse_stream_reader.close() + self._clean_up_memory_streams(GET_STREAM_KEY, request_streams) async def _handle_delete_request(self, request: Request, send: Send) -> None: """Handle DELETE requests for explicit session termination.""" @@ -856,15 +822,8 @@ async def terminate(self) -> None: self._terminated = True logger.info(f"Terminating session: {self.mcp_session_id}") - # We need a copy of the keys to avoid modification during iteration - request_stream_keys = list(self._request_streams.keys()) - - # Close all request streams asynchronously - for key in request_stream_keys: - await self._clean_up_memory_streams(key) - - # Clear the request streams dictionary immediately - self._request_streams.clear() + for request_id, streams in list(self._request_streams.items()): + self._clean_up_memory_streams(request_id, streams) try: if self._read_stream_writer is not None: # pragma: no branch await self._read_stream_writer.aclose() @@ -956,6 +915,7 @@ 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 try: async with sse_stream_writer: # Define an async callback for sending events @@ -968,11 +928,12 @@ async def send_event(event_message: EventMessage) -> None: stream_id = await event_store.replay_events_after(last_event_id, send_event) if not stream_id or stream_id in self._request_streams: return - self._sse_stream_writers[stream_id] = sse_stream_writer - self._request_streams[stream_id] = anyio.create_memory_object_stream[EventMessage]( + request_streams = anyio.create_memory_object_stream[EventMessage]( REQUEST_STREAM_BUFFER_SIZE ) - msg_reader = self._request_streams[stream_id][1] + self._sse_stream_writers[stream_id] = sse_stream_writer + self._request_streams[stream_id] = request_streams + msg_reader = request_streams[1] 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) @@ -988,7 +949,9 @@ async def send_event(event_message: EventMessage) -> None: finally: if stream_id is not None and self._sse_stream_writers.get(stream_id) is sse_stream_writer: self._sse_stream_writers.pop(stream_id, None) - await self._clean_up_memory_streams(stream_id) + if request_streams is not None: + assert stream_id is not None + self._clean_up_memory_streams(stream_id, request_streams) # Create and start EventSourceResponse response = EventSourceResponse( @@ -1093,9 +1056,8 @@ async def message_router(): try: # Send both the message and the event ID await target[0].send(EventMessage(message, event_id)) - except (anyio.BrokenResourceError, anyio.ClosedResourceError): # pragma: no cover - if self._request_streams.get(request_stream_id) is target: - self._request_streams.pop(request_stream_id, None) + except (anyio.BrokenResourceError, anyio.ClosedResourceError): + self._clean_up_memory_streams(request_stream_id, target) else: logger.debug( f"""Request stream {request_stream_id} not found @@ -1117,9 +1079,8 @@ async def message_router(): # Yield the streams for the caller to use yield read_stream, write_stream finally: - for stream_id in list(self._request_streams.keys()): - await self._clean_up_memory_streams(stream_id) - self._request_streams.clear() + for stream_id, streams in list(self._request_streams.items()): + self._clean_up_memory_streams(stream_id, streams) # Clean up the read and write streams try: diff --git a/src/mcp/shared/_httpx_utils.py b/src/mcp/shared/_httpx_utils.py index 6e5486a1ea..b5b0a6d331 100644 --- a/src/mcp/shared/_httpx_utils.py +++ b/src/mcp/shared/_httpx_utils.py @@ -167,7 +167,7 @@ async def sse_within_origin( @asynccontextmanager async def sse_events(source: httpx2.EventSource) -> AsyncGenerator[AsyncIterator[httpx2.ServerSentEvent]]: - """Close the SSE iterator when its consumer stops before the response ends.""" + """Close the outer EventSource iterator; HTTPX2 owns its nested iterators.""" events = source.__aiter__() try: yield events diff --git a/tests/server/test_streamable_http_router.py b/tests/server/test_streamable_http_router.py index 7f93a89253..30ecd1b883 100644 --- a/tests/server/test_streamable_http_router.py +++ b/tests/server/test_streamable_http_router.py @@ -1,9 +1,11 @@ """Regression coverage for the StreamableHTTP per-session response router.""" +import gc + import anyio import httpx2 import pytest -from mcp_types import JSONRPCMessage, JSONRPCNotification, 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 ( @@ -293,3 +295,123 @@ async def send(message: Message) -> None: assert transport.is_terminated assert sent[0]["type"] == "http.response.start" assert sent[0]["status"] == 200 + + +@pytest.mark.anyio +@pytest.mark.parametrize("json_response", [False, True], ids=["sse", "json"]) +@pytest.mark.parametrize("pause_priming", [False, True], ids=["live", "priming"]) +async def test_closing_a_replay_preserves_a_later_post_reusing_its_request_id( + json_response: bool, pause_priming: bool +) -> None: + """SDK-defined: a replay closes its own streams, not a later POST's streams for the same ID. + + A raw ASGI peer controls the old HTTP connection's lifetime after receiving its completed response. + """ + replay_registered = anyio.Event() + replay_finished = anyio.Event() + post_finished = anyio.Event() + replay_scope = anyio.CancelScope() + request = JSONRPCRequest(jsonrpc="2.0", id="request", method="tools/list") + previous = JSONRPCResponse(jsonrpc="2.0", id="request", result={"value": "previous"}) + expected = JSONRPCResponse(jsonrpc="2.0", id="request", result={"value": "current"}) + replay_wire: list[bytes] = [] + post_wire: list[bytes] = [] + + class RecordingStore(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)) + event_id = str(len(self.events)) + if len(self.events) == 3: + replay_registered.set() + if pause_priming: + await anyio.sleep_forever() + return event_id + + async def replay_events_after(self, last_event_id: EventId, send_callback: EventCallback) -> StreamId | None: + assert last_event_id == "1" + await send_callback(EventMessage(previous, "2")) + return "request" + + store = RecordingStore() + cursor = await store.store_event("request", None) + await store.store_event("request", previous) + transport = StreamableHTTPServerTransport( + mcp_session_id=None, event_store=store, is_json_response_enabled=json_response + ) + scope: Scope = { + "type": "http", + "method": "GET", + "path": "/mcp", + "query_string": b"", + "headers": [ + (b"accept", b"application/json, text/event-stream"), + (b"content-type", b"application/json"), + (b"mcp-protocol-version", b"2025-11-25"), + ], + } + + async def get_receive() -> Message: + await anyio.sleep_forever() + raise NotImplementedError + + async def get_send(message: Message) -> None: + if message["type"] == "http.response.body": + replay_wire.append(message.get("body", b"")) + + async def get() -> None: + with replay_scope: + await transport.handle_request( + scope | {"headers": [*scope["headers"], (b"last-event-id", cursor.encode())]}, get_receive, get_send + ) + replay_finished.set() + + async def post_send(message: Message) -> None: + if message["type"] == "http.response.body": + post_wire.append(message.get("body", b"")) + + outgoing, incoming = anyio.create_memory_object_stream[Message](1) + + async def post() -> None: + await transport.handle_request(scope | {"method": "POST"}, incoming.receive, post_send) + post_finished.set() + + with anyio.fail_after(5): + async with ( + outgoing, + incoming, + transport.connect() as (read_stream, write_stream), + anyio.create_task_group() as tg, + ): + tg.start_soon(get) + await replay_registered.wait() + await anyio.wait_all_tasks_blocked() + assert previous.model_dump_json(by_alias=True, exclude_unset=True).encode() in b"".join(replay_wire) + await outgoing.send( + {"type": "http.request", "body": request.model_dump_json(by_alias=True, exclude_none=True).encode()} + ) + tg.start_soon(post) + forwarded = await read_stream.receive() + assert isinstance(forwarded, SessionMessage) + assert forwarded.message == request + replay_scope.cancel() + await replay_finished.wait() + gc.collect() + await write_stream.send(SessionMessage(expected)) + await post_finished.wait() + + if json_response: + assert jsonrpc_message_adapter.validate_json(b"".join(post_wire)) == expected + else: + received = httpx2.Response( + 200, + headers={"content-type": "text/event-stream"}, + content=b"".join(post_wire), + request=httpx2.Request("POST", "http://localhost/mcp"), + ) + assert [ + jsonrpc_message_adapter.validate_json(event.data) for event in httpx2.EventSource(received) if event.data + ] == [expected] + gc.collect() diff --git a/tests/shared/test_streamable_http.py b/tests/shared/test_streamable_http.py index 655d9941dc..08b4e3c541 100644 --- a/tests/shared/test_streamable_http.py +++ b/tests/shared/test_streamable_http.py @@ -21,7 +21,6 @@ import httpx2 import mcp_types as types import pytest -from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream from httpx2 import ServerSentEvent from mcp_types import ( DEFAULT_NEGOTIATED_VERSION, @@ -50,7 +49,6 @@ from mcp.client.streamable_http import StreamableHTTPTransport, streamable_http_client from mcp.server import Server, ServerRequestContext from mcp.server.streamable_http import ( - GET_STREAM_KEY, MCP_PROTOCOL_VERSION_HEADER, MCP_SESSION_ID_HEADER, SESSION_ID_PATTERN, @@ -2279,7 +2277,7 @@ async def message_handler(message: IncomingMessage) -> None: await notified.wait() # Tear the standalone stream down while the writer is parked on it. (transport,) = session_manager._server_instances.values() # pyright: ignore[reportPrivateUsage] - await transport._clean_up_memory_streams(GET_STREAM_KEY) # pyright: ignore[reportPrivateUsage] + transport.close_standalone_sse_stream() assert "Error in standalone SSE writer" not in caplog.text @@ -2297,38 +2295,19 @@ async def test_standalone_stream_teardown_between_dequeues_is_not_an_error( mcp_session_id=None, security_settings=TransportSecuritySettings(enable_dns_rebinding_protection=False), ) - # The GET handler only checks that a read-stream writer exists; it is never written to. - read_stream_writer, read_stream = create_context_streams[SessionMessage | Exception](0) - transport._read_stream_writer = read_stream_writer # pyright: ignore[reportPrivateUsage] - - stream_registered = anyio.Event() - - class SignalingStreams( - dict[types.RequestId, tuple[MemoryObjectSendStream[EventMessage], MemoryObjectReceiveStream[EventMessage]]] - ): - # Only the GET handler inserts here, so any insert is the standalone stream registration. - def __setitem__( - self, - key: types.RequestId, - value: tuple[MemoryObjectSendStream[EventMessage], MemoryObjectReceiveStream[EventMessage]], - ) -> None: - super().__setitem__(key, value) - stream_registered.set() - - transport._request_streams = SignalingStreams() # pyright: ignore[reportPrivateUsage] - + headers_sent = anyio.Event() gate = anyio.Event() sent: list[Message] = [] async def asgi_send(message: Message) -> None: sent.append(message) + if message["type"] == "http.response.start": + headers_sent.set() await gate.wait() - # Never delivers anything, parking the response's disconnect listener. - disconnect_send, disconnect_receive = anyio.create_memory_object_stream[Message](0) - async def asgi_receive() -> Message: - return await disconnect_receive.receive() + await anyio.sleep_forever() + raise NotImplementedError scope: Scope = { "type": "http", @@ -2339,18 +2318,24 @@ async def asgi_receive() -> Message: } notification = types.JSONRPCNotification(jsonrpc="2.0", method="notifications/initialized") - async with read_stream_writer, read_stream, disconnect_send, disconnect_receive: - with anyio.fail_after(5): - async with anyio.create_task_group() as tg: # pragma: no branch + with anyio.fail_after(5): + async with transport.connect() as (_, write_stream): + + async def send_notifications() -> None: + while True: + await write_stream.send(SessionMessage(notification)) + + async with anyio.create_task_group() as tg: tg.start_soon(transport.handle_request, scope, asgi_receive, asgi_send) - await stream_registered.wait() - standalone_send = transport._request_streams[GET_STREAM_KEY][0] # pyright: ignore[reportPrivateUsage] - # Zero-buffer rendezvous: once send() returns, the writer has dequeued the event - # and is blocked forwarding it past the closed gate — the between-dequeues window. - await standalone_send.send(EventMessage(notification)) - await transport._clean_up_memory_streams(GET_STREAM_KEY) # pyright: ignore[reportPrivateUsage] - # Unblock the response; the writer's next dequeue hits its closed stream. + await headers_sent.wait() + async with anyio.create_task_group() as senders: + senders.start_soon(send_notifications) + # Fill the route until the router and SSE writer are both blocked. + await anyio.wait_all_tasks_blocked() + transport.close_standalone_sse_stream() + senders.cancel_scope.cancel() gate.set() + await transport.terminate() assert sent[0]["type"] == "http.response.start" assert sent[0]["status"] == 200 @@ -2359,3 +2344,4 @@ async def asgi_receive() -> Message: assert body_chunks[-1] == {"type": "http.response.body", "body": b"", "more_body": False} assert "Error in standalone SSE writer" not in caplog.text assert "Error in standalone SSE response" not in caplog.text + assert "Error in message router" not in caplog.text From 64554aa2930fe334ee2a16074f519e9d4938a93c Mon Sep 17 00:00:00 2001 From: Marcelo Trylesinski Date: Wed, 16 Sep 2026 14:13:08 +0200 Subject: [PATCH 5/7] Accept class-based SSE iterators during cleanup --- src/mcp/shared/_httpx_utils.py | 14 ++++-- tests/client/test_streamable_http.py | 73 ++++++++++++++++++++++++++++ 2 files changed, 83 insertions(+), 4 deletions(-) diff --git a/src/mcp/shared/_httpx_utils.py b/src/mcp/shared/_httpx_utils.py index b5b0a6d331..1800a61154 100644 --- a/src/mcp/shared/_httpx_utils.py +++ b/src/mcp/shared/_httpx_utils.py @@ -3,9 +3,10 @@ from abc import ABC, abstractmethod from collections.abc import AsyncGenerator, AsyncIterator from contextlib import asynccontextmanager -from typing import Any, Protocol +from typing import Any import httpx2 +from typing_extensions import Protocol, runtime_checkable __all__ = ["create_mcp_http_client", "MCP_DEFAULT_TIMEOUT", "MCP_DEFAULT_SSE_READ_TIMEOUT"] @@ -165,15 +166,20 @@ async def sse_within_origin( yield httpx2.EventSource(response) +@runtime_checkable +class _AsyncClosable(Protocol): + async def aclose(self) -> None: ... + + @asynccontextmanager async def sse_events(source: httpx2.EventSource) -> AsyncGenerator[AsyncIterator[httpx2.ServerSentEvent]]: - """Close the outer EventSource iterator; HTTPX2 owns its nested iterators.""" + """Close the outer EventSource iterator if supported; HTTPX2 owns its nested iterators.""" events = source.__aiter__() try: yield events finally: - assert isinstance(events, AsyncGenerator) - await events.aclose() + if isinstance(events, _AsyncClosable): + await events.aclose() def redirect_location(response: httpx2.Response) -> httpx2.URL | None: diff --git a/tests/client/test_streamable_http.py b/tests/client/test_streamable_http.py index c6e62ad94a..77e091cafb 100644 --- a/tests/client/test_streamable_http.py +++ b/tests/client/test_streamable_http.py @@ -917,6 +917,79 @@ def handler(request: httpx2.Request) -> httpx2.Response: assert seen == [("GET http://test/mcp", "evt-41")] +@pytest.mark.anyio +@pytest.mark.parametrize("iterator_kind", ["generator", "closable", "plain"]) +async def test_resumed_response_accepts_async_iterators_and_closes_them_when_supported( + monkeypatch: pytest.MonkeyPatch, iterator_kind: str +) -> None: + """SDK-defined: resumption accepts any EventSource async iterator and closes it when supported. + + Substitute the public iterator boundary to isolate representation from HTTPX2's nested-generator cleanup. + """ + expected = JSONRPCResponse(jsonrpc="2.0", id="resume-1", result={"ok": True}) + event = httpx2.ServerSentEvent(data=expected.model_dump_json(by_alias=True)) + body = f"data: {event.data}\n\n" + closed: list[bool] = [] + + async def generate() -> AsyncIterator[httpx2.ServerSentEvent]: + try: + yield event + finally: + closed.append(True) + + class EventIterator: + def __aiter__(self) -> AsyncIterator[httpx2.ServerSentEvent]: + return self + + async def __anext__(self) -> httpx2.ServerSentEvent: + return event + + class ClosingEventIterator: + def __aiter__(self) -> AsyncIterator[httpx2.ServerSentEvent]: + return self + + async def __anext__(self) -> httpx2.ServerSentEvent: + return event + + async def aclose(self) -> None: + closed.append(True) + + iterators: dict[str, AsyncIterator[httpx2.ServerSentEvent]] = { + "generator": generate(), + "closable": ClosingEventIterator(), + "plain": EventIterator(), + } + + def iterate(source: httpx2.EventSource) -> AsyncIterator[httpx2.ServerSentEvent]: + assert source.response.text == body + return iterators[iterator_kind] + + monkeypatch.setattr(httpx2.EventSource, "__aiter__", iterate) + token = "evt-41" + + def handler(request: httpx2.Request) -> httpx2.Response: + assert request.method == "GET" + assert request.headers["last-event-id"] == token + return httpx2.Response(200, headers={"content-type": "text/event-stream"}, text=body) + + with anyio.fail_after(5): + async with ( + httpx2.AsyncClient(transport=httpx2.MockTransport(handler)) as http, + streamable_http_client("http://test/mcp", http_client=http) as (read, write), + ): + await write.send( + SessionMessage( + JSONRPCRequest(jsonrpc="2.0", id=expected.id, method="tools/call", params={}), + metadata=ClientMessageMetadata(resumption_token=token), + ) + ) + reply = await read.receive() + + assert isinstance(reply, SessionMessage) + assert reply.message == expected + assert closed == ([] if iterator_kind == "plain" else [True]) + + async def _redirected_call_error(url: str, location: str) -> str: """Send one request through streamable_http_client to a server answering `url` with a 307 to `location`, and return the message of the error that resolves it.""" From 8d8614ad3601ab55507b1d5495d236c2d83ba44e Mon Sep 17 00:00:00 2001 From: Marcelo Trylesinski Date: Wed, 16 Sep 2026 14:13:08 +0200 Subject: [PATCH 6/7] Avoid waiting for remote subscription cancellation writes --- docs/client/subscriptions.md | 4 +- src/mcp/client/subscriptions.py | 13 +++--- tests/client/test_subscriptions.py | 67 ++++++++++++++++++++++++++++++ 3 files changed, 77 insertions(+), 7 deletions(-) diff --git a/docs/client/subscriptions.md b/docs/client/subscriptions.md index dddb52abf2..90478bc32a 100644 --- a/docs/client/subscriptions.md +++ b/docs/client/subscriptions.md @@ -59,7 +59,9 @@ Requests run freely beside an open stream, from the watcher task or any other, o To stop watching, leave the block: there is no `unsubscribe` call. Cancelling the task that owns the block does that for you, and the SDK cancels the listen request the way the transport expects: over streamable HTTP, by closing that request's stream. A watcher that runs for the life of your app never returns on its own, so cancel it, or its task group's scope, at shutdown. -Exit waits up to five seconds for the SDK's listen task to finish, even if the caller is cancelled. With an in-memory server, this lets cooperative handler cleanup finish before you open another subscription. The limit prevents an uncooperative handler from blocking subscription exit indefinitely. It is not an acknowledgment that a remote server has finished its cleanup. +With a direct in-memory connection (`Client(server)` or a `ClientSession` using `DirectDispatcher`), exit waits up to five seconds for the listen task to finish, even if the caller is cancelled. This lets cooperative handler cleanup release its subscription slot before you open another subscription. The limit prevents an uncooperative handler from blocking exit indefinitely. + +With a stream-backed connection, exit cancels the listen task without waiting for its courtesy cancellation write. The session still owns that task and its cleanup. A slow transport does not delay each subscription's exit, and exit does not acknowledge that the remote server has finished its cleanup. ## Streams end diff --git a/src/mcp/client/subscriptions.py b/src/mcp/client/subscriptions.py index f8a1af262c..a6ba307223 100644 --- a/src/mcp/client/subscriptions.py +++ b/src/mcp/client/subscriptions.py @@ -18,6 +18,7 @@ import mcp_types as types from mcp_types.version import MODERN_PROTOCOL_VERSIONS +from mcp.shared.direct_dispatcher import DirectDispatcher from mcp.shared.dispatcher import CallOptions from mcp.shared.exceptions import MCPError from mcp.shared.subscriptions import ( @@ -241,6 +242,7 @@ async def listen( data = request.model_dump(by_alias=True, mode="json", exclude_none=True) opts: CallOptions = {"request_id": request_id} session._stamp(data, opts) # pyright: ignore[reportPrivateUsage] + dispatcher = session._dispatcher # pyright: ignore[reportPrivateUsage] driver_scope = anyio.CancelScope() driver_done = anyio.Event() @@ -249,9 +251,7 @@ async def drive() -> None: try: with driver_scope: try: - await session._dispatcher.send_raw_request( # pyright: ignore[reportPrivateUsage] - data["method"], data.get("params"), opts - ) + await dispatcher.send_raw_request(data["method"], data.get("params"), opts) except MCPError as error: route.settle("lost", error=error) return @@ -284,8 +284,9 @@ async def drive() -> None: finally: route.settle("local") driver_scope.cancel() - # Direct handlers unwind in the driver; remote cancellation has no acknowledgment. - with anyio.move_on_after(5, shield=True): - await driver_done.wait() + # Only direct drivers own handler cleanup; remote courtesy writes remain session-owned. + if isinstance(dispatcher, DirectDispatcher): + with anyio.move_on_after(5, shield=True): + await driver_done.wait() finally: session._unregister_listen_route(request_id) # pyright: ignore[reportPrivateUsage] diff --git a/tests/client/test_subscriptions.py b/tests/client/test_subscriptions.py index d862561850..aa71a52639 100644 --- a/tests/client/test_subscriptions.py +++ b/tests/client/test_subscriptions.py @@ -35,6 +35,8 @@ ) from mcp.shared.direct_dispatcher import create_direct_dispatcher_pair from mcp.shared.dispatcher import CallOptions +from mcp.shared.jsonrpc_dispatcher import JSONRPCDispatcher +from mcp.shared.message import SessionMessage pytestmark = pytest.mark.anyio @@ -322,6 +324,71 @@ async def shielded_cleanup( await cleanup_finished.wait() +@pytest.mark.parametrize( + "anyio_backend", + [pytest.param(("trio", {"clock": MockClock(autojump_threshold=0)}), id="trio-mockclock")], +) +@pytest.mark.parametrize("cancelled", [False, True], ids=["normal-exit", "cancelled-exit"]) +async def test_sequential_remote_subscription_exits_do_not_wait_for_courtesy_writes( + monkeypatch: pytest.MonkeyPatch, cancelled: bool +) -> None: + """SDK-defined: remote exits leave courtesy writes to session-owned drivers, even when cancelled. + + Block the public stream `send` boundary; typed server handlers cannot wedge a client's transport write. + """ + server = Server("subs", on_subscriptions_listen=ListenHandler(InMemorySubscriptionBus())) + client_write, server_read = anyio.create_memory_object_stream[SessionMessage | Exception]() + server_write, client_read = anyio.create_memory_object_stream[SessionMessage | Exception]() + release_writes = anyio.Event() + attempted: list[types.RequestId] = [] + delivered: list[types.RequestId] = [] + send = client_write.send + + async def block_courtesy_write(item: SessionMessage | Exception) -> None: + assert isinstance(item, SessionMessage) + message = item.message + if isinstance(message, types.JSONRPCNotification) and message.method == "notifications/cancelled": + assert message.params is not None + request_id = message.params["requestId"] + attempted.append(request_id) + with anyio.fail_after(5): + await release_writes.wait() + await send(item) + delivered.append(request_id) + else: + await send(item) + + monkeypatch.setattr(client_write, "send", block_courtesy_write) + dispatcher = JSONRPCDispatcher(client_read, client_write) + with anyio.fail_after(5): + async with client_write, server_read, server_write, client_read, anyio.create_task_group() as tg: + tg.start_soon(server.run, server_read, server_write, server.create_initialization_options()) + async with ClientSession(dispatcher=dispatcher) as session: + await session.discover() + subscription_ids: list[types.RequestId] = [] + started = anyio.current_time() + try: + for _ in range(2): + subscription = listen(session, tools_list_changed=True) + sub = await subscription.__aenter__() + subscription_ids.append(sub.subscription_id) + with anyio.CancelScope() as scope: + if cancelled: + scope.cancel() + await subscription.__aexit__(None, None, None) + assert anyio.current_time() == started # MockClock time, never wall-clock time. + with pytest.raises(StopAsyncIteration): + await anext(sub) + await anyio.wait_all_tasks_blocked() + assert attempted == subscription_ids + assert delivered == [] + finally: + release_writes.set() + await anyio.wait_all_tasks_blocked() + assert set(delivered) == set(subscription_ids) + tg.cancel_scope.cancel() + + async def test_concurrent_subscriptions_demux_independently(): """Two open subscriptions each receive only their own filter's events.""" bus = InMemorySubscriptionBus() From d0608103213061f348cacd19e77e585d727f2197 Mon Sep 17 00:00:00 2001 From: Marcelo Trylesinski Date: Wed, 16 Sep 2026 14:13:08 +0200 Subject: [PATCH 7/7] Keep replay network delivery outside the event-store lock --- docs/run/legacy-clients.md | 12 +- src/mcp/server/streamable_http.py | 73 +++-- tests/server/test_streamable_http_router.py | 310 +++++++++++++++++++- 3 files changed, 360 insertions(+), 35 deletions(-) diff --git a/docs/run/legacy-clients.md b/docs/run/legacy-clients.md index 80773e7270..a5532f3e5d 100644 --- a/docs/run/legacy-clients.md +++ b/docs/run/legacy-clients.md @@ -66,11 +66,13 @@ 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 "Slow replay delays messages in the same session" - With `event_store=`, the session's message router waits until replay finishes and the live - stream is registered. This prevents newly produced responses from being lost between replay - and live delivery. It does not serialize incoming POSTs or reserve JSON-RPC request IDs. - A slow replay can delay other requests in that session, but not other sessions. +!!! 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. ## Session lifetime and limits diff --git a/src/mcp/server/streamable_http.py b/src/mcp/server/streamable_http.py index 1f1547b129..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 @@ -216,6 +217,7 @@ def __init__( 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 @@ -916,29 +918,41 @@ 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: - # Define an async callback for sending events + async def send_event(event_message: EventMessage) -> None: event_data = self._create_event_data(event_message) - await sse_stream_writer.send(event_data) + # 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") + ) - # ponytail: replay stalls this session's router; per-stream cursors would narrow the lock. 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 not stream_id or stream_id in self._request_streams: + if self._terminated or not self._connected: return - 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 - msg_reader = request_streams[1] - 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 msg_reader: + 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 @@ -952,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( @@ -1006,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 @@ -1045,10 +1063,25 @@ async def message_router(): # messages will be replayed on the re-connect event_id = None if self._event_store: - async with self._event_store_lock: + 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) @@ -1076,9 +1109,10 @@ async def message_router(): tg.start_soon(message_router) try: - # Yield the streams for the caller to use + 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) @@ -1091,5 +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}") - if self._event_store is not None: - tg.cancel_scope.cancel() + # 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 30ecd1b883..fdebf32e3e 100644 --- a/tests/server/test_streamable_http_router.py +++ b/tests/server/test_streamable_http_router.py @@ -251,9 +251,13 @@ async def send(message: Message) -> None: @pytest.mark.anyio -async def test_transport_shutdown_cancels_a_router_waiting_for_replay() -> None: - """SDK-defined: a stalled replay cannot keep the session's response router alive after termination.""" +@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): @@ -262,7 +266,8 @@ async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) async def replay_events_after(self, last_event_id: EventId, send_callback: EventCallback) -> StreamId | None: replay_started.set() - return await anyio.sleep_forever() + await release_replay.wait() + return "request" store = BlockingStore() cursor = await store.store_event("request", None) @@ -272,7 +277,11 @@ async def replay_events_after(self, last_event_id: EventId, send_callback: Event "method": "GET", "path": "/mcp", "query_string": b"", - "headers": [(b"accept", b"text/event-stream"), (b"last-event-id", cursor.encode())], + "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: @@ -282,19 +291,297 @@ async def receive() -> Message: 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(transport.handle_request, scope, receive, send) + requests.start_soon(replay) await replay_started.wait() - await write_stream.send(SessionMessage(JSONRPCResponse(jsonrpc="2.0", id="request", result={}))) await anyio.wait_all_tasks_blocked() - await transport.terminate() - requests.cancel_scope.cancel() + 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 + 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 @@ -305,7 +592,7 @@ async def test_closing_a_replay_preserves_a_later_post_reusing_its_request_id( ) -> None: """SDK-defined: a replay closes its own streams, not a later POST's streams for the same ID. - A raw ASGI peer controls the old HTTP connection's lifetime after receiving its completed response. + A raw ASGI peer controls the old HTTP connection's lifetime after the original request has completed. """ replay_registered = anyio.Event() replay_finished = anyio.Event() @@ -388,7 +675,8 @@ async def post() -> None: tg.start_soon(get) await replay_registered.wait() await anyio.wait_all_tasks_blocked() - assert previous.model_dump_json(by_alias=True, exclude_unset=True).encode() in b"".join(replay_wire) + if not pause_priming: + assert previous.model_dump_json(by_alias=True, exclude_unset=True).encode() in b"".join(replay_wire) await outgoing.send( {"type": "http.request", "body": request.model_dump_json(by_alias=True, exclude_none=True).encode()} )