diff --git a/docs/run/legacy-clients.md b/docs/run/legacy-clients.md index 15809fd608..b19364bc09 100644 --- a/docs/run/legacy-clients.md +++ b/docs/run/legacy-clients.md @@ -66,6 +66,11 @@ 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 "Request cleanup preserves newer streams" + Closing an HTTP request releases its own streams, including during cancellation. + If you reuse a JSON-RPC request ID after the previous request completes, cleanup from an + older connection does not close the newer request's streams. This does not reserve request IDs. + ## 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/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..5f86df316e 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): @@ -211,13 +212,7 @@ def __init__( self._event_store = event_store 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 @@ -355,12 +350,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: @@ -375,8 +369,6 @@ 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) def _create_error_response( self, @@ -457,19 +449,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.""" @@ -642,34 +626,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: @@ -679,40 +657,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") @@ -773,19 +749,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: @@ -803,24 +772,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.""" @@ -854,15 +821,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() @@ -953,49 +913,38 @@ 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 + request_streams: _RequestStreams | None = None 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) + await sse_stream_writer.send(self._create_event_data(event_message)) - # 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) + 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) + if priming_event is not None: + await sse_stream_writer.send(priming_event) + async with request_streams[1] as msg_reader: + async for event_message in msg_reader: + await sse_stream_writer.send(self._create_event_data(event_message)) except anyio.ClosedResourceError: # pragma: lax no cover # Expected when close_sse_stream() is called logger.debug("Replay SSE stream closed by close_sse_stream()") 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) + 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( @@ -1092,13 +1041,13 @@ async def message_router(): event_id = await self._event_store.store_event(request_stream_id, message) logger.debug(f"Stored {event_id} from {request_stream_id}") - if request_stream_id in self._request_streams: + target = self._request_streams.get(request_stream_id) + 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)) - except (anyio.BrokenResourceError, anyio.ClosedResourceError): # pragma: no cover - # Stream might be closed, remove from registry - self._request_streams.pop(request_stream_id, None) + await target[0].send(EventMessage(message, event_id)) + 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,12 +1066,10 @@ async def message_router(): tg.start_soon(message_router) try: - # 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/tests/server/test_streamable_http_router.py b/tests/server/test_streamable_http_router.py index 0c5796c1f5..bd10a2b647 100644 --- a/tests/server/test_streamable_http_router.py +++ b/tests/server/test_streamable_http_router.py @@ -1,8 +1,11 @@ """Regression coverage for the StreamableHTTP per-session response router.""" +import gc + import anyio +import httpx2 import pytest -from mcp_types import JSONRPCMessage, JSONRPCResponse +from mcp_types import JSONRPCMessage, JSONRPCRequest, JSONRPCResponse, jsonrpc_message_adapter from starlette.types import Message, Scope from mcp.server.streamable_http import ( @@ -17,6 +20,11 @@ from mcp.shared.message import SessionMessage +@pytest.fixture(scope="module", params=["asyncio", "trio"]) +def anyio_backend(request: pytest.FixtureRequest) -> str: + return request.param + + class _PrimingFailingStore(EventStore): async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId: raise RuntimeError("backend unavailable") @@ -157,3 +165,124 @@ 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 +@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 the original request has completed. + """ + 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() + 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()} + ) + 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