diff --git a/docs/servers/middleware.mdx b/docs/servers/middleware.mdx index 5ef8d8fe1..a57882016 100644 --- a/docs/servers/middleware.mdx +++ b/docs/servers/middleware.mdx @@ -93,6 +93,7 @@ This hierarchy allows you to target your middleware logic with the right level o - `on_message`: Called for all MCP messages (requests and notifications) - `on_request`: Called specifically for MCP requests (that expect responses) - `on_notification`: Called specifically for MCP notifications (fire-and-forget) +- `on_initialize`: Called when a client connects and initializes the session (returns `None`) - `on_call_tool`: Called when tools are being executed - `on_read_resource`: Called when resources are being read - `on_get_prompt`: Called when prompts are being retrieved @@ -101,6 +102,10 @@ This hierarchy allows you to target your middleware logic with the right level o - `on_list_resource_templates`: Called when listing resource templates - `on_list_prompts`: Called when listing available prompts + +The `on_initialize` hook receives the client's initialization request but **returns `None`** rather than a result. The initialization response is handled internally by the MCP protocol and cannot be modified by middleware. This hook is useful for client detection, logging connections, or initializing session state, but not for modifying the initialization handshake itself. + + ## Component Access in Middleware Understanding how to access component information (tools, resources, prompts) in middleware is crucial for building powerful middleware functionality. The access patterns differ significantly between listing operations and execution operations. diff --git a/src/fastmcp/server/low_level.py b/src/fastmcp/server/low_level.py index eb9e87c2a..4251b3b80 100644 --- a/src/fastmcp/server/low_level.py +++ b/src/fastmcp/server/low_level.py @@ -1,5 +1,12 @@ -from typing import Any +from __future__ import annotations +import weakref +from contextlib import AsyncExitStack +from typing import TYPE_CHECKING, Any + +import anyio +import mcp.types +from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream from mcp.server.lowlevel.server import ( LifespanResultT, NotificationOptions, @@ -9,11 +16,82 @@ from mcp.server.lowlevel.server import ( Server as _Server, ) from mcp.server.models import InitializationOptions +from mcp.server.session import ServerSession +from mcp.server.stdio import stdio_server as stdio_server +from mcp.shared.message import SessionMessage +from mcp.shared.session import RequestResponder + +from fastmcp.utilities.logging import get_logger + +if TYPE_CHECKING: + from fastmcp.server.server import FastMCP + +logger = get_logger(__name__) + + +class MiddlewareServerSession(ServerSession): + """ServerSession that routes initialization requests through FastMCP middleware.""" + + def __init__(self, fastmcp: FastMCP, *args, **kwargs): + super().__init__(*args, **kwargs) + self._fastmcp_ref: weakref.ref[FastMCP] = weakref.ref(fastmcp) + + @property + def fastmcp(self) -> FastMCP: + """Get the FastMCP instance.""" + fastmcp = self._fastmcp_ref() + if fastmcp is None: + raise RuntimeError("FastMCP instance is no longer available") + return fastmcp + + async def _received_request( + self, + responder: RequestResponder[mcp.types.ClientRequest, mcp.types.ServerResult], + ): + """ + Override the _received_request method to route initialization requests + through FastMCP middleware. + + These are not handled by routes that FastMCP typically overrides and + require special handling. + """ + import fastmcp.server.context + from fastmcp.server.middleware.middleware import MiddlewareContext + + if isinstance(responder.request.root, mcp.types.InitializeRequest): + + async def call_original_handler( + ctx: MiddlewareContext, + ) -> None: + return await super(MiddlewareServerSession, self)._received_request( + responder + ) + + async with fastmcp.server.context.Context( + fastmcp=self.fastmcp + ) as fastmcp_ctx: + # Create the middleware context. + mw_context = MiddlewareContext( + message=responder.request.root, + source="client", + type="request", + method="initialize", + fastmcp_context=fastmcp_ctx, + ) + + return await self.fastmcp._apply_middleware( + mw_context, call_original_handler + ) + else: + return await super()._received_request(responder) class LowLevelServer(_Server[LifespanResultT, RequestT]): - def __init__(self, *args: Any, **kwargs: Any): + def __init__(self, fastmcp: FastMCP, *args: Any, **kwargs: Any): super().__init__(*args, **kwargs) + # Store a weak reference to FastMCP to avoid circular references + self._fastmcp_ref: weakref.ref[FastMCP] = weakref.ref(fastmcp) + # FastMCP servers support notifications for all components self.notification_options = NotificationOptions( prompts_changed=True, @@ -21,6 +99,14 @@ class LowLevelServer(_Server[LifespanResultT, RequestT]): tools_changed=True, ) + @property + def fastmcp(self) -> FastMCP: + """Get the FastMCP instance.""" + fastmcp = self._fastmcp_ref() + if fastmcp is None: + raise RuntimeError("FastMCP instance is no longer available") + return fastmcp + def create_initialization_options( self, notification_options: NotificationOptions | None = None, @@ -35,3 +121,36 @@ class LowLevelServer(_Server[LifespanResultT, RequestT]): experimental_capabilities=experimental_capabilities, **kwargs, ) + + async def run( + self, + read_stream: MemoryObjectReceiveStream[SessionMessage | Exception], + write_stream: MemoryObjectSendStream[SessionMessage], + initialization_options: InitializationOptions, + raise_exceptions: bool = False, + stateless: bool = False, + ): + """ + Overrides the run method to use the MiddlewareServerSession. + """ + async with AsyncExitStack() as stack: + lifespan_context = await stack.enter_async_context(self.lifespan(self)) + session = await stack.enter_async_context( + MiddlewareServerSession( + self.fastmcp, + read_stream, + write_stream, + initialization_options, + stateless=stateless, + ) + ) + + async with anyio.create_task_group() as tg: + async for message in session.incoming_messages: + tg.start_soon( + self._handle_message, + message, + session, + lifespan_context, + raise_exceptions, + ) diff --git a/src/fastmcp/server/middleware/middleware.py b/src/fastmcp/server/middleware/middleware.py index 8b262d4f5..0b78e4866 100644 --- a/src/fastmcp/server/middleware/middleware.py +++ b/src/fastmcp/server/middleware/middleware.py @@ -99,6 +99,8 @@ class Middleware: handler = call_next match context.method: + case "initialize": + handler = partial(self.on_initialize, call_next=handler) case "tools/call": handler = partial(self.on_call_tool, call_next=handler) case "resources/read": @@ -145,6 +147,13 @@ class Middleware: ) -> Any: return await call_next(context) + async def on_initialize( + self, + context: MiddlewareContext[mt.InitializeRequestParams], + call_next: CallNext[mt.InitializeRequestParams, None], + ) -> None: + return await call_next(context) + async def on_call_tool( self, context: MiddlewareContext[mt.CallToolRequestParams], diff --git a/src/fastmcp/server/server.py b/src/fastmcp/server/server.py index 22b434d06..887d6471b 100644 --- a/src/fastmcp/server/server.py +++ b/src/fastmcp/server/server.py @@ -195,6 +195,7 @@ class FastMCP(Generic[LifespanResultT]): self._has_lifespan = True # Generate random ID if no name provided self._mcp_server = LowLevelServer[LifespanResultT]( + fastmcp=self, name=name or self.generate_name(), version=version, instructions=instructions, diff --git a/tests/server/middleware/test_initialization_middleware.py b/tests/server/middleware/test_initialization_middleware.py new file mode 100644 index 000000000..06776e16e --- /dev/null +++ b/tests/server/middleware/test_initialization_middleware.py @@ -0,0 +1,251 @@ +"""Tests for middleware support during initialization.""" + +from typing import Any + +import mcp.types as mt + +from fastmcp import Client, FastMCP +from fastmcp.server.middleware import CallNext, Middleware, MiddlewareContext + + +class InitializationMiddleware(Middleware): + """Middleware that captures initialization details.""" + + def __init__(self): + super().__init__() + self.initialized = False + self.client_info = None + self.session_data = {} + + async def on_initialize( + self, + context: MiddlewareContext[mt.InitializeRequest], + call_next: CallNext[mt.InitializeRequest, None], + ) -> None: + """Capture initialization details and store session data.""" + self.initialized = True + + # Extract client info from the initialize params + if hasattr(context.message, "params") and hasattr( + context.message.params, "clientInfo" + ): + self.client_info = context.message.params.clientInfo + + # Store data in the context state for cross-request access + if context.fastmcp_context: + context.fastmcp_context.set_state("client_initialized", True) + if self.client_info: + context.fastmcp_context.set_state( + "client_name", getattr(self.client_info, "name", "unknown") + ) + + return await call_next(context) + + +class ClientDetectionMiddleware(Middleware): + """Middleware that detects specific clients and modifies behavior. + + This demonstrates storing data in the middleware instance itself + for cross-request access, since context state is request-scoped. + """ + + def __init__(self): + super().__init__() + self.is_test_client = False + self.tools_modified = False + self.initialization_called = False + + async def on_initialize( + self, + context: MiddlewareContext[mt.InitializeRequest], + call_next: CallNext[mt.InitializeRequest, None], + ) -> None: + """Detect test client during initialization.""" + self.initialization_called = True + + # For testing purposes, always set it to true + # Store in instance variable for cross-request access + self.is_test_client = True + + return await call_next(context) + + async def on_list_tools( + self, + context: MiddlewareContext[mt.ListToolsRequest], + call_next: CallNext[mt.ListToolsRequest, list], + ) -> list: + """Modify tools based on client detection.""" + tools = await call_next(context) + + # Use the instance variable set during initialization + if self.is_test_client: + # Add a special annotation to tools for test clients + for tool in tools: + if not hasattr(tool, "annotations"): + tool.annotations = mt.ToolAnnotations() + if tool.annotations is None: + tool.annotations = mt.ToolAnnotations() + # Mark as read-only for test clients + tool.annotations.readOnlyHint = True + self.tools_modified = True + + return tools + + +async def test_simple_initialization_hook(): + """Test that the on_initialize hook is called.""" + server = FastMCP("TestServer") + + class SimpleInitMiddleware(Middleware): + def __init__(self): + super().__init__() + self.called = False + + async def on_initialize( + self, + context: MiddlewareContext[mt.InitializeRequest], + call_next: CallNext[mt.InitializeRequest, None], + ) -> None: + self.called = True + return await call_next(context) + + middleware = SimpleInitMiddleware() + server.add_middleware(middleware) + + # Connect client + async with Client(server): + # Middleware should have been called + assert middleware.called is True, "on_initialize was not called" + + +async def test_middleware_receives_initialization(): + """Test that middleware can intercept initialization requests.""" + server = FastMCP("TestServer") + middleware = InitializationMiddleware() + server.add_middleware(middleware) + + @server.tool + def test_tool(x: int) -> str: + return f"Result: {x}" + + # Connect client + async with Client(server) as client: + # Middleware should have been called during initialization + assert middleware.initialized is True + + # Test that the tool still works + result = await client.call_tool("test_tool", {"x": 42}) + assert result.content[0].text == "Result: 42" # type: ignore[attr-defined] + + +async def test_client_detection_middleware(): + """Test middleware that detects specific clients and modifies behavior.""" + server = FastMCP("TestServer") + middleware = ClientDetectionMiddleware() + server.add_middleware(middleware) + + @server.tool + def example_tool() -> str: + return "example" + + # Connect with a client + async with Client(server) as client: + # Middleware should have been called during initialization + assert middleware.initialization_called is True + assert middleware.is_test_client is True + + # List tools to trigger modification + tools = await client.list_tools() + assert len(tools) == 1 + assert middleware.tools_modified is True + + # Check that the tool has the modified annotation + tool = tools[0] + assert tool.annotations is not None + assert tool.annotations.readOnlyHint is True + + +async def test_multiple_middleware_initialization(): + """Test that multiple middleware can handle initialization.""" + server = FastMCP("TestServer") + + init_mw = InitializationMiddleware() + detect_mw = ClientDetectionMiddleware() + + server.add_middleware(init_mw) + server.add_middleware(detect_mw) + + @server.tool + def test_tool() -> str: + return "test" + + async with Client(server) as client: + # Both middleware should have processed initialization + assert init_mw.initialized is True + assert detect_mw.initialization_called is True + assert detect_mw.is_test_client is True + + # List tools to check detection worked + await client.list_tools() + assert detect_mw.tools_modified is True + + +async def test_initialization_middleware_with_state_sharing(): + """Test that state set during initialization is available in later requests.""" + server = FastMCP("TestServer") + + class StateTrackingMiddleware(Middleware): + def __init__(self): + super().__init__() + self.init_state = {} + self.tool_state = {} + + async def on_initialize( + self, + context: MiddlewareContext[mt.InitializeRequest], + call_next: CallNext[mt.InitializeRequest, None], + ) -> None: + # Store some state during initialization + if context.fastmcp_context: + context.fastmcp_context.set_state("init_timestamp", "2024-01-01") + context.fastmcp_context.set_state("client_id", "test-123") + self.init_state["timestamp"] = "2024-01-01" + self.init_state["client_id"] = "test-123" + + return await call_next(context) + + async def on_call_tool( + self, + context: MiddlewareContext[mt.CallToolRequestParams], + call_next: CallNext[mt.CallToolRequestParams, Any], + ) -> Any: + # Try to access state from initialization + if context.fastmcp_context: + timestamp = context.fastmcp_context.get_state("init_timestamp") + client_id = context.fastmcp_context.get_state("client_id") + self.tool_state["timestamp"] = timestamp + self.tool_state["client_id"] = client_id + + return await call_next(context) + + middleware = StateTrackingMiddleware() + server.add_middleware(middleware) + + @server.tool + def test_tool() -> str: + return "success" + + async with Client(server) as client: + # Initialization should have set state + assert middleware.init_state["timestamp"] == "2024-01-01" + assert middleware.init_state["client_id"] == "test-123" + + # Call a tool - state should be accessible + result = await client.call_tool("test_tool", {}) + assert result.content[0].text == "success" # type: ignore[attr-defined] + + # State should have been accessible during tool call + # Note: State is request-scoped, so it won't persist across requests + # This test shows the pattern, but actual cross-request state would need + # external storage (Redis, DB, etc.) + # The middleware.tool_state might be None if state doesn't persist diff --git a/tests/server/middleware/test_middleware.py b/tests/server/middleware/test_middleware.py index 8469df424..d93b6a202 100644 --- a/tests/server/middleware/test_middleware.py +++ b/tests/server/middleware/test_middleware.py @@ -293,6 +293,17 @@ class TestMiddlewareHooks: result = list_prompts_calls[0].result assert isinstance(result, list) + async def test_initialize( + self, mcp_server: FastMCP, recording_middleware: RecordingMiddleware + ): + async with Client(mcp_server) as client: + await client.ping() + + assert recording_middleware.assert_called(at_least=1) + assert recording_middleware.assert_called(hook="on_message", at_least=1) + assert recording_middleware.assert_called(hook="on_request", at_least=1) + assert recording_middleware.assert_called(hook="on_initialize", at_least=1) + async def test_list_tools_filtering_middleware(self): """Test that middleware can filter tools.""" diff --git a/tests/server/middleware/test_rate_limiting.py b/tests/server/middleware/test_rate_limiting.py index 0a4e1d67b..1f17d7095 100644 --- a/tests/server/middleware/test_rate_limiting.py +++ b/tests/server/middleware/test_rate_limiting.py @@ -306,9 +306,10 @@ class TestRateLimitingMiddlewareIntegration: async def test_rate_limiting_blocks_rapid_requests(self, rate_limit_server): """Test that rate limiting blocks rapid successive requests.""" - # Very restrictive rate limit (accounting for extra list_tools calls per tool call) + # Very restrictive rate limit (accounting for initialization and list_tools calls) + # Requests: 1 initialize + 1 list_tools + 4 call_tools = 6 total before limit rate_limit_server.add_middleware( - RateLimitingMiddleware(max_requests_per_second=10.0, burst_capacity=5) + RateLimitingMiddleware(max_requests_per_second=10.0, burst_capacity=6) ) async with Client(rate_limit_server) as client: @@ -356,7 +357,7 @@ class TestRateLimitingMiddlewareIntegration: """Test sliding window rate limiting implementation.""" rate_limit_server.add_middleware( SlidingWindowRateLimitingMiddleware( - max_requests=5, # Accounting for extra list_tools calls + max_requests=6, # 1 init + 1 list_tools + 3 calls + 1 to fail window_minutes=1, # 1-minute window ) ) @@ -374,7 +375,7 @@ class TestRateLimitingMiddlewareIntegration: async def test_rate_limiting_with_different_operations(self, rate_limit_server): """Test that rate limiting applies to all types of operations.""" rate_limit_server.add_middleware( - RateLimitingMiddleware(max_requests_per_second=9.0, burst_capacity=4) + RateLimitingMiddleware(max_requests_per_second=9.0, burst_capacity=5) ) async with Client(rate_limit_server) as client: @@ -395,8 +396,8 @@ class TestRateLimitingMiddlewareIntegration: rate_limit_server.add_middleware( RateLimitingMiddleware( - max_requests_per_second=6.0, # Accounting for extra list_tools calls - burst_capacity=3, + max_requests_per_second=6.0, # Accounting for initialization and list_tools calls + burst_capacity=4, get_client_id=get_client_id, ) ) @@ -416,8 +417,8 @@ class TestRateLimitingMiddlewareIntegration: rate_limit_server.add_middleware( RateLimitingMiddleware( max_requests_per_second=6.0, - burst_capacity=4, - global_limit=True, # Accounting for extra list_tools calls + burst_capacity=5, # 1 init + 2 list_tools + 2 calls before limit + global_limit=True, # Accounting for initialization and list_tools calls ) ) @@ -435,7 +436,7 @@ class TestRateLimitingMiddlewareIntegration: rate_limit_server.add_middleware( RateLimitingMiddleware( max_requests_per_second=10.0, # 10 per second = 1 every 100ms - burst_capacity=3, + burst_capacity=4, ) )