Incorporate streamable HTTP server changes

This commit is contained in:
Jeremiah Lowin 2025-05-07 12:13:17 -04:00
commit cabe668e41
5 changed files with 336 additions and 11 deletions

View file

@ -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"<WebSocket(url='{self.url}')>"
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"<StreamableHttp(url='{self.url}')>"
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"):

View file

@ -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,
)

View file

@ -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,

View file

@ -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."""

View file

@ -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"