Skip to content

Commit 69a7593

Browse files
committed
Close active SSE streams on server shutdown.
Track per-session SSE writers and close them from the SSE app lifespan so SIGINT after traffic can finish EventSourceResponse tasks and exit cleanly.
1 parent 6affe5c commit 69a7593

3 files changed

Lines changed: 73 additions & 1 deletion

File tree

‎src/mcp/server/mcpserver/server.py‎

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1273,8 +1273,20 @@ async def sse_endpoint(request: Request) -> Response: # pragma: no cover
12731273
# mount these routes last, so they have the lowest route matching precedence
12741274
routes.extend(self._custom_starlette_routes)
12751275

1276+
@asynccontextmanager
1277+
async def sse_lifespan(_app: Starlette):
1278+
try:
1279+
yield
1280+
finally:
1281+
await sse.close()
1282+
12761283
# Create Starlette app with routes and middleware
1277-
return Starlette(debug=self.settings.debug, routes=routes, middleware=middleware)
1284+
return Starlette(
1285+
debug=self.settings.debug,
1286+
routes=routes,
1287+
middleware=middleware,
1288+
lifespan=sse_lifespan,
1289+
)
12781290

12791291
def streamable_http_app(
12801292
self,

‎src/mcp/server/sse.py‎

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -130,6 +130,8 @@ def __init__(
130130
self._endpoint = endpoint
131131
self._read_stream_writers = {}
132132
self._session_owners = {}
133+
# SSE body writers; closed on shutdown so EventSourceResponse can finish.
134+
self._sse_stream_writers: dict[UUID, Any] = {}
133135
self._security = TransportSecurityMiddleware(security_settings)
134136
self._post_message_app = RequestBodyLimitMiddleware(self._handle_post_message, max_request_body_size)
135137
logger.debug(f"SseServerTransport initialized with endpoint: {endpoint}")
@@ -175,6 +177,7 @@ async def connect_sse(self, scope: Scope, receive: Receive, send: Send):
175177
client_post_uri_data = f"{quote(full_message_path_for_client)}?session_id={session_id.hex}"
176178

177179
sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[dict[str, Any]](0)
180+
self._sse_stream_writers[session_id] = sse_stream_writer
178181

179182
async def sse_writer():
180183
logger.debug("Starting SSE writer")
@@ -214,8 +217,26 @@ async def response_wrapper(scope: Scope, receive: Receive, send: Send):
214217
yield (read_stream, write_stream)
215218
finally:
216219
self._read_stream_writers.pop(session_id, None)
220+
self._sse_stream_writers.pop(session_id, None)
217221
self._session_owners.pop(session_id, None)
218222

223+
async def close(self) -> None:
224+
"""Close all active SSE sessions so the ASGI server can shut down.
225+
226+
Uvicorn waits for outstanding streaming responses on SIGINT. Closing the
227+
per-session SSE and read streams unblocks EventSourceResponse and the
228+
MCP session task so the process can exit.
229+
"""
230+
session_ids = set(self._read_stream_writers) | set(self._sse_stream_writers)
231+
for session_id in session_ids:
232+
read_writer = self._read_stream_writers.pop(session_id, None)
233+
sse_writer = self._sse_stream_writers.pop(session_id, None)
234+
self._session_owners.pop(session_id, None)
235+
if read_writer is not None:
236+
await read_writer.aclose()
237+
if sse_writer is not None:
238+
await sse_writer.aclose()
239+
219240
async def handle_post_message(self, scope: Scope, receive: Receive, send: Send) -> None:
220241
"""ASGI application for the message endpoint.
221242

‎tests/shared/test_sse.py‎

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -523,3 +523,42 @@ async def test_sse_session_cleanup_on_disconnect() -> None:
523523
headers={"Content-Type": "application/json"},
524524
)
525525
assert response.status_code == 404
526+
527+
528+
@pytest.mark.anyio
529+
async def test_sse_transport_close_unblocks_active_session() -> None:
530+
"""Closing the transport ends active SSE streams so the server can shut down."""
531+
sse = SseServerTransport(
532+
"/messages/", security_settings=TransportSecuritySettings(enable_dns_rebinding_protection=False)
533+
)
534+
server = Server(SERVER_NAME)
535+
536+
async def handle_sse(request: Request) -> Response:
537+
async with sse.connect_sse(request.scope, request.receive, request._send) as (read_stream, write_stream):
538+
await server.run(read_stream, write_stream, server.create_initialization_options())
539+
return Response()
540+
541+
app = Starlette(routes=[Route("/sse", endpoint=handle_sse), Mount("/messages/", app=sse.handle_post_message)])
542+
http_client = httpx2.AsyncClient(
543+
transport=StreamingASGITransport(app, cancel_on_close=False), base_url=BASE_URL
544+
)
545+
546+
async with http_client:
547+
async with anyio.create_task_group() as tg:
548+
connected = anyio.Event()
549+
550+
async def hold_sse() -> None:
551+
async with http_client.stream("GET", "/sse") as response:
552+
assert response.status_code == 200
553+
lines = response.aiter_lines()
554+
assert await anext(lines) == "event: endpoint"
555+
connected.set()
556+
# Stay connected until the transport is closed.
557+
async for _ in lines:
558+
pass
559+
560+
tg.start_soon(hold_sse)
561+
await connected.wait()
562+
assert sse._sse_stream_writers # noqa: SLF001
563+
await sse.close()
564+

0 commit comments

Comments
 (0)