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"