Jeremiah Lowin commited on
Commit
f14a5c8
·
1 Parent(s): 987be7e

Allow passing custom middleware

Browse files
Files changed (1) hide show
  1. tests/server/test_http_middleware.py +219 -0
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"