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(