From 6bdf15601b0ce6eebd6519a3568eba17a6c999e0 Mon Sep 17 00:00:00 2001
From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com>
Date: Sat, 10 May 2025 09:19:29 -0400
Subject: [PATCH 1/3] Add middleware kwarg
---
src/fastmcp/server/http.py | 78 ++++++++++++++++++++++--------------
src/fastmcp/server/server.py | 18 ++++++---
test.py | 12 ------
3 files changed, 61 insertions(+), 47 deletions(-)
delete mode 100644 test.py
diff --git a/src/fastmcp/server/http.py b/src/fastmcp/server/http.py
index a6996bdeb..5257c24b7 100644
--- a/src/fastmcp/server/http.py
+++ b/src/fastmcp/server/http.py
@@ -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,
)
diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py
index 748a53f39..85be37412 100644
--- a/src/fastmcp/server/server.py
+++ b/src/fastmcp/server/server.py
@@ -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,6 +745,7 @@ 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."""
return create_sse_app(
@@ -753,10 +755,15 @@ 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:
+ def streamable_http_app(
+ self,
+ path: str | None = None,
+ middleware: list[Middleware] | None = None,
+ ) -> Starlette:
"""Return an instance of the StreamableHTTP server app."""
from fastmcp.server.http import create_streamable_http_app
@@ -769,7 +776,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(
diff --git a/test.py b/test.py
deleted file mode 100644
index e5931b8e6..000000000
--- a/test.py
+++ /dev/null
@@ -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",
- )
From 7f5c5878ebf10924b72de944a37fbb269aa83ff8 Mon Sep 17 00:00:00 2001
From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com>
Date: Sat, 10 May 2025 09:26:10 -0400
Subject: [PATCH 2/3] Allow passing custom middleware
---
tests/server/test_http_middleware.py | 219 +++++++++++++++++++++++++++
1 file changed, 219 insertions(+)
create mode 100644 tests/server/test_http_middleware.py
diff --git a/tests/server/test_http_middleware.py b/tests/server/test_http_middleware.py
new file mode 100644
index 000000000..2ff4127be
--- /dev/null
+++ b/tests/server/test_http_middleware.py
@@ -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"
From 042a3852b1b1e9749fd716c5f267896fed16287d Mon Sep 17 00:00:00 2001
From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com>
Date: Sat, 10 May 2025 09:29:58 -0400
Subject: [PATCH 3/3] Add docs
---
docs/deployment/asgi.mdx | 26 ++++++++++++++++++++++++++
src/fastmcp/server/server.py | 17 +++++++++++++++--
2 files changed, 41 insertions(+), 2 deletions(-)
diff --git a/docs/deployment/asgi.mdx b/docs/deployment/asgi.mdx
index f68f19435..e3a66d5ea 100644
--- a/docs/deployment/asgi.mdx
+++ b/docs/deployment/asgi.mdx
@@ -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
+
+
+
+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
+
+
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
## FastAPI Integration
+
+
FastAPI is built on Starlette, so you can mount your FastMCP server in a similar way:
```python
diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py
index 85be37412..37e97d79f 100644
--- a/src/fastmcp/server/server.py
+++ b/src/fastmcp/server/server.py
@@ -747,7 +747,14 @@ class FastMCP(Generic[LifespanResultT]):
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,
@@ -764,7 +771,13 @@ class FastMCP(Generic[LifespanResultT]):
path: str | None = None,
middleware: list[Middleware] | None = None,
) -> Starlette:
- """Return an instance of the StreamableHTTP server app."""
+ """
+ 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(