diff --git a/src/mcp/server/mcpserver/server.py b/src/mcp/server/mcpserver/server.py index fbd2c26dd8..370eea6d89 100644 --- a/src/mcp/server/mcpserver/server.py +++ b/src/mcp/server/mcpserver/server.py @@ -1273,8 +1273,20 @@ async def sse_endpoint(request: Request) -> Response: # pragma: no cover # mount these routes last, so they have the lowest route matching precedence routes.extend(self._custom_starlette_routes) + @asynccontextmanager + async def sse_lifespan(_app: Starlette): + try: + yield + finally: + await sse.close() + # Create Starlette app with routes and middleware - return Starlette(debug=self.settings.debug, routes=routes, middleware=middleware) + return Starlette( + debug=self.settings.debug, + routes=routes, + middleware=middleware, + lifespan=sse_lifespan, + ) def streamable_http_app( self, diff --git a/src/mcp/server/sse.py b/src/mcp/server/sse.py index d71ef25004..37409bdc76 100644 --- a/src/mcp/server/sse.py +++ b/src/mcp/server/sse.py @@ -130,6 +130,8 @@ def __init__( self._endpoint = endpoint self._read_stream_writers = {} self._session_owners = {} + # SSE body writers; closed on shutdown so EventSourceResponse can finish. + self._sse_stream_writers: dict[UUID, Any] = {} self._security = TransportSecurityMiddleware(security_settings) self._post_message_app = RequestBodyLimitMiddleware(self._handle_post_message, max_request_body_size) logger.debug(f"SseServerTransport initialized with endpoint: {endpoint}") @@ -175,6 +177,7 @@ async def connect_sse(self, scope: Scope, receive: Receive, send: Send): client_post_uri_data = f"{quote(full_message_path_for_client)}?session_id={session_id.hex}" sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[dict[str, Any]](0) + self._sse_stream_writers[session_id] = sse_stream_writer async def sse_writer(): logger.debug("Starting SSE writer") @@ -214,8 +217,26 @@ async def response_wrapper(scope: Scope, receive: Receive, send: Send): yield (read_stream, write_stream) finally: self._read_stream_writers.pop(session_id, None) + self._sse_stream_writers.pop(session_id, None) self._session_owners.pop(session_id, None) + async def close(self) -> None: + """Close all active SSE sessions so the ASGI server can shut down. + + Uvicorn waits for outstanding streaming responses on SIGINT. Closing the + per-session SSE and read streams unblocks EventSourceResponse and the + MCP session task so the process can exit. + """ + session_ids = set(self._read_stream_writers) | set(self._sse_stream_writers) + for session_id in session_ids: + read_writer = self._read_stream_writers.pop(session_id, None) + sse_writer = self._sse_stream_writers.pop(session_id, None) + self._session_owners.pop(session_id, None) + if read_writer is not None: + await read_writer.aclose() + if sse_writer is not None: + await sse_writer.aclose() + async def handle_post_message(self, scope: Scope, receive: Receive, send: Send) -> None: """ASGI application for the message endpoint. diff --git a/tests/shared/test_sse.py b/tests/shared/test_sse.py index 77d1b28a0a..6bfa7b024d 100644 --- a/tests/shared/test_sse.py +++ b/tests/shared/test_sse.py @@ -523,3 +523,42 @@ async def test_sse_session_cleanup_on_disconnect() -> None: headers={"Content-Type": "application/json"}, ) assert response.status_code == 404 + + +@pytest.mark.anyio +async def test_sse_transport_close_unblocks_active_session() -> None: + """Closing the transport ends active SSE streams so the server can shut down.""" + sse = SseServerTransport( + "/messages/", security_settings=TransportSecuritySettings(enable_dns_rebinding_protection=False) + ) + server = Server(SERVER_NAME) + + async def handle_sse(request: Request) -> Response: + async with sse.connect_sse(request.scope, request.receive, request._send) as (read_stream, write_stream): + await server.run(read_stream, write_stream, server.create_initialization_options()) + return Response() + + app = Starlette(routes=[Route("/sse", endpoint=handle_sse), Mount("/messages/", app=sse.handle_post_message)]) + http_client = httpx2.AsyncClient( + transport=StreamingASGITransport(app, cancel_on_close=False), base_url=BASE_URL + ) + + async with http_client: + async with anyio.create_task_group() as tg: + connected = anyio.Event() + + async def hold_sse() -> None: + async with http_client.stream("GET", "/sse") as response: + assert response.status_code == 200 + lines = response.aiter_lines() + assert await anext(lines) == "event: endpoint" + connected.set() + # Stay connected until the transport is closed. + async for _ in lines: + pass + + tg.start_soon(hold_sse) + await connected.wait() + assert sse._sse_stream_writers # noqa: SLF001 + await sse.close() +