From cabe668e41934a2d09060e124f4f6673646dc64e Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Wed, 7 May 2025 12:13:17 -0400 Subject: [PATCH] Incorporate streamable HTTP server changes --- src/fastmcp/client/transports.py | 33 ++++++- src/fastmcp/server/http.py | 140 ++++++++++++++++++++++++++- src/fastmcp/server/server.py | 65 +++++++++++-- src/fastmcp/settings.py | 7 ++ tests/client/test_streamable_http.py | 102 +++++++++++++++++++ 5 files changed, 336 insertions(+), 11 deletions(-) create mode 100644 tests/client/test_streamable_http.py diff --git a/src/fastmcp/client/transports.py b/src/fastmcp/client/transports.py index 4f73e7c80..efcdba9b0 100644 --- a/src/fastmcp/client/transports.py +++ b/src/fastmcp/client/transports.py @@ -18,6 +18,7 @@ from mcp.client.session import ( ) from mcp.client.sse import sse_client from mcp.client.stdio import stdio_client +from mcp.client.streamable_http import streamablehttp_client from mcp.client.websocket import websocket_client from mcp.shared.memory import create_connected_server_and_client_session from pydantic import AnyUrl @@ -98,6 +99,33 @@ class WSTransport(ClientTransport): return f"" +class StreamableHttpTransport(ClientTransport): + """Transport implementation that connects to an MCP server via Streamable HTTP Requests.""" + + def __init__(self, url: str | AnyUrl, headers: dict[str, str] | None = None): + if isinstance(url, AnyUrl): + url = str(url) + if not isinstance(url, str) or not url.startswith("http"): + raise ValueError("Invalid HTTP/S URL provided for Streamable HTTP.") + self.url = url + self.headers = headers or {} + + @contextlib.asynccontextmanager + async def connect_session( + self, **session_kwargs: Unpack[SessionKwargs] + ) -> AsyncIterator[ClientSession]: + async with streamablehttp_client(self.url, headers=self.headers) as transport: + read_stream, write_stream, _ = transport + async with ClientSession( + read_stream, write_stream, **session_kwargs + ) as session: + await session.initialize() + yield session + + def __repr__(self) -> str: + return f"" + + class SSETransport(ClientTransport): """Transport implementation that connects to an MCP server via Server-Sent Events.""" @@ -442,7 +470,10 @@ def infer_transport( # the transport is an http(s) URL elif isinstance(transport, AnyUrl | str) and str(transport).startswith("http"): - return SSETransport(url=transport) + if str(transport).endswith("/sse"): + return SSETransport(url=transport) + else: + return StreamableHttpTransport(url=transport) # the transport is a websocket URL elif isinstance(transport, AnyUrl | str) and str(transport).startswith("ws"): diff --git a/src/fastmcp/server/http.py b/src/fastmcp/server/http.py index 86dd668a7..dad09253f 100644 --- a/src/fastmcp/server/http.py +++ b/src/fastmcp/server/http.py @@ -1,7 +1,7 @@ from __future__ import annotations -from collections.abc import Generator -from contextlib import contextmanager +from collections.abc import AsyncGenerator, Generator +from contextlib import asynccontextmanager, contextmanager from contextvars import ContextVar from typing import TYPE_CHECKING @@ -24,8 +24,18 @@ from starlette.types import Receive, Scope, Send from fastmcp.utilities.logging import get_logger +# Import these conditionally to handle case where they might not be available +try: + from mcp.server.streamable_http import EventStore + from mcp.server.streamable_http_manager import StreamableHTTPSessionManager + + STREAMABLE_HTTP_AVAILABLE = True +except ImportError: + STREAMABLE_HTTP_AVAILABLE = False + + if TYPE_CHECKING: - from fastmcp import FastMCP + from fastmcp.server.server import FastMCP logger = get_logger(__name__) @@ -53,7 +63,10 @@ class RequestContextMiddleware: self.app = app async def __call__(self, scope, receive, send): - with set_http_request(Request(scope)): + if scope["type"] == "http": + with set_http_request(Request(scope)): + await self.app(scope, receive, send) + else: await self.app(scope, receive, send) @@ -170,3 +183,122 @@ def create_sse_app( # Create and return the Starlette app with middleware return Starlette(debug=debug, routes=routes, middleware=middleware) + + +def create_streamable_http_app( + server: FastMCP, + streamable_http_path: str, + event_store: EventStore | None = None, + auth_server_provider: OAuthAuthorizationServerProvider | None = None, + auth_settings: AuthSettings | None = None, + json_response: bool = False, + stateless_http: bool = False, + debug: bool = False, + additional_routes: list[Route] | list[Mount] | list[Route | Mount] | None = None, +) -> Starlette: + """Return an instance of the StreamableHTTP server app. + + Args: + server: The FastMCP server instance + streamable_http_path: Path for StreamableHTTP connections + event_store: Optional event store for session management + auth_server_provider: Optional auth provider + auth_settings: Optional auth settings + json_response: Whether to use JSON response format + stateless_http: Whether to use stateless mode (new transport per request) + debug: Whether to enable debug mode + additional_routes: Optional list of custom routes + + Returns: + A Starlette application with StreamableHTTP support + """ + if not STREAMABLE_HTTP_AVAILABLE: + raise ImportError( + "StreamableHTTP transport is not available. Make sure your version of `mcp` is up-to-date." + ) + + # Create session manager using the provided event store + session_manager = StreamableHTTPSessionManager( + app=server._mcp_server, + event_store=event_store, + json_response=json_response, + stateless=stateless_http, + ) + + # Create the ASGI handler + async def handle_streamable_http( + scope: Scope, receive: Receive, send: Send + ) -> None: + await session_manager.handle_request(scope, receive, send) + + # Configure routes and middleware + routes: list[Route | Mount] = [] + middleware: list[Middleware] = [] + + # Handle authentication configuration + if auth_server_provider: + # Ensure auth settings are provided when auth provider is present + if not auth_settings: + raise ValueError( + "auth_settings must be provided when auth_server_provider is specified" + ) + + # Configure auth middleware + middleware = [ + Middleware( + AuthenticationMiddleware, + backend=BearerAuthBackend(provider=auth_server_provider), + ), + Middleware(AuthContextMiddleware), + ] + + # Get required scopes for authentication + required_scopes = auth_settings.required_scopes or [] + + # Add auth routes + routes.extend( + create_auth_routes( + provider=auth_server_provider, + issuer_url=auth_settings.issuer_url, + service_documentation_url=auth_settings.service_documentation_url, + client_registration_options=auth_settings.client_registration_options, + revocation_options=auth_settings.revocation_options, + ) + ) + + # Add authenticated route + routes.append( + Mount( + streamable_http_path, + app=RequireAuthMiddleware(handle_streamable_http, required_scopes), + ) + ) + else: + # No authentication required + routes.append( + Mount( + streamable_http_path, + app=handle_streamable_http, + ) + ) + + # Add custom routes with lowest precedence + if additional_routes: + routes.extend(additional_routes) + + # Add RequestContextMiddleware as the outermost middleware + middleware.append(Middleware(RequestContextMiddleware)) + + # Create a lifespan manager to start and stop the session manager + @asynccontextmanager + async def lifespan(app: Starlette) -> AsyncGenerator[None, None]: + async with session_manager.run(): + yield + + # Create and return the Starlette app with middleware + return Starlette( + debug=debug, + routes=routes, + middleware=middleware, + lifespan=lifespan, + ) diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index f7073f5c7..0fcbb15c5 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -146,6 +146,7 @@ class FastMCP(Generic[LifespanResultT]): "is specified" ) self._auth_server_provider = auth_server_provider + self._additional_http_routes: list[Route] = [] self.dependencies = self.settings.dependencies @@ -167,30 +168,36 @@ class FastMCP(Generic[LifespanResultT]): return self._mcp_server.instructions async def run_async( - self, transport: Literal["stdio", "sse"] | None = None, **transport_kwargs: Any + self, + transport: Literal["stdio", "sse", "streamable-http"] | None = None, + **transport_kwargs: Any, ) -> None: """Run the FastMCP server asynchronously. Args: - transport: Transport protocol to use ("stdio" or "sse") + transport: Transport protocol to use ("stdio", "sse", or "streamable-http") """ if transport is None: transport = "stdio" - if transport not in ["stdio", "sse"]: + if transport not in ["stdio", "sse", "streamable-http"]: raise ValueError(f"Unknown transport: {transport}") if transport == "stdio": await self.run_stdio_async(**transport_kwargs) - else: # transport == "sse" + elif transport == "sse": await self.run_sse_async(**transport_kwargs) + else: # transport == "streamable-http" + await self.run_streamable_http_async(**transport_kwargs) def run( - self, transport: Literal["stdio", "sse"] | None = None, **transport_kwargs: Any + self, + transport: Literal["stdio", "sse", "streamable-http"] | None = None, + **transport_kwargs: Any, ) -> None: """Run the FastMCP server. Note this is a synchronous function. Args: - transport: Transport protocol to use ("stdio" or "sse") + transport: Transport protocol to use ("stdio", "sse", or "streamable-http") """ logger.info(f'Starting server "{self.name}"...') @@ -743,6 +750,52 @@ class FastMCP(Generic[LifespanResultT]): additional_routes=self._additional_http_routes, ) + def streamable_http_app(self) -> Starlette: + """Return an instance of the StreamableHTTP server app.""" + try: + from fastmcp.server.http import create_streamable_http_app + + return create_streamable_http_app( + server=self, + streamable_http_path=self.settings.streamable_http_path, + event_store=None, + auth_server_provider=self._auth_server_provider, + auth_settings=self.settings.auth, + json_response=self.settings.json_response, + stateless_http=self.settings.stateless_http, + debug=self.settings.debug, + additional_routes=self._additional_http_routes, + ) + except ImportError as e: + logger.error(f"Failed to create StreamableHTTP app: {e}") + raise ImportError( + "StreamableHTTP transport is not available. Make sure your version of `mcp` is up-to-date." + ) from e + + async def run_streamable_http_async( + self, + host: str | None = None, + port: int | None = None, + log_level: str | None = None, + uvicorn_config: dict | None = None, + ) -> None: + """Run the server using StreamableHTTP transport.""" + uvicorn_config = uvicorn_config or {} + uvicorn_config.setdefault("timeout_graceful_shutdown", 0) + + app = self.streamable_http_app() + + config = uvicorn.Config( + app, + host=host or self.settings.host, + port=port or self.settings.port, + log_level=log_level or self.settings.log_level.lower(), + lifespan="on", + **uvicorn_config, + ) + server = uvicorn.Server(config) + await server.serve() + def mount( self, prefix: str, diff --git a/src/fastmcp/settings.py b/src/fastmcp/settings.py index 808f8a2a9..21395ce7c 100644 --- a/src/fastmcp/settings.py +++ b/src/fastmcp/settings.py @@ -61,6 +61,7 @@ class ServerSettings(BaseSettings): port: int = 8000 sse_path: str = "/sse" message_path: str = "/messages/" + streamable_http_path: str = "/mcp" debug: bool = False # resource settings @@ -82,6 +83,12 @@ class ServerSettings(BaseSettings): auth: AuthSettings | None = None + # StreamableHTTP settings + json_response: bool = False + stateless_http: bool = ( + False # If True, uses true stateless mode (new transport per request) + ) + class ClientSettings(BaseSettings): """FastMCP client settings.""" diff --git a/tests/client/test_streamable_http.py b/tests/client/test_streamable_http.py new file mode 100644 index 000000000..18455c3ff --- /dev/null +++ b/tests/client/test_streamable_http.py @@ -0,0 +1,102 @@ +import json +import sys +from collections.abc import Generator + +import pytest +import uvicorn +from mcp.types import TextResourceContents + +from fastmcp.client import Client +from fastmcp.client.transports import StreamableHttpTransport +from fastmcp.server.dependencies import get_http_request +from fastmcp.server.server import FastMCP +from fastmcp.utilities.tests import run_server_in_process + + +def fastmcp_server(): + """Fixture that creates a FastMCP server with tools, resources, and prompts.""" + server = FastMCP("TestServer") + + # Add a tool + @server.tool() + def greet(name: str) -> str: + """Greet someone by name.""" + return f"Hello, {name}!" + + # Add a second tool + @server.tool() + def add(a: int, b: int) -> int: + """Add two numbers together.""" + return a + b + + # Add a resource + @server.resource(uri="data://users") + async def get_users(): + return ["Alice", "Bob", "Charlie"] + + # Add a resource template + @server.resource(uri="data://user/{user_id}") + async def get_user(user_id: str): + return {"id": user_id, "name": f"User {user_id}", "active": True} + + @server.resource(uri="request://headers") + async def get_headers() -> dict[str, str]: + request = get_http_request() + + return dict(request.headers) + + # Add a prompt + @server.prompt() + def welcome(name: str) -> str: + """Example greeting prompt.""" + return f"Welcome to FastMCP, {name}!" + + return server + + +def run_server(host: str, port: int) -> None: + try: + app = fastmcp_server().streamable_http_app() + server = uvicorn.Server( + config=uvicorn.Config( + app=app, + host=host, + port=port, + log_level="error", + lifespan="on", + ) + ) + server.run() + except Exception as e: + print(f"Server error: {e}") + sys.exit(1) + sys.exit(0) + + +@pytest.fixture(scope="module") +def streamable_http_server() -> Generator[str, None, None]: + with run_server_in_process(run_server) as url: + yield f"{url}/mcp" + + +async def test_ping(streamable_http_server: str): + """Test pinging the server.""" + async with Client( + transport=StreamableHttpTransport(streamable_http_server) + ) as client: + result = await client.ping() + assert result is True + + +async def test_http_headers(streamable_http_server: str): + """Test getting HTTP headers from the server.""" + async with Client( + transport=StreamableHttpTransport( + streamable_http_server, headers={"X-DEMO-HEADER": "ABC"} + ) + ) as client: + raw_result = await client.read_resource("request://headers") + assert isinstance(raw_result[0], TextResourceContents) + json_result = json.loads(raw_result[0].text) + assert "x-demo-header" in json_result + assert json_result["x-demo-header"] == "ABC"