Jeremiah Lowin commited on
Commit
9663571
·
unverified ·
2 Parent(s): 75ee53f1f22024

Merge pull request #398 from jlowin/middleware

Browse files

Allow users to pass middleware to starlette app constructors

docs/deployment/asgi.mdx CHANGED
@@ -63,10 +63,34 @@ Or, from the command line:
63
  uvicorn path.to.your.app:http_app --host 0.0.0.0 --port 8000
64
  ```
65
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
66
 
67
 
68
  ## Starlette Integration
69
 
 
 
70
  You can mount your FastMCP server in another Starlette application using the `Mount` class.
71
 
72
  ```python
@@ -127,6 +151,8 @@ For Streamable HTTP transport, you **must** pass the lifespan context from the F
127
  </Warning>
128
  ## FastAPI Integration
129
 
 
 
130
  FastAPI is built on Starlette, so you can mount your FastMCP server in a similar way:
131
 
132
  ```python
 
63
  uvicorn path.to.your.app:http_app --host 0.0.0.0 --port 8000
64
  ```
65
 
66
+ ### Custom Middleware
67
+
68
+ <VersionBadge version="2.3.2" />
69
+
70
+ You can add custom Starlette middleware to your FastMCP ASGI apps by passing a list of middleware instances to the app creation methods:
71
+
72
+ ```python
73
+ from fastmcp import FastMCP
74
+ from starlette.middleware import Middleware
75
+ from starlette.middleware.cors import CORSMiddleware
76
+
77
+ # Create your FastMCP server
78
+ mcp = FastMCP("MyServer")
79
+
80
+ # Define custom middleware
81
+ custom_middleware = [
82
+ Middleware(CORSMiddleware, allow_origins=["*"]),
83
+ ]
84
+
85
+ # Create ASGI app with custom middleware
86
+ http_app = mcp.streamable_http_app(middleware=custom_middleware)
87
+ ```
88
 
89
 
90
  ## Starlette Integration
91
 
92
+ <VersionBadge version="2.3.1" />
93
+
94
  You can mount your FastMCP server in another Starlette application using the `Mount` class.
95
 
96
  ```python
 
151
  </Warning>
152
  ## FastAPI Integration
153
 
154
+ <VersionBadge version="2.3.1" />
155
+
156
  FastAPI is built on Starlette, so you can mount your FastMCP server in a similar way:
157
 
158
  ```python
src/fastmcp/server/http.py CHANGED
@@ -3,7 +3,7 @@ from __future__ import annotations
3
  from collections.abc import AsyncGenerator, Callable, Generator
4
  from contextlib import asynccontextmanager, contextmanager
5
  from contextvars import ContextVar
6
- from typing import TYPE_CHECKING, cast
7
 
8
  from mcp.server.auth.middleware.auth_context import AuthContextMiddleware
