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