Merge pull request #398 from jlowin/middleware

Allow users to pass middleware to starlette app constructors
This commit is contained in:
Jeremiah Lowin 2025-05-10 09:31:12 -04:00 committed by GitHub
commit 3fd015f180
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 321 additions and 49 deletions

View file

@ -63,10 +63,34 @@ Or, from the command line:
uvicorn path.to.your.app:http_app --host 0.0.0.0 --port 8000
```
### Custom Middleware
<VersionBadge version="2.3.2" />
You can add custom Starlette middleware to your FastMCP ASGI apps by passing a list of middleware instances to the app creation methods:
```python
from fastmcp import FastMCP
from starlette.middleware import Middleware
from starlette.middleware.cors import CORSMiddleware
# Create your FastMCP server
mcp = FastMCP("MyServer")
# Define custom middleware
custom_middleware = [
Middleware(CORSMiddleware, allow_origins=["*"]),
]
# Create ASGI app with custom middleware
http_app = mcp.streamable_http_app(middleware=custom_middleware)
```
## Starlette Integration
<VersionBadge version="2.3.1" />
You can mount your FastMCP server in another Starlette application using the `Mount` class.
```python
@ -127,6 +151,8 @@ For Streamable HTTP transport, you **must** pass the lifespan context from the F
</Warning>
## FastAPI Integration
<VersionBadge version="2.3.1" />
FastAPI is built on Starlette, so you can mount your FastMCP server in a similar way:
```python

View file

@ -3,7 +3,7 @@ from __future__ import annotations
from collections.abc import AsyncGenerator, Callable, Generator
from contextlib import asynccontextmanager, contextmanager
from contextvars import ContextVar
from typing import TYPE_CHECKING, cast
from typing import TYPE_CHECKING
from mcp.server.auth.middleware.auth_context import AuthContextMiddleware
from mcp.server.auth.middleware.bearer_auth import (
@ -19,7 +19,7 @@ from starlette.middleware import Middleware
from starlette.middleware.authentication import AuthenticationMiddleware
from starlette.requests import Request
from starlette.responses import Response
from starlette.routing import Mount, Route
from starlette.routing import BaseRoute, Mount, Route
from starlette.types import Receive, Scope, Send
from fastmcp.low_level.sse_server_transport import SseServerTransport
@ -64,7 +64,7 @@ class RequestContextMiddleware:
def setup_auth_middleware_and_routes(
auth_server_provider: OAuthAuthorizationServerProvider | None,
auth_settings: AuthSettings | None,
) -> tuple[list[Middleware], list[Route | Mount], list[str]]:
) -> tuple[list[Middleware], list[BaseRoute], list[str]]:
"""Set up authentication middleware and routes if auth is enabled.
Args:
@ -75,7 +75,7 @@ def setup_auth_middleware_and_routes(
Tuple of (middleware, auth_routes, required_scopes)
"""
middleware: list[Middleware] = []
auth_routes: list[Route | Mount] = []
auth_routes: list[BaseRoute] = []
required_scopes: list[str] = []
if auth_server_provider:
@ -108,7 +108,7 @@ def setup_auth_middleware_and_routes(
def create_base_app(
routes: list[Route | Mount],
routes: list[BaseRoute],
middleware: list[Middleware],
debug: bool = False,
lifespan: Callable | None = None,
@ -142,7 +142,8 @@ def create_sse_app(
auth_server_provider: OAuthAuthorizationServerProvider | None = None,
auth_settings: AuthSettings | None = None,
debug: bool = False,
additional_routes: list[Route] | list[Mount] | list[Route | Mount] | None = None,
routes: list[BaseRoute] | None = None,
middleware: list[Middleware] | None = None,
) -> Starlette:
"""Return an instance of the SSE server app.
@ -153,11 +154,15 @@ def create_sse_app(
auth_server_provider: Optional auth provider
auth_settings: Optional auth settings
debug: Whether to enable debug mode
additional_routes: Optional list of custom routes
routes: Optional list of custom routes
middleware: Optional list of middleware
Returns:
A Starlette application with RequestContextMiddleware
"""
server_routes: list[BaseRoute] = []
server_middleware: list[Middleware] = []
# Set up SSE transport
sse = SseServerTransport(message_path)
@ -172,24 +177,24 @@ def create_sse_app(
return Response()
# Get auth middleware and routes
middleware, auth_routes, required_scopes = setup_auth_middleware_and_routes(
auth_middleware, auth_routes, required_scopes = setup_auth_middleware_and_routes(
auth_server_provider, auth_settings
)
# Initialize routes with auth routes
routes: list[Route | Mount] = auth_routes.copy()
server_routes.extend(auth_routes)
server_middleware.extend(auth_middleware)
# Add SSE routes with or without auth
if auth_server_provider:
# Auth is enabled, wrap endpoints with RequireAuthMiddleware
routes.append(
server_routes.append(
Route(
sse_path,
endpoint=RequireAuthMiddleware(handle_sse, required_scopes),
methods=["GET"],
)
)
routes.append(
server_routes.append(
Mount(
message_path,
app=RequireAuthMiddleware(sse.handle_post_message, required_scopes),
@ -200,14 +205,14 @@ def create_sse_app(
async def sse_endpoint(request: Request) -> Response:
return await handle_sse(request.scope, request.receive, request._send) # type: ignore[reportPrivateUsage]
routes.append(
server_routes.append(
Route(
sse_path,
endpoint=sse_endpoint,
methods=["GET"],
)
)
routes.append(
server_routes.append(
Mount(
message_path,
app=sse.handle_post_message,
@ -215,13 +220,17 @@ def create_sse_app(
)
# Add custom routes with lowest precedence
if additional_routes:
routes.extend(cast(list[Route | Mount], additional_routes))
if routes:
server_routes.extend(routes)
# Add middleware
if middleware:
server_middleware.extend(middleware)
# Create and return the app
return create_base_app(
routes=routes,
middleware=middleware,
routes=server_routes,
middleware=server_middleware,
debug=debug,
)
@ -235,7 +244,8 @@ def create_streamable_http_app(
json_response: bool = False,
stateless_http: bool = False,
debug: bool = False,
additional_routes: list[Route] | list[Mount] | list[Route | Mount] | None = None,
routes: list[BaseRoute] | None = None,
middleware: list[Middleware] | None = None,
) -> Starlette:
"""Return an instance of the StreamableHTTP server app.
@ -248,11 +258,15 @@ def create_streamable_http_app(
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
routes: Optional list of custom routes
middleware: Optional list of middleware
Returns:
A Starlette application with StreamableHTTP support
"""
server_routes: list[BaseRoute] = []
server_middleware: list[Middleware] = []
# Create session manager using the provided event store
session_manager = StreamableHTTPSessionManager(
app=server._mcp_server,
@ -268,17 +282,17 @@ def create_streamable_http_app(
await session_manager.handle_request(scope, receive, send)
# Get auth middleware and routes
middleware, auth_routes, required_scopes = setup_auth_middleware_and_routes(
auth_middleware, auth_routes, required_scopes = setup_auth_middleware_and_routes(
auth_server_provider, auth_settings
)
# Initialize routes with auth routes
routes: list[Route | Mount] = auth_routes.copy()
server_routes.extend(auth_routes)
server_middleware.extend(auth_middleware)
# Add StreamableHTTP routes with or without auth
if auth_server_provider:
# Auth is enabled, wrap endpoint with RequireAuthMiddleware
routes.append(
server_routes.append(
Mount(
streamable_http_path,
app=RequireAuthMiddleware(handle_streamable_http, required_scopes),
@ -286,7 +300,7 @@ def create_streamable_http_app(
)
else:
# No auth required
routes.append(
server_routes.append(
Mount(
streamable_http_path,
app=handle_streamable_http,
@ -294,8 +308,12 @@ def create_streamable_http_app(
)
# Add custom routes with lowest precedence
if additional_routes:
routes.extend(cast(list[Route | Mount], additional_routes))
if routes:
server_routes.extend(routes)
# Add middleware
if middleware:
server_middleware.extend(middleware)
# Create a lifespan manager to start and stop the session manager
@asynccontextmanager
@ -305,8 +323,8 @@ def create_streamable_http_app(
# Create and return the app with lifespan
return create_base_app(
routes=routes,
middleware=middleware,
routes=server_routes,
middleware=server_middleware,
debug=debug,
lifespan=lifespan,
)

View file

@ -35,9 +35,10 @@ from mcp.types import ResourceTemplate as MCPResourceTemplate
from mcp.types import Tool as MCPTool
from pydantic import AnyUrl
from starlette.applications import Starlette
from starlette.middleware import Middleware
from starlette.requests import Request
from starlette.responses import Response
from starlette.routing import Route
from starlette.routing import BaseRoute, Route
import fastmcp.server
import fastmcp.settings
@ -147,7 +148,7 @@ class FastMCP(Generic[LifespanResultT]):
)
self._auth_server_provider = auth_server_provider
self._additional_http_routes: list[Route] = []
self._additional_http_routes: list[BaseRoute] = []
self.dependencies = self.settings.dependencies
# Set up MCP protocol handlers
@ -744,8 +745,16 @@ class FastMCP(Generic[LifespanResultT]):
self,
path: str | None = None,
message_path: str | None = None,
middleware: list[Middleware] | None = None,
) -> Starlette:
"""Return an instance of the SSE server app."""
"""
Create a Starlette app for the SSE server.
Args:
path: The path to the SSE endpoint
message_path: The path to the message endpoint
middleware: A list of middleware to apply to the app
"""
return create_sse_app(
server=self,
message_path=message_path or self.settings.message_path,
@ -753,11 +762,22 @@ class FastMCP(Generic[LifespanResultT]):
auth_server_provider=self._auth_server_provider,
auth_settings=self.settings.auth,
debug=self.settings.debug,
additional_routes=self._additional_http_routes,
routes=self._additional_http_routes,
middleware=middleware,
)
def streamable_http_app(self, path: str | None = None) -> Starlette:
"""Return an instance of the StreamableHTTP server app."""
def streamable_http_app(
self,
path: str | None = None,
middleware: list[Middleware] | None = None,
) -> Starlette:
"""
Create a Starlette app for the StreamableHTTP server.
Args:
path: The path to the StreamableHTTP endpoint
middleware: A list of middleware to apply to the app
"""
from fastmcp.server.http import create_streamable_http_app
return create_streamable_http_app(
@ -769,7 +789,8 @@ class FastMCP(Generic[LifespanResultT]):
json_response=self.settings.json_response,
stateless_http=self.settings.stateless_http,
debug=self.settings.debug,
additional_routes=self._additional_http_routes,
routes=self._additional_http_routes,
middleware=middleware,
)
async def run_streamable_http_async(

12
test.py
View file

@ -1,12 +0,0 @@
from fastmcp import FastMCP
mcp = FastMCP()
if __name__ == "__main__":
mcp.run(
transport="streamable-http",
host="127.0.0.1",
port=4200,
path="/my-custom-path/",
log_level="debug",
)

View file

@ -0,0 +1,219 @@
"""Tests for custom middleware in HTTP servers."""
from collections.abc import Callable
from typing import Any
import httpx
import pytest
from httpx import ASGITransport
from starlette.middleware import Middleware
from starlette.middleware.base import BaseHTTPMiddleware
from starlette.requests import Request
from starlette.responses import JSONResponse
from starlette.routing import BaseRoute, Route
from starlette.types import ASGIApp
from fastmcp.server import FastMCP
from fastmcp.server.http import create_sse_app, create_streamable_http_app
class HeaderMiddleware(BaseHTTPMiddleware):
"""Simple middleware that adds a custom header to responses."""
def __init__(self, app: ASGIApp, header_name: str, header_value: str):
super().__init__(app)
self.header_name = header_name
self.header_value = header_value
async def dispatch(self, request: Request, call_next: Callable):
response = await call_next(request)
response.headers[self.header_name] = self.header_value
return response
class RequestModifierMiddleware(BaseHTTPMiddleware):
"""Middleware that adds a value to request state."""
def __init__(self, app: ASGIApp, key: str, value: Any):
super().__init__(app)
self.key = key
self.value = value
async def dispatch(self, request: Request, call_next: Callable):
request.state.custom_value = {self.key: self.value}
return await call_next(request)
async def endpoint_handler(request: Request):
"""Endpoint that returns request state or headers."""
if hasattr(request.state, "custom_value"):
return JSONResponse({"state": request.state.custom_value})
return JSONResponse({"message": "Hello, world!"})
@pytest.mark.asyncio
async def test_sse_app_with_custom_middleware():
"""Test that custom middleware works with SSE app."""
server = FastMCP(name="TestServer")
# Create custom middleware
custom_middleware = [
Middleware(
HeaderMiddleware, header_name="X-Custom-Header", header_value="test-value"
)
]
# Add a test route to server's additional routes
routes: list[BaseRoute] = [Route("/test", endpoint_handler)]
server._additional_http_routes = routes
# Create the app with custom middleware
app = server.sse_app(middleware=custom_middleware)
# Create a test client
transport = ASGITransport(app=app)
async with httpx.AsyncClient(
transport=transport, base_url="http://testserver"
) as client:
response = await client.get("/test")
# Verify middleware was applied
assert response.status_code == 200
assert response.headers["X-Custom-Header"] == "test-value"
@pytest.mark.asyncio
async def test_streamable_http_app_with_custom_middleware():
"""Test that custom middleware works with StreamableHTTP app."""
server = FastMCP(name="TestServer")
# Create custom middleware
custom_middleware = [
Middleware(
HeaderMiddleware, header_name="X-Custom-Header", header_value="test-value"
)
]
# Add a test route to server's additional routes
routes: list[BaseRoute] = [Route("/test", endpoint_handler)]
server._additional_http_routes = routes
# Create the app with custom middleware
app = server.streamable_http_app(middleware=custom_middleware)
# Create a test client
transport = ASGITransport(app=app)
async with httpx.AsyncClient(
transport=transport, base_url="http://testserver"
) as client:
response = await client.get("/test")
# Verify middleware was applied
assert response.status_code == 200
assert response.headers["X-Custom-Header"] == "test-value"
@pytest.mark.asyncio
async def test_create_sse_app_with_custom_middleware():
"""Test that custom middleware works with create_sse_app function."""
server = FastMCP(name="TestServer")
# Create custom middleware
custom_middleware = [
Middleware(RequestModifierMiddleware, key="modified_by", value="middleware")
]
# Add a test route
additional_routes: list[BaseRoute] = [Route("/test", endpoint_handler)]
# Create the app with custom middleware
app = create_sse_app(
server=server,
message_path="/message",
sse_path="/sse",
middleware=custom_middleware,
routes=additional_routes,
)
# Create a test client
transport = ASGITransport(app=app)
async with httpx.AsyncClient(
transport=transport, base_url="http://testserver"
) as client:
response = await client.get("/test")
# Verify middleware was applied
assert response.status_code == 200
data = response.json()
assert "state" in data
assert data["state"]["modified_by"] == "middleware"
@pytest.mark.asyncio
async def test_create_streamable_http_app_with_custom_middleware():
"""Test that custom middleware works with create_streamable_http_app function."""
server = FastMCP(name="TestServer")
# Create custom middleware
custom_middleware = [
Middleware(RequestModifierMiddleware, key="modified_by", value="middleware")
]
# Add a test route
additional_routes: list[BaseRoute] = [Route("/test", endpoint_handler)]
# Create the app with custom middleware
app = create_streamable_http_app(
server=server,
streamable_http_path="/streamable",
middleware=custom_middleware,
routes=additional_routes,
)
# Create a test client
transport = ASGITransport(app=app)
async with httpx.AsyncClient(
transport=transport, base_url="http://testserver"
) as client:
response = await client.get("/test")
# Verify middleware was applied
assert response.status_code == 200
data = response.json()
assert "state" in data
assert data["state"]["modified_by"] == "middleware"
@pytest.mark.asyncio
async def test_multiple_middleware_ordering():
"""Test that multiple middleware are applied in the correct order."""
server = FastMCP(name="TestServer")
# Create multiple middleware
custom_middleware = [
Middleware(
HeaderMiddleware, header_name="X-First-Header", header_value="first"
),
Middleware(
HeaderMiddleware, header_name="X-Second-Header", header_value="second"
),
]
# Add a test route to server's additional routes
routes: list[BaseRoute] = [Route("/test", endpoint_handler)]
server._additional_http_routes = routes
# Create the app with custom middleware
app = server.sse_app(middleware=custom_middleware)
# Create a test client
transport = ASGITransport(app=app)
async with httpx.AsyncClient(
transport=transport, base_url="http://testserver"
) as client:
response = await client.get("/test")
# Verify both middleware were applied
assert response.status_code == 200
assert response.headers["X-First-Header"] == "first"
assert response.headers["X-Second-Header"] == "second"