From d83afcc9eff8dba1ed94369b0da75fe04f5dcd65 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Fri, 9 May 2025 16:07:25 -0400 Subject: [PATCH] Ensure we handle nested SSE apps --- src/fastmcp/low_level/README.md | 1 + src/fastmcp/low_level/__init__.py | 0 src/fastmcp/low_level/sse_server_transport.py | 104 ++++++++++++++++++ src/fastmcp/server/http.py | 2 +- tests/client/test_sse.py | 29 +++++ 5 files changed, 135 insertions(+), 1 deletion(-) create mode 100644 src/fastmcp/low_level/README.md create mode 100644 src/fastmcp/low_level/__init__.py create mode 100644 src/fastmcp/low_level/sse_server_transport.py diff --git a/src/fastmcp/low_level/README.md b/src/fastmcp/low_level/README.md new file mode 100644 index 000000000..626d69f9c --- /dev/null +++ b/src/fastmcp/low_level/README.md @@ -0,0 +1 @@ +Patched low-level objects. When possisble, we prefer the official SDK, but we patch bugs here if necessary. \ No newline at end of file diff --git a/src/fastmcp/low_level/__init__.py b/src/fastmcp/low_level/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/src/fastmcp/low_level/sse_server_transport.py b/src/fastmcp/low_level/sse_server_transport.py new file mode 100644 index 000000000..21df959e7 --- /dev/null +++ b/src/fastmcp/low_level/sse_server_transport.py @@ -0,0 +1,104 @@ +import logging +from contextlib import asynccontextmanager +from typing import Any +from urllib.parse import quote +from uuid import uuid4 + +import anyio +from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream +from mcp.server.sse import SseServerTransport as LowLevelSSEServerTransport +from mcp.shared.message import SessionMessage +from sse_starlette import EventSourceResponse +from starlette.types import Receive, Scope, Send + +logger = logging.getLogger(__name__) + + +class SseServerTransport(LowLevelSSEServerTransport): + """ + Patched SSE server transport + """ + + @asynccontextmanager + async def connect_sse(self, scope: Scope, receive: Receive, send: Send): + """ + See https://github.com/modelcontextprotocol/python-sdk/pull/659/ + """ + if scope["type"] != "http": + logger.error("connect_sse received non-HTTP request") + raise ValueError("connect_sse can only handle HTTP requests") + + logger.debug("Setting up SSE connection") + read_stream: MemoryObjectReceiveStream[SessionMessage | Exception] + read_stream_writer: MemoryObjectSendStream[SessionMessage | Exception] + + write_stream: MemoryObjectSendStream[SessionMessage] + write_stream_reader: MemoryObjectReceiveStream[SessionMessage] + + read_stream_writer, read_stream = anyio.create_memory_object_stream(0) + write_stream, write_stream_reader = anyio.create_memory_object_stream(0) + + session_id = uuid4() + self._read_stream_writers[session_id] = read_stream_writer + logger.debug(f"Created new session with ID: {session_id}") + + # Determine the full path for the message endpoint to be sent to the client. + # scope['root_path'] is the prefix where the current Starlette app + # instance is mounted. + # e.g., "" if top-level, or "/api_prefix" if mounted under "/api_prefix". + root_path = scope.get("root_path", "") + + # self._endpoint is the path *within* this app, e.g., "/messages". + # Concatenating them gives the full absolute path from the server root. + # e.g., "" + "/messages" -> "/messages" + # e.g., "/api_prefix" + "/messages" -> "/api_prefix/messages" + full_message_path_for_client = root_path.rstrip("/") + self._endpoint + + # This is the URI (path + query) the client will use to POST messages. + 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) + + async def sse_writer(): + logger.debug("Starting SSE writer") + async with sse_stream_writer, write_stream_reader: + await sse_stream_writer.send( + {"event": "endpoint", "data": client_post_uri_data} + ) + logger.debug(f"Sent endpoint event: {client_post_uri_data}") + + async for session_message in write_stream_reader: + logger.debug(f"Sending message via SSE: {session_message}") + await sse_stream_writer.send( + { + "event": "message", + "data": session_message.message.model_dump_json( + by_alias=True, exclude_none=True + ), + } + ) + + async with anyio.create_task_group() as tg: + + async def response_wrapper(scope: Scope, receive: Receive, send: Send): + """ + The EventSourceResponse returning signals a client close / disconnect. + In this case we close our side of the streams to signal the client that + the connection has been closed. + """ + await EventSourceResponse( + content=sse_stream_reader, data_sender_callable=sse_writer + )(scope, receive, send) + await read_stream_writer.aclose() + await write_stream_reader.aclose() + logging.debug(f"Client session disconnected {session_id}") + + logger.debug("Starting SSE response task") + tg.start_soon(response_wrapper, scope, receive, send) + + logger.debug("Yielding read and write streams") + yield (read_stream, write_stream) diff --git a/src/fastmcp/server/http.py b/src/fastmcp/server/http.py index 385ea5280..547399559 100644 --- a/src/fastmcp/server/http.py +++ b/src/fastmcp/server/http.py @@ -13,7 +13,6 @@ from mcp.server.auth.middleware.bearer_auth import ( from mcp.server.auth.provider import OAuthAuthorizationServerProvider from mcp.server.auth.routes import create_auth_routes from mcp.server.auth.settings import AuthSettings -from mcp.server.sse import SseServerTransport from mcp.server.streamable_http_manager import StreamableHTTPSessionManager from starlette.applications import Starlette from starlette.middleware import Middleware @@ -23,6 +22,7 @@ from starlette.responses import Response from starlette.routing import Mount, Route from starlette.types import Receive, Scope, Send +from fastmcp.low_level.sse_server_transport import SseServerTransport from fastmcp.utilities.logging import get_logger if TYPE_CHECKING: diff --git a/tests/client/test_sse.py b/tests/client/test_sse.py index 58ba87c79..9c6b87e20 100644 --- a/tests/client/test_sse.py +++ b/tests/client/test_sse.py @@ -5,6 +5,8 @@ from collections.abc import Generator import pytest import uvicorn from mcp.types import TextResourceContents +from starlette.applications import Starlette +from starlette.routing import Mount from fastmcp.client import Client from fastmcp.client.transports import SSETransport @@ -90,3 +92,30 @@ async def test_http_headers(sse_server: str): json_result = json.loads(raw_result[0].text) assert "x-demo-header" in json_result assert json_result["x-demo-header"] == "ABC" + + +def run_nested_server(host: str, port: int) -> None: + try: + app = fastmcp_server().sse_app() + mount = Starlette(routes=[Mount("/nest-inner", app=app)]) + mount2 = Starlette(routes=[Mount("/nest-outer", app=mount)]) + server = uvicorn.Server( + config=uvicorn.Config(app=mount2, host=host, port=port, log_level="error") + ) + server.run() + except Exception as e: + print(f"Server error: {e}") + sys.exit(1) + sys.exit(0) + + +async def test_nested_sse_server_resolves_correctly(): + # tests patch for + # https://github.com/modelcontextprotocol/python-sdk/pull/659 + + with run_server_in_process(run_nested_server) as url: + async with Client( + transport=SSETransport(f"{url}/nest-outer/nest-inner/sse") + ) as client: + result = await client.ping() + assert result is True