9
  from mcp.server.auth.middleware.bearer_auth import (
@@ -19,7 +19,7 @@ from starlette.middleware import Middleware
19
  from starlette.middleware.authentication import AuthenticationMiddleware
20
  from starlette.requests import Request
21
  from starlette.responses import Response
22
- from starlette.routing import Mount, Route
23
  from starlette.types import Receive, Scope, Send
24
 
25
  from fastmcp.low_level.sse_server_transport import SseServerTransport
@@ -64,7 +64,7 @@ class RequestContextMiddleware:
64
  def setup_auth_middleware_and_routes(
65
  auth_server_provider: OAuthAuthorizationServerProvider | None,
66
  auth_settings: AuthSettings | None,
67
- ) -> tuple[list[Middleware], list[Route | Mount], list[str]]:
68
  """Set up authentication middleware and routes if auth is enabled.
69
 
70
  Args:
@@ -75,7 +75,7 @@ def setup_auth_middleware_and_routes(
75
  Tuple of (middleware, auth_routes, required_scopes)
76
  """
77
  middleware: list[Middleware] = []
78
- auth_routes: list[Route | Mount] = []
79
  required_scopes: list[str] = []
80
 
81
  if auth_server_provider:
@@ -108,7 +108,7 @@ def setup_auth_middleware_and_routes(
108
 
109
 
110
  def create_base_app(
111
- routes: list[Route | Mount],
112
  middleware: list[Middleware],
113
  debug: bool = False,
114
  lifespan: Callable | None = None,
@@ -142,7 +142,8 @@ def create_sse_app(
142
  auth_server_provider: OAuthAuthorizationServerProvider | None = None,
143
  auth_settings: AuthSettings | None = None,
144
  debug: bool = False,
145
- additional_routes: list[Route] | list[Mount] | list[Route | Mount] | None = None,
 
146
  ) -> Starlette:
147
  """Return an instance of the SSE server app.
148
 
@@ -153,11 +154,15 @@ def create_sse_app(
153
  auth_server_provider: Optional auth provider
154
  auth_settings: Optional auth settings
155
  debug: Whether to enable debug mode
156
- additional_routes: Optional list of custom routes
157
-
158
  Returns:
159
  A Starlette application with RequestContextMiddleware
160
  """
 
 
 
 
161
  # Set up SSE transport
162
  sse = SseServerTransport(message_path)
163
 
@@ -172,24 +177,24 @@ def create_sse_app(
172
  return Response()
173
 
174
  # Get auth middleware and routes
175
- middleware, auth_routes, required_scopes = setup_auth_middleware_and_routes(
176
  auth_server_provider, auth_settings
177
  )
178
 
179
- # Initialize routes with auth routes
180
- routes: list[Route | Mount] = auth_routes.copy()
181
 
182
  # Add SSE routes with or without auth
183
  if auth_server_provider:
184
  # Auth is enabled, wrap endpoints with RequireAuthMiddleware
185
- routes.append(
186
  Route(
187
  sse_path,
188
  endpoint=RequireAuthMiddleware(handle_sse, required_scopes),
189
  methods=["GET"],
190
  )
191
  )
192
- routes.append(
193
  Mount(
194
  message_path,
195
  app=RequireAuthMiddleware(sse.handle_post_message, required_scopes),
@@ -200,14 +205,14 @@ def create_sse_app(
200
  async def sse_endpoint(request: Request) -> Response:
201
  return await handle_sse(request.scope, request.receive, request._send) # type: ignore[reportPrivateUsage]
202
 
203
- routes.append(
204
  Route(
205
  sse_path,
206
  endpoint=sse_endpoint,
207
  methods=["GET"],
208
  )
209
  )
210
- routes.append(
211
  Mount(
212
  message_path,
213
  app=sse.handle_post_message,
@@ -215,13 +220,17 @@ def create_sse_app(
215
  )
216
 
217
  # Add custom routes with lowest precedence
218
- if additional_routes:
219
- routes.extend(cast(list[Route | Mount], additional_routes))
 
 
 
 
220
 
221
  # Create and return the app
222
  return create_base_app(
223
- routes=routes,
224
- middleware=middleware,
225
  debug=debug,
226
  )
227
 
@@ -235,7 +244,8 @@ def create_streamable_http_app(
235
  json_response: bool = False,
236
  stateless_http: bool = False,
237
  debug: bool = False,
238
- additional_routes: list[Route] | list[Mount] | list[Route | Mount] | None = None,
 
239
  ) -> Starlette:
240
  """Return an instance of the StreamableHTTP server app.
241
 
@@ -248,11 +258,15 @@ def create_streamable_http_app(
248
  json_response: Whether to use JSON response format
249
  stateless_http: Whether to use stateless mode (new transport per request)
250
  debug: Whether to enable debug mode
251
- additional_routes: Optional list of custom routes
 
252
 
253
  Returns:
254
  A Starlette application with StreamableHTTP support
255
  """
 
 
 
256
  # Create session manager using the provided event store
257
  session_manager = StreamableHTTPSessionManager(
258
  app=server._mcp_server,
@@ -268,17 +282,17 @@ def create_streamable_http_app(
268
  await session_manager.handle_request(scope, receive, send)
269
 
270
  # Get auth middleware and routes
271
- middleware, auth_routes, required_scopes = setup_auth_middleware_and_routes(
272
  auth_server_provider, auth_settings
273
  )
274
 
275
- # Initialize routes with auth routes
276
- routes: list[Route | Mount] = auth_routes.copy()
277
 
278
  # Add StreamableHTTP routes with or without auth
279
  if auth_server_provider:
280
  # Auth is enabled, wrap endpoint with RequireAuthMiddleware
281
- routes.append(
282
  Mount(
283
  streamable_http_path,
284
  app=RequireAuthMiddleware(handle_streamable_http, required_scopes),
@@ -286,7 +300,7 @@ def create_streamable_http_app(
286
  )
287
  else:
288
  # No auth required
289
- routes.append(
290
  Mount(
291
  streamable_http_path,
292
  app=handle_streamable_http,
@@ -294,8 +308,12 @@ def create_streamable_http_app(
294
  )
295
 
296
  # Add custom routes with lowest precedence
297
- if additional_routes:
298
- routes.extend(cast(list[Route | Mount], additional_routes))
 
 
 
 
299
 
300
  # Create a lifespan manager to start and stop the session manager
301
  @asynccontextmanager
@@ -305,8 +323,8 @@ def create_streamable_http_app(
305
 
306
  # Create and return the app with lifespan
307
  return create_base_app(
308
- routes=routes,
309
- middleware=middleware,
310
  debug=debug,
311
  lifespan=lifespan,
312
  )
 
3
  from collections.abc import AsyncGenerator, Callable, Generator
4
  from contextlib import asynccontextmanager, contextmanager
5
  from contextvars import ContextVar
6
+ from typing import TYPE_CHECKING
7
 
8
  from mcp.server.auth.middleware.auth_context import AuthContextMiddleware
9
  from mcp.server.auth.middleware.bearer_auth import (
 
19
  from starlette.middleware.authentication import AuthenticationMiddleware
20
  from starlette.requests import Request
21
  from starlette.responses import Response
22
+ from starlette.routing import BaseRoute, Mount, Route
23
  from starlette.types import Receive, Scope, Send
24
 
25
  from fastmcp.low_level.sse_server_transport import SseServerTransport
 
64
  def setup_auth_middleware_and_routes(
65
  auth_server_provider: OAuthAuthorizationServerProvider | None,
66
  auth_settings: AuthSettings | None,
67
+ ) -> tuple[list[Middleware], list[BaseRoute], list[str]]:
68
  """Set up authentication middleware and routes if auth is enabled.
69
 
70
  Args:
 
75
  Tuple of (middleware, auth_routes, required_scopes)
76
  """
77
  middleware: list[Middleware] = []
78
+ auth_routes: list[BaseRoute] = []
79
  required_scopes: list[str] = []
80
 
81
  if auth_server_provider:
 
108
 
109
 
110
  def create_base_app(
111
+ routes: list[BaseRoute],
112
  middleware: list[Middleware],
113
  debug: bool = False,
114
  lifespan: Callable | None = None,
 
142
  auth_server_provider: OAuthAuthorizationServerProvider | None = None,
143
  auth_settings: AuthSettings | None = None,
144
  debug: bool = False,
145
+ routes: list[BaseRoute] | None = None,
146
+ middleware: list[Middleware] | None = None,
147
  ) -> Starlette:
148
  """Return an instance of the SSE server app.
149
 
 
154
  auth_server_provider: Optional auth provider
155
  auth_settings: Optional auth settings
156
  debug: Whether to enable debug mode
157
+ routes: Optional list of custom routes
158
+ middleware: Optional list of middleware
159
  Returns:
160
  A Starlette application with RequestContextMiddleware
161
  """
162
+
163
+ server_routes: list[BaseRoute] = []
164
+ server_middleware: list[Middleware] = []
165
+
166
  # Set up SSE transport
167
  sse = SseServerTransport(message_path)
168
 
 
177
  return Response()
178
 
179
  # Get auth middleware and routes
180
+ auth_middleware, auth_routes, required_scopes = setup_auth_middleware_and_routes(
181
  auth_server_provider, auth_settings
182
  )
183
 
184
+ server_routes.extend(auth_routes)
185
+ server_middleware.extend(auth_middleware)
186
 
187
  # Add SSE routes with or without auth
188
  if auth_server_provider:
189
  # Auth is enabled, wrap endpoints with RequireAuthMiddleware
190
+ server_routes.append(
191
  Route(
192
  sse_path,
193
  endpoint=RequireAuthMiddleware(handle_sse, required_scopes),
194
  methods=["GET"],
195
  )
196
  )
197
+ server_routes.append(
198
  Mount(
199
  message_path,
200
  app=RequireAuthMiddleware(sse.handle_post_message, required_scopes),
 
205
  async def sse_endpoint(request: Request) -> Response:
206
  return await handle_sse(request.scope, request.receive, request._send) # type: ignore[reportPrivateUsage]
207
 
208
+ server_routes.append(
209
  Route(
210
  sse_path,
211
  endpoint=sse_endpoint,
212
  methods=["GET"],
213
  )
214
  )
215
+ server_routes.append(
216
  Mount(
217
  message_path,
218
  app=sse.handle_post_message,
 
220
  )
221
 
222
  # Add custom routes with lowest precedence
223
+ if routes:
224
+ server_routes.extend(routes)
225
+
226
+ # Add middleware
227
+ if middleware:
228
+ server_middleware.extend(middleware)
229
 
230
  # Create and return the app
231
  return create_base_app(
232
+ routes=server_routes,
233
+ middleware=server_middleware,
234
  debug=debug,
235
  )
236
 
 
244
  json_response: bool = False,
245
  stateless_http: bool = False,
246
  debug: bool = False,
247
+ routes: list[BaseRoute] | None = None,
248
+ middleware: list[Middleware] | None = None,
249
  ) -> Starlette:
250
  """Return an instance of the StreamableHTTP server app.
251
 
 
258
  json_response: Whether to use JSON response format
259
  stateless_http: Whether to use stateless mode (new transport per request)
260
  debug: Whether to enable debug mode
261
+ routes: Optional list of custom routes
262
+ middleware: Optional list of middleware
263
 
264
  Returns:
265
  A Starlette application with StreamableHTTP support
266
  """
267
+ server_routes: list[BaseRoute] = []
268
+ server_middleware: list[Middleware] = []
269
+
270
  # Create session manager using the provided event store
271
  session_manager = StreamableHTTPSessionManager(
272
  app=server._mcp_server,
 
282
  await session_manager.handle_request(scope, receive, send)
283
 
284
  # Get auth middleware and routes
285
+ auth_middleware, auth_routes, required_scopes = setup_auth_middleware_and_routes(
286
  auth_server_provider, auth_settings
287
  )
288
 
289
+ server_routes.extend(auth_routes)
290
+ server_middleware.extend(auth_middleware)
291
 
292
  # Add StreamableHTTP routes with or without auth
293
  if auth_server_provider:
294
  # Auth is enabled, wrap endpoint with RequireAuthMiddleware
295
+ server_routes.append(
296
  Mount(
297
  streamable_http_path,
298
  app=RequireAuthMiddleware(handle_streamable_http, required_scopes),
 
300
  )
301
  else:
302
  # No auth required
303
+ server_routes.append(
304
  Mount(
305
  streamable_http_path,
306
  app=handle_streamable_http,
 
308
  )
309
 
310
  # Add custom routes with lowest precedence
311
+ if routes:
312
+ server_routes.extend(routes)
313
+
314
+ # Add middleware
315
+ if middleware:
316
+ server_middleware.extend(middleware)
317
 
318
  # Create a lifespan manager to start and stop the session manager
319
  @asynccontextmanager
 
323
 
324
  # Create and return the app with lifespan
325
  return create_base_app(
326
+ routes=server_routes,
327
+ middleware=server_middleware,
328
  debug=debug,
329
  lifespan=lifespan,
330
  )
src/fastmcp/server/server.py CHANGED
@@ -35,9 +35,10 @@ from mcp.types import ResourceTemplate as MCPResourceTemplate
35
  from mcp.types import Tool as MCPTool
36
  from pydantic import AnyUrl
37
  from starlette.applications import Starlette
 
38
  from starlette.requests import Request
39
  from starlette.responses import Response
40
- from starlette.routing import Route
41
 
42
  import fastmcp.server
43
  import fastmcp.settings
@@ -147,7 +148,7 @@ class FastMCP(Generic[LifespanResultT]):
147
  )
148
  self._auth_server_provider = auth_server_provider
149
 
150
- self._additional_http_routes: list[Route] = []
151
  self.dependencies = self.settings.dependencies
152
 
153
  # Set up MCP protocol handlers
@@ -744,8 +745,16 @@ class FastMCP(Generic[LifespanResultT]):
744
  self,
745
  path: str | None = None,
746
  message_path: str | None = None,
 
747
  ) -> Starlette:
748
- """Return an instance of the SSE server app."""
 
 
 
 
 
 
 
749
  return create_sse_app(
750
  server=self,
751
  message_path=message_path or self.settings.message_path,
@@ -753,11 +762,22 @@ class FastMCP(Generic[LifespanResultT]):
753
  auth_server_provider=self._auth_server_provider,
754
  auth_settings=self.settings.auth,
755
  debug=self.settings.debug,
756
- additional_routes=self._additional_http_routes,
 
757
  )
758
 
759
- def streamable_http_app(self, path: str | None = None) -> Starlette:
760
- """Return an instance of the StreamableHTTP server app."""
 
 
 
 
 
 
 
 
 
 
761
  from fastmcp.server.http import create_streamable_http_app
762
 
763
  return create_streamable_http_app(
@@ -769,7 +789,8 @@ class FastMCP(Generic[LifespanResultT]):
769
  json_response=self.settings.json_response,
770
  stateless_http=self.settings.stateless_http,
771
  debug=self.settings.debug,
772
- additional_routes=self._additional_http_routes,
 
773
  )
774
 
775
  async def run_streamable_http_async(
 
35
  from mcp.types import Tool as MCPTool
36
  from pydantic import AnyUrl
37
  from starlette.applications import Starlette
38
+ from starlette.middleware import Middleware
39
  from starlette.requests import Request
40
  from starlette.responses import Response
41
+ from starlette.routing import BaseRoute, Route
42
 
43
  import fastmcp.server
44
  import fastmcp.settings
 
148
  )
149
  self._auth_server_provider = auth_server_provider
150
 
151
+ self._additional_http_routes: list[BaseRoute] = []
152
  self.dependencies = self.settings.dependencies
153
 
154
  # Set up MCP protocol handlers
 
745
  self,
746
  path: str | None = None,
747
  message_path: str | None = None,
748
+ middleware: list[Middleware] | None = None,
749
  ) -> Starlette:
750
+ """
751
+ Create a Starlette app for the SSE server.
752
+
753
+ Args:
754
+ path: The path to the SSE endpoint
755
+ message_path: The path to the message endpoint
756
+ middleware: A list of middleware to apply to the app
757
+ """
758
  return create_sse_app(
759
  server=self,
760
  message_path=message_path or self.settings.message_path,
 
762
  auth_server_provider=self._auth_server_provider,
763
  auth_settings=self.settings.auth,
764
  debug=self.settings.debug,
765
+ routes=self._additional_http_routes,
766
+ middleware=middleware,
767
  )
768
 
769
+ def streamable_http_app(
770
+ self,
771
+ path: str | None = None,
772
+ middleware: list[Middleware] | None = None,
773
+ ) -> Starlette:
774
+ """
775
+ Create a Starlette app for the StreamableHTTP server.
776
+
777
+ Args:
778
+ path: The path to the StreamableHTTP endpoint
779
+ middleware: A list of middleware to apply to the app
780
+ """
781
  from fastmcp.server.http import create_streamable_http_app
782
 
783
  return create_streamable_http_app(
 
789
  json_response=self.settings.json_response,
790
  stateless_http=self.settings.stateless_http,
791
  debug=self.settings.debug,
792
+ routes=self._additional_http_routes,
793
+ middleware=middleware,
794
  )
795
 
796
  async def run_streamable_http_async(
test.py DELETED
@@ -1,12 +0,0 @@
1
- from fastmcp import FastMCP
2
-
3
- mcp = FastMCP()
4
-
5
- if __name__ == "__main__":
6
- mcp.run(
7
- transport="streamable-http",
8
- host="127.0.0.1",
9
- port=4200,
10
- path="/my-custom-path/",
11
- log_level="debug",
12
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
tests/server/test_http_middleware.py ADDED
@@ -0,0 +1,219 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Tests for custom middleware in HTTP servers."""
2
+
3
+ from collections.abc import Callable
4
+ from typing import Any
5
+
6
+ import httpx
7
+ import pytest
8
+ from httpx import ASGITransport
9
+ from starlette.middleware import Middleware
10
+ from starlette.middleware.base import BaseHTTPMiddleware
11
+ from starlette.requests import Request
12
+ from starlette.responses import JSONResponse
13
+ from starlette.routing import BaseRoute, Route
14
+ from starlette.types import ASGIApp
15
+
16
+ from fastmcp.server import FastMCP
17
+ from fastmcp.server.http import create_sse_app, create_streamable_http_app
18
+
19
+
20
+ class HeaderMiddleware(BaseHTTPMiddleware):
21
+ """Simple middleware that adds a custom header to responses."""
22
+
23
+ def __init__(self, app: ASGIApp, header_name: str, header_value: str):
24
+ super().__init__(app)
25
+ self.header_name = header_name
26
+ self.header_value = header_value
27
+
28
+ async def dispatch(self, request: Request, call_next: Callable):
29
+ response = await call_next(request)
30
+ response.headers[self.header_name] = self.header_value
31
+ return response
32
+
33
+
34
+ class RequestModifierMiddleware(BaseHTTPMiddleware):
35
+ """Middleware that adds a value to request state."""
36
+
37
+ def __init__(self, app: ASGIApp, key: str, value: Any):
38
+ super().__init__(app)
39
+ self.key = key
40
+ self.value = value
41
+
42
+ async def dispatch(self, request: Request, call_next: Callable):
43
+ request.state.custom_value = {self.key: self.value}
44
+ return await call_next(request)
45
+
46
+
47
+ async def endpoint_handler(request: Request):
48
+ """Endpoint that returns request state or headers."""
49
+ if hasattr(request.state, "custom_value"):
50
+ return JSONResponse({"state": request.state.custom_value})
51
+ return JSONResponse({"message": "Hello, world!"})
52
+
53
+
54
+ @pytest.mark.asyncio
55
+ async def test_sse_app_with_custom_middleware():
56
+ """Test that custom middleware works with SSE app."""
57
+ server = FastMCP(name="TestServer")
58
+
59
+ # Create custom middleware
60
+ custom_middleware = [
61
+ Middleware(
62
+ HeaderMiddleware, header_name="X-Custom-Header", header_value="test-value"
63
+ )
64
+ ]
65
+
66
+ # Add a test route to server's additional routes
67
+ routes: list[BaseRoute] = [Route("/test", endpoint_handler)]
68
+ server._additional_http_routes = routes
69
+
70
+ # Create the app with custom middleware
71
+ app = server.sse_app(middleware=custom_middleware)
72
+
73
+ # Create a test client
74
+ transport = ASGITransport(app=app)
75
+ async with httpx.AsyncClient(
76
+ transport=transport, base_url="http://testserver"
77
+ ) as client:
78
+ response = await client.get("/test")
79
+
80
+ # Verify middleware was applied
81
+ assert response.status_code == 200
82
+ assert response.headers["X-Custom-Header"] == "test-value"
83
+
84
+
85
+ @pytest.mark.asyncio
86
+ async def test_streamable_http_app_with_custom_middleware():
87
+ """Test that custom middleware works with StreamableHTTP app."""
88
+ server = FastMCP(name="TestServer")
89
+
90
+ # Create custom middleware
91
+ custom_middleware = [
92
+ Middleware(
93
+ HeaderMiddleware, header_name="X-Custom-Header", header_value="test-value"
94
+ )
95
+ ]
96
+
97
+ # Add a test route to server's additional routes
98
+ routes: list[BaseRoute] = [Route("/test", endpoint_handler)]
99
+ server._additional_http_routes = routes
100
+
101
+ # Create the app with custom middleware
102
+ app = server.streamable_http_app(middleware=custom_middleware)
103
+
104
+ # Create a test client
105
+ transport = ASGITransport(app=app)
106
+ async with httpx.AsyncClient(
107
+ transport=transport, base_url="http://testserver"
108
+ ) as client:
109
+ response = await client.get("/test")
110
+
111
+ # Verify middleware was applied
112
+ assert response.status_code == 200
113
+ assert response.headers["X-Custom-Header"] == "test-value"
114
+
115
+
116
+ @pytest.mark.asyncio
117
+ async def test_create_sse_app_with_custom_middleware():
118
+ """Test that custom middleware works with create_sse_app function."""
119
+ server = FastMCP(name="TestServer")
120
+
121
+ # Create custom middleware
122
+ custom_middleware = [
123
+ Middleware(RequestModifierMiddleware, key="modified_by", value="middleware")
124
+ ]
125
+
126
+ # Add a test route
127
+ additional_routes: list[BaseRoute] = [Route("/test", endpoint_handler)]
128
+
129
+ # Create the app with custom middleware
130
+ app = create_sse_app(
131
+ server=server,
132
+ message_path="/message",
133
+ sse_path="/sse",
134
+ middleware=custom_middleware,
135
+ routes=additional_routes,
136
+ )
137
+
138
+ # Create a test client
139
+ transport = ASGITransport(app=app)
140
+ async with httpx.AsyncClient(
141
+ transport=transport, base_url="http://testserver"
142
+ ) as client:
143
+ response = await client.get("/test")
144
+
145
+ # Verify middleware was applied
146
+ assert response.status_code == 200
147
+ data = response.json()
148
+ assert "state" in data
149
+ assert data["state"]["modified_by"] == "middleware"
150
+
151
+
152
+ @pytest.mark.asyncio
153
+ async def test_create_streamable_http_app_with_custom_middleware():
154
+ """Test that custom middleware works with create_streamable_http_app function."""
155
+ server = FastMCP(name="TestServer")
156
+
157
+ # Create custom middleware
158
+ custom_middleware = [
159
+ Middleware(RequestModifierMiddleware, key="modified_by", value="middleware")
160
+ ]
161
+
162
+ # Add a test route
163
+ additional_routes: list[BaseRoute] = [Route("/test", endpoint_handler)]
164
+
165
+ # Create the app with custom middleware
166
+ app = create_streamable_http_app(
167
+ server=server,
168
+ streamable_http_path="/streamable",
169
+ middleware=custom_middleware,
170
+ routes=additional_routes,
171
+ )
172
+
173
+ # Create a test client
174
+ transport = ASGITransport(app=app)
175
+ async with httpx.AsyncClient(
176
+ transport=transport, base_url="http://testserver"
177
+ ) as client:
178
+ response = await client.get("/test")
179
+
180
+ # Verify middleware was applied
181
+ assert response.status_code == 200
182
+ data = response.json()
183
+ assert "state" in data
184
+ assert data["state"]["modified_by"] == "middleware"
185
+
186
+
187
+ @pytest.mark.asyncio
188
+ async def test_multiple_middleware_ordering():
189
+ """Test that multiple middleware are applied in the correct order."""
190
+ server = FastMCP(name="TestServer")
191
+
192
+ # Create multiple middleware
193
+ custom_middleware = [
194
+ Middleware(
195
+ HeaderMiddleware, header_name="X-First-Header", header_value="first"
196
+ ),
197
+ Middleware(
198
+ HeaderMiddleware, header_name="X-Second-Header", header_value="second"
199
+ ),
200
+ ]
201
+
202
+ # Add a test route to server's additional routes
203
+ routes: list[BaseRoute] = [Route("/test", endpoint_handler)]
204
+ server._additional_http_routes = routes
205
+
206
+ # Create the app with custom middleware
207
+ app = server.sse_app(middleware=custom_middleware)
208
+
209
+ # Create a test client
210
+ transport = ASGITransport(app=app)
211
+ async with httpx.AsyncClient(
212
+ transport=transport, base_url="http://testserver"
213
+ ) as client:
214
+ response = await client.get("/test")
215
+
216
+ # Verify both middleware were applied
217
+ assert response.status_code == 200
218
+ assert response.headers["X-First-Header"] == "first"
219
+ assert response.headers["X-Second-Header"] == "second"