mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-21 13:04:18 +02:00
Incorporate streamable HTTP server changes
This commit is contained in:
parent
9c936168cd
commit
cabe668e41
5 changed files with 336 additions and 11 deletions
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
102
tests/client/test_streamable_http.py
Normal file
102
tests/client/test_streamable_http.py
Normal 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"
|
||||
Loading…
Add table
Add a link
Reference in a new issue