Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 13 additions & 1 deletion src/mcp/server/mcpserver/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
21 changes: 21 additions & 0 deletions src/mcp/server/sse.py
Original file line number Diff line number Diff line change
Expand Up @@ -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}")
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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.

Expand Down
39 changes: 39 additions & 0 deletions tests/shared/test_sse.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Loading