diff --git a/src/mcp/server/fastmcp/server.py b/src/mcp/server/fastmcp/server.py index bf0ce880a5..15bbea1fe0 100644 --- a/src/mcp/server/fastmcp/server.py +++ b/src/mcp/server/fastmcp/server.py @@ -75,6 +75,7 @@ class Settings(BaseSettings, Generic[LifespanResultT]): port: int = 8000 sse_path: str = "/sse" message_path: str = "/messages/" + base_path: str = "/" # resource settings warn_on_duplicate_resources: bool = True @@ -479,7 +480,7 @@ async def run_sse_async(self) -> None: def sse_app(self) -> Starlette: """Return an instance of the SSE server app.""" - sse = SseServerTransport(self.settings.message_path) + sse = SseServerTransport(self.settings.message_path, self.settings.base_path) async def handle_sse(request: Request) -> None: async with sse.connect_sse( diff --git a/src/mcp/server/sse.py b/src/mcp/server/sse.py index d051c25bf6..bb0911cc40 100644 --- a/src/mcp/server/sse.py +++ b/src/mcp/server/sse.py @@ -36,6 +36,8 @@ async def handle_sse(request): from typing import Any from urllib.parse import quote from uuid import UUID, uuid4 +import os +import re import anyio from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream @@ -67,7 +69,7 @@ class SseServerTransport: UUID, MemoryObjectSendStream[types.JSONRPCMessage | Exception] ] - def __init__(self, endpoint: str) -> None: + def __init__(self, endpoint: str, base_path: str) -> None: """ Creates a new SSE server transport, which will direct the client to POST messages to the relative or absolute URL given. @@ -75,6 +77,7 @@ def __init__(self, endpoint: str) -> None: super().__init__() self._endpoint = endpoint + self._base_path = base_path self._read_stream_writers = {} logger.debug(f"SseServerTransport initialized with endpoint: {endpoint}") @@ -96,6 +99,8 @@ async def connect_sse(self, scope: Scope, receive: Receive, send: Send): session_id = uuid4() session_uri = f"{quote(self._endpoint)}?session_id={session_id.hex}" + # join the base path and endpoint + session_full_url = os.path.join(self._base_path, re.sub("^/", "", session_uri, 1)) self._read_stream_writers[session_id] = read_stream_writer logger.debug(f"Created new session with ID: {session_id}") @@ -106,8 +111,8 @@ async def connect_sse(self, scope: Scope, receive: Receive, send: Send): 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": session_uri}) - logger.debug(f"Sent endpoint event: {session_uri}") + await sse_stream_writer.send({"event": "endpoint", "data": session_full_url}) + logger.debug(f"Sent endpoint event: {session_full_url}") async for message in write_stream_reader: logger.debug(f"Sending message via SSE: {message}")