Spaces:
Running
Running
Jeremiah Lowin commited on
Commit ·
f827f24
1
Parent(s): b0ed233
Ensure we handle nested SSE apps
Browse files
src/fastmcp/low_level/README.md
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
Patched low-level objects. When possisble, we prefer the official SDK, but we patch bugs here if necessary.
|
src/fastmcp/low_level/__init__.py
ADDED
|
File without changes
|
src/fastmcp/low_level/sse_server_transport.py
ADDED
|
@@ -0,0 +1,104 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import logging
|
| 2 |
+
from contextlib import asynccontextmanager
|
| 3 |
+
from typing import Any
|
| 4 |
+
from urllib.parse import quote
|
| 5 |
+
from uuid import uuid4
|
| 6 |
+
|
| 7 |
+
import anyio
|
| 8 |
+
from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream
|
| 9 |
+
from mcp.server.sse import SseServerTransport as LowLevelSSEServerTransport
|
| 10 |
+
from mcp.shared.message import SessionMessage
|
| 11 |
+
from sse_starlette import EventSourceResponse
|
| 12 |
+
from starlette.types import Receive, Scope, Send
|
| 13 |
+
|
| 14 |
+
logger = logging.getLogger(__name__)
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
class SseServerTransport(LowLevelSSEServerTransport):
|
| 18 |
+
"""
|
| 19 |
+
Patched SSE server transport
|
| 20 |
+
"""
|
| 21 |
+
|
| 22 |
+
@asynccontextmanager
|
| 23 |
+
async def connect_sse(self, scope: Scope, receive: Receive, send: Send):
|
| 24 |
+
"""
|
| 25 |
+
See https://github.com/modelcontextprotocol/python-sdk/pull/659/
|
| 26 |
+
"""
|
| 27 |
+
if scope["type"] != "http":
|
| 28 |
+
logger.error("connect_sse received non-HTTP request")
|
| 29 |
+
raise ValueError("connect_sse can only handle HTTP requests")
|
| 30 |
+
|
| 31 |
+
logger.debug("Setting up SSE connection")
|
| 32 |
+
read_stream: MemoryObjectReceiveStream[SessionMessage | Exception]
|
| 33 |
+
read_stream_writer: MemoryObjectSendStream[SessionMessage | Exception]
|
| 34 |
+
|
| 35 |
+
write_stream: MemoryObjectSendStream[SessionMessage]
|
| 36 |
+
write_stream_reader: MemoryObjectReceiveStream[SessionMessage]
|
| 37 |
+
|
| 38 |
+
read_stream_writer, read_stream = anyio.create_memory_object_stream(0)
|
| 39 |
+
write_stream, write_stream_reader = anyio.create_memory_object_stream(0)
|
| 40 |
+
|
| 41 |
+
session_id = uuid4()
|
| 42 |
+
self._read_stream_writers[session_id] = read_stream_writer
|
| 43 |
+
logger.debug(f"Created new session with ID: {session_id}")
|
| 44 |
+
|
| 45 |
+
# Determine the full path for the message endpoint to be sent to the client.
|
| 46 |
+
# scope['root_path'] is the prefix where the current Starlette app
|
| 47 |
+
# instance is mounted.
|
| 48 |
+
# e.g., "" if top-level, or "/api_prefix" if mounted under "/api_prefix".
|
| 49 |
+
root_path = scope.get("root_path", "")
|
| 50 |
+
|
| 51 |
+
# self._endpoint is the path *within* this app, e.g., "/messages".
|
| 52 |
+
# Concatenating them gives the full absolute path from the server root.
|
| 53 |
+
# e.g., "" + "/messages" -> "/messages"
|
| 54 |
+
# e.g., "/api_prefix" + "/messages" -> "/api_prefix/messages"
|
| 55 |
+
full_message_path_for_client = root_path.rstrip("/") + self._endpoint
|
| 56 |
+
|
| 57 |
+
# This is the URI (path + query) the client will use to POST messages.
|
| 58 |
+
client_post_uri_data = (
|
| 59 |
+
f"{quote(full_message_path_for_client)}?session_id={session_id.hex}"
|
| 60 |
+
)
|
| 61 |
+
|
| 62 |
+
sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[
|
| 63 |
+
dict[str, Any]
|
| 64 |
+
](0)
|
| 65 |
+
|
| 66 |
+
async def sse_writer():
|
| 67 |
+
logger.debug("Starting SSE writer")
|
| 68 |
+
async with sse_stream_writer, write_stream_reader:
|
| 69 |
+
await sse_stream_writer.send(
|
| 70 |
+
{"event": "endpoint", "data": client_post_uri_data}
|
| 71 |
+
)
|
| 72 |
+
logger.debug(f"Sent endpoint event: {client_post_uri_data}")
|
| 73 |
+
|
| 74 |
+
async for session_message in write_stream_reader:
|
| 75 |
+
logger.debug(f"Sending message via SSE: {session_message}")
|
| 76 |
+
await sse_stream_writer.send(
|
| 77 |
+
{
|
| 78 |
+
"event": "message",
|
| 79 |
+
"data": session_message.message.model_dump_json(
|
| 80 |
+
by_alias=True, exclude_none=True
|
| 81 |
+
),
|
| 82 |
+
}
|
| 83 |
+
)
|
| 84 |
+
|
| 85 |
+
async with anyio.create_task_group() as tg:
|
| 86 |
+
|
| 87 |
+
async def response_wrapper(scope: Scope, receive: Receive, send: Send):
|
| 88 |
+
"""
|
| 89 |
+
The EventSourceResponse returning signals a client close / disconnect.
|
| 90 |
+
In this case we close our side of the streams to signal the client that
|
| 91 |
+
the connection has been closed.
|
| 92 |
+
"""
|
| 93 |
+
await EventSourceResponse(
|
| 94 |
+
content=sse_stream_reader, data_sender_callable=sse_writer
|
| 95 |
+
)(scope, receive, send)
|
| 96 |
+
await read_stream_writer.aclose()
|
| 97 |
+
await write_stream_reader.aclose()
|
| 98 |
+
logging.debug(f"Client session disconnected {session_id}")
|
| 99 |
+
|
| 100 |
+
logger.debug("Starting SSE response task")
|
| 101 |
+
tg.start_soon(response_wrapper, scope, receive, send)
|
| 102 |
+
|
| 103 |
+
logger.debug("Yielding read and write streams")
|
| 104 |
+
yield (read_stream, write_stream)
|
src/fastmcp/server/http.py
CHANGED
|
@@ -13,7 +13,6 @@ from mcp.server.auth.middleware.bearer_auth import (
|
|
| 13 |
from mcp.server.auth.provider import OAuthAuthorizationServerProvider
|
| 14 |
from mcp.server.auth.routes import create_auth_routes
|
| 15 |
from mcp.server.auth.settings import AuthSettings
|
| 16 |
-
from mcp.server.sse import SseServerTransport
|
| 17 |
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
|
| 18 |
from starlette.applications import Starlette
|
| 19 |
from starlette.middleware import Middleware
|
|
@@ -23,6 +22,7 @@ from starlette.responses import Response
|
|
| 23 |
from starlette.routing import Mount, Route
|
| 24 |
from starlette.types import Receive, Scope, Send
|
| 25 |
|
|
|
|
| 26 |
from fastmcp.utilities.logging import get_logger
|
| 27 |
|
| 28 |
if TYPE_CHECKING:
|
|
|
|
| 13 |
from mcp.server.auth.provider import OAuthAuthorizationServerProvider
|
| 14 |
from mcp.server.auth.routes import create_auth_routes
|
| 15 |
from mcp.server.auth.settings import AuthSettings
|
|
|
|
| 16 |
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
|
| 17 |
from starlette.applications import Starlette
|
| 18 |
from starlette.middleware import Middleware
|
|
|
|
| 22 |
from starlette.routing import Mount, Route
|
| 23 |
from starlette.types import Receive, Scope, Send
|
| 24 |
|
| 25 |
+
from fastmcp.low_level.sse_server_transport import SseServerTransport
|
| 26 |
from fastmcp.utilities.logging import get_logger
|
| 27 |
|
| 28 |
if TYPE_CHECKING:
|
tests/client/test_sse.py
CHANGED
|
@@ -5,6 +5,8 @@ from collections.abc import Generator
|
|
| 5 |
import pytest
|
| 6 |
import uvicorn
|
| 7 |
from mcp.types import TextResourceContents
|
|
|
|
|
|
|
| 8 |
|
| 9 |
from fastmcp.client import Client
|
| 10 |
from fastmcp.client.transports import SSETransport
|
|
@@ -90,3 +92,30 @@ async def test_http_headers(sse_server: str):
|
|
| 90 |
json_result = json.loads(raw_result[0].text)
|
| 91 |
assert "x-demo-header" in json_result
|
| 92 |
assert json_result["x-demo-header"] == "ABC"
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 5 |
import pytest
|
| 6 |
import uvicorn
|
| 7 |
from mcp.types import TextResourceContents
|
| 8 |
+
from starlette.applications import Starlette
|
| 9 |
+
from starlette.routing import Mount
|
| 10 |
|
| 11 |
from fastmcp.client import Client
|
| 12 |
from fastmcp.client.transports import SSETransport
|
|
|
|
| 92 |
json_result = json.loads(raw_result[0].text)
|
| 93 |
assert "x-demo-header" in json_result
|
| 94 |
assert json_result["x-demo-header"] == "ABC"
|
| 95 |
+
|
| 96 |
+
|
| 97 |
+
def run_nested_server(host: str, port: int) -> None:
|
| 98 |
+
try:
|
| 99 |
+
app = fastmcp_server().sse_app()
|
| 100 |
+
mount = Starlette(routes=[Mount("/nest-inner", app=app)])
|
| 101 |
+
mount2 = Starlette(routes=[Mount("/nest-outer", app=mount)])
|
| 102 |
+
server = uvicorn.Server(
|
| 103 |
+
config=uvicorn.Config(app=mount2, host=host, port=port, log_level="error")
|
| 104 |
+
)
|
| 105 |
+
server.run()
|
| 106 |
+
except Exception as e:
|
| 107 |
+
print(f"Server error: {e}")
|
| 108 |
+
sys.exit(1)
|
| 109 |
+
sys.exit(0)
|
| 110 |
+
|
| 111 |
+
|
| 112 |
+
async def test_nested_sse_server_resolves_correctly():
|
| 113 |
+
# tests patch for
|
| 114 |
+
# https://github.com/modelcontextprotocol/python-sdk/pull/659
|
| 115 |
+
|
| 116 |
+
with run_server_in_process(run_nested_server) as url:
|
| 117 |
+
async with Client(
|
| 118 |
+
transport=SSETransport(f"{url}/nest-outer/nest-inner/sse")
|
| 119 |
+
) as client:
|
| 120 |
+
result = await client.ping()
|
| 121 |
+
assert result is True
|