Skip to content

Commit 99cd5e0

Browse files
committed
Move RequestBodyLimitMiddleware out of the Streamable HTTP manager module
The middleware and DEFAULT_MAX_REQUEST_BODY_SIZE are now used by the SSE transport and the OAuth routes as well, so they move next to the other shared HTTP request checks in mcp.server.transport_security. Both names remain importable from mcp.server.streamable_http_manager. The middleware's own unit tests move with it; no behaviour change.
1 parent f27c1ed commit 99cd5e0

10 files changed

Lines changed: 201 additions & 197 deletions

File tree

src/mcp/server/auth/routes.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,7 @@
1717
from mcp.server.auth.middleware.client_auth import ClientAuthenticator
1818
from mcp.server.auth.provider import OAuthAuthorizationServerProvider
1919
from mcp.server.auth.settings import ClientRegistrationOptions, RevocationOptions
20-
from mcp.server.streamable_http_manager import DEFAULT_MAX_REQUEST_BODY_SIZE, RequestBodyLimitMiddleware
20+
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE, RequestBodyLimitMiddleware
2121
from mcp.shared.auth import JWT_BEARER_GRANT_TYPE, OAuthMetadata, ProtectedResourceMetadata
2222
from mcp.shared.inbound import MCP_PROTOCOL_VERSION_HEADER
2323

src/mcp/server/lowlevel/server.py

Lines changed: 2 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -65,12 +65,8 @@ async def main():
6565
from mcp.server.models import InitializationOptions
6666
from mcp.server.runner import serve_dual_era_loop
6767
from mcp.server.streamable_http import EventStore
68-
from mcp.server.streamable_http_manager import (
69-
DEFAULT_MAX_REQUEST_BODY_SIZE,
70-
StreamableHTTPASGIApp,
71-
StreamableHTTPSessionManager,
72-
)
73-
from mcp.server.transport_security import TransportSecuritySettings
68+
from mcp.server.streamable_http_manager import StreamableHTTPASGIApp, StreamableHTTPSessionManager
69+
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE, TransportSecuritySettings
7470
from mcp.shared._stream_protocols import ReadStream, WriteStream
7571
from mcp.shared.exceptions import MCPDeprecationWarning
7672
from mcp.shared.message import SessionMessage

src/mcp/server/mcpserver/server.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -87,9 +87,9 @@
8787
from mcp.server.sse import SseServerTransport
8888
from mcp.server.stdio import stdio_server
8989
from mcp.server.streamable_http import EventStore
90-
from mcp.server.streamable_http_manager import DEFAULT_MAX_REQUEST_BODY_SIZE, StreamableHTTPSessionManager
90+
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
9191
from mcp.server.subscriptions import InMemorySubscriptionBus, ListenHandler, SubscriptionBus
92-
from mcp.server.transport_security import TransportSecuritySettings
92+
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE, TransportSecuritySettings
9393
from mcp.shared.exceptions import MCPError
9494
from mcp.shared.uri_template import UriTemplate
9595

src/mcp/server/sse.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -51,8 +51,9 @@ async def handle_sse(request):
5151
from starlette.types import Receive, Scope, Send
5252

5353
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser, AuthorizationContext, authorization_context
54-
from mcp.server.streamable_http_manager import DEFAULT_MAX_REQUEST_BODY_SIZE, RequestBodyLimitMiddleware
5554
from mcp.server.transport_security import (
55+
DEFAULT_MAX_REQUEST_BODY_SIZE,
56+
RequestBodyLimitMiddleware,
5657
TransportSecurityMiddleware,
5758
TransportSecuritySettings,
5859
)

src/mcp/server/streamable_http_manager.py

Lines changed: 4 additions & 67 deletions
Original file line numberDiff line numberDiff line change
@@ -4,25 +4,25 @@
44

55
import contextlib
66
import logging
7-
from collections import deque
87
from collections.abc import AsyncIterator
9-
from typing import TYPE_CHECKING, Any, Final
8+
from typing import TYPE_CHECKING, Any
109
from uuid import uuid4
1110

1211
import anyio
1312
from anyio.abc import TaskStatus
1413
from mcp_types import DEFAULT_NEGOTIATED_VERSION, INVALID_REQUEST, ErrorData, JSONRPCError
1514
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS
16-
from starlette.datastructures import Headers
1715
from starlette.requests import Request
1816
from starlette.responses import Response
19-
from starlette.types import ASGIApp, Message, Receive, Scope, Send
17+
from starlette.types import Receive, Scope, Send
2018

2119
from mcp.server._streamable_http_modern import handle_modern_request
2220
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser, AuthorizationContext, authorization_context
2321
from mcp.server.connection import Connection
2422
from mcp.server.runner import serve_connection, serve_loop
2523
from mcp.server.streamable_http import MCP_SESSION_ID_HEADER, EventStore, StreamableHTTPServerTransport
24+
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE as DEFAULT_MAX_REQUEST_BODY_SIZE
25+
from mcp.server.transport_security import RequestBodyLimitMiddleware as RequestBodyLimitMiddleware
2626
from mcp.server.transport_security import TransportSecuritySettings
2727
from mcp.shared._compat import resync_tracer
2828
from mcp.shared.inbound import MCP_PROTOCOL_VERSION_HEADER
@@ -34,9 +34,6 @@
3434

3535
logger = logging.getLogger(__name__)
3636

37-
DEFAULT_MAX_REQUEST_BODY_SIZE: Final = 4 * 1024 * 1024
38-
"""Default maximum HTTP request body size in bytes (4 MiB)."""
39-
4037

4138
class StreamableHTTPSessionManager:
4239
"""Manages StreamableHTTP sessions with optional resumability via event store.
@@ -371,66 +368,6 @@ async def run_server(*, task_status: TaskStatus[None] = anyio.TASK_STATUS_IGNORE
371368
await response(scope, receive, send)
372369

373370

374-
class RequestBodyLimitMiddleware:
375-
"""Reject oversized HTTP request bodies before invoking an ASGI application."""
376-
377-
def __init__(self, app: ASGIApp, max_body_size: int) -> None:
378-
self.app = app
379-
self.max_body_size = max_body_size
380-
381-
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
382-
if scope["type"] != "http":
383-
await self.app(scope, receive, send)
384-
return
385-
386-
headers = Headers(scope=scope)
387-
content_length = headers.get("content-length")
388-
if content_length is not None:
389-
try:
390-
declared_size = int(content_length)
391-
except ValueError:
392-
pass
393-
else:
394-
if declared_size > self.max_body_size:
395-
response = Response("Request body too large", status_code=413)
396-
return await response(scope, receive, send)
397-
398-
received_body = bytearray()
399-
received_request = False
400-
body_complete = False
401-
trailing_message: Message | None = None
402-
while True:
403-
message = await receive()
404-
if message["type"] != "http.request":
405-
trailing_message = message
406-
break
407-
408-
received_request = True
409-
body = message.get("body", b"")
410-
if len(received_body) + len(body) > self.max_body_size:
411-
response = Response("Request body too large", status_code=413)
412-
return await response(scope, receive, send)
413-
received_body.extend(body)
414-
if not message.get("more_body", False):
415-
body_complete = True
416-
break
417-
418-
cached_messages: deque[Message] = deque()
419-
if received_request:
420-
cached_messages.append(
421-
{"type": "http.request", "body": bytes(received_body), "more_body": not body_complete}
422-
)
423-
if trailing_message is not None:
424-
cached_messages.append(trailing_message)
425-
426-
async def replay() -> Message:
427-
if cached_messages:
428-
return cached_messages.popleft()
429-
return await receive()
430-
431-
await self.app(scope, replay, send)
432-
433-
434371
class StreamableHTTPASGIApp:
435372
"""ASGI application for Streamable HTTP server transport."""
436373

src/mcp/server/transport_security.py

Lines changed: 68 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,13 +1,20 @@
1-
"""DNS rebinding protection for MCP server transports."""
1+
"""Request checks shared by the HTTP server transports: Host/Origin header validation and body size limits."""
22

33
import logging
4+
from collections import deque
5+
from typing import Final
46

57
from pydantic import BaseModel, Field
8+
from starlette.datastructures import Headers
69
from starlette.requests import Request
710
from starlette.responses import Response
11+
from starlette.types import ASGIApp, Message, Receive, Scope, Send
812

913
logger = logging.getLogger(__name__)
1014

15+
DEFAULT_MAX_REQUEST_BODY_SIZE: Final = 4 * 1024 * 1024
16+
"""Default maximum HTTP request body size in bytes (4 MiB)."""
17+
1118

1219
# TODO(Marcelo): We should flatten these settings. To be fair, I don't think we should even have this middleware.
1320
class TransportSecuritySettings(BaseModel):
@@ -114,3 +121,63 @@ async def validate_request(self, request: Request, is_post: bool = False) -> Res
114121
return Response("Invalid Origin header", status_code=403)
115122

116123
return None
124+
125+
126+
class RequestBodyLimitMiddleware:
127+
"""Reject oversized HTTP request bodies before invoking an ASGI application."""
128+
129+
def __init__(self, app: ASGIApp, max_body_size: int) -> None:
130+
self.app = app
131+
self.max_body_size = max_body_size
132+
133+
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
134+
if scope["type"] != "http":
135+
await self.app(scope, receive, send)
136+
return
137+
138+
headers = Headers(scope=scope)
139+
content_length = headers.get("content-length")
140+
if content_length is not None:
141+
try:
142+
declared_size = int(content_length)
143+
except ValueError:
144+
pass
145+
else:
146+
if declared_size > self.max_body_size:
147+
response = Response("Request body too large", status_code=413)
148+
return await response(scope, receive, send)
149+
150+
received_body = bytearray()
151+
received_request = False
152+
body_complete = False
153+
trailing_message: Message | None = None
154+
while True:
155+
message = await receive()
156+
if message["type"] != "http.request":
157+
trailing_message = message
158+
break
159+
160+
received_request = True
161+
body = message.get("body", b"")
162+
if len(received_body) + len(body) > self.max_body_size:
163+
response = Response("Request body too large", status_code=413)
164+
return await response(scope, receive, send)
165+
received_body.extend(body)
166+
if not message.get("more_body", False):
167+
body_complete = True
168+
break
169+
170+
cached_messages: deque[Message] = deque()
171+
if received_request:
172+
cached_messages.append(
173+
{"type": "http.request", "body": bytes(received_body), "more_body": not body_complete}
174+
)
175+
if trailing_message is not None:
176+
cached_messages.append(trailing_message)
177+
178+
async def replay() -> Message:
179+
if cached_messages:
180+
return cached_messages.popleft()
181+
return await receive()
182+
183+
await self.app(scope, replay, send)

tests/server/auth/test_error_handling.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -16,7 +16,7 @@
1616
from mcp.server.auth.provider import AuthorizeError, RegistrationError, TokenError
1717
from mcp.server.auth.routes import create_auth_routes
1818
from mcp.server.auth.settings import ClientRegistrationOptions, RevocationOptions
19-
from mcp.server.streamable_http_manager import DEFAULT_MAX_REQUEST_BODY_SIZE
19+
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE
2020
from tests.server.mcpserver.auth.test_auth_integration import MockOAuthProvider
2121

2222

tests/server/test_sse_security.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -18,8 +18,7 @@
1818
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser
1919
from mcp.server.auth.provider import AccessToken
2020
from mcp.server.sse import SseServerTransport
21-
from mcp.server.streamable_http_manager import DEFAULT_MAX_REQUEST_BODY_SIZE
22-
from mcp.server.transport_security import TransportSecuritySettings
21+
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE, TransportSecuritySettings
2322
from mcp.shared._stream_protocols import WriteStream
2423
from mcp.shared.message import SessionMessage
2524
from tests.interaction.transports import StreamingASGITransport

tests/server/test_streamable_http_manager.py

Lines changed: 2 additions & 114 deletions
Original file line numberDiff line numberDiff line change
@@ -10,19 +10,15 @@
1010
import httpx2
1111
import pytest
1212
from mcp_types import INVALID_REQUEST, ListToolsResult, PaginatedRequestParams
13-
from starlette.types import Message, Receive, Scope, Send
13+
from starlette.types import Message, Scope
1414

1515
from mcp import Client
1616
from mcp.client.streamable_http import streamable_http_client
1717
from mcp.server import Server, ServerRequestContext, streamable_http_manager
1818
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser
1919
from mcp.server.auth.provider import AccessToken
2020
from mcp.server.streamable_http import MCP_SESSION_ID_HEADER, StreamableHTTPServerTransport
21-
from mcp.server.streamable_http_manager import (
22-
DEFAULT_MAX_REQUEST_BODY_SIZE,
23-
RequestBodyLimitMiddleware,
24-
StreamableHTTPSessionManager,
25-
)
21+
from mcp.server.streamable_http_manager import DEFAULT_MAX_REQUEST_BODY_SIZE, StreamableHTTPSessionManager
2622

2723

2824
@pytest.mark.anyio
@@ -146,114 +142,6 @@ async def send(message: Message) -> None:
146142
assert response_start["status"] == 413
147143

148144

149-
@pytest.mark.anyio
150-
async def test_client_disconnect_while_streaming_request_body_is_replayed() -> None:
151-
"""SDK-defined: raw ASGI is required to prove a disconnect before body completion reaches the transport."""
152-
disconnect: Message = {"type": "http.disconnect"}
153-
request_messages: Iterator[Message] = iter(
154-
[{"type": "http.request", "body": b"1234", "more_body": True}, disconnect]
155-
)
156-
received_messages: list[Message] = []
157-
158-
async def receive() -> Message:
159-
return next(request_messages)
160-
161-
async def app(scope: Scope, receive: Receive, send: Send) -> None:
162-
received_messages.append(await receive())
163-
received_messages.append(await receive())
164-
165-
scope: Scope = {"type": "http", "method": "POST", "path": "/mcp", "headers": []}
166-
middleware = RequestBodyLimitMiddleware(app, max_body_size=8)
167-
168-
await middleware(scope, receive, AsyncMock())
169-
170-
assert received_messages == [
171-
{"type": "http.request", "body": b"1234", "more_body": True},
172-
disconnect,
173-
]
174-
175-
176-
@pytest.mark.anyio
177-
async def test_client_disconnect_before_request_body_is_replayed() -> None:
178-
"""SDK-defined: raw ASGI proves a disconnect before the first body message reaches the transport."""
179-
disconnect: Message = {"type": "http.disconnect"}
180-
received_messages: list[Message] = []
181-
182-
async def receive() -> Message:
183-
return disconnect
184-
185-
async def app(scope: Scope, receive: Receive, send: Send) -> None:
186-
received_messages.append(await receive())
187-
188-
scope: Scope = {"type": "http", "method": "POST", "path": "/mcp", "headers": []}
189-
middleware = RequestBodyLimitMiddleware(app, max_body_size=8)
190-
191-
await middleware(scope, receive, AsyncMock())
192-
193-
assert received_messages == [disconnect]
194-
195-
196-
@pytest.mark.anyio
197-
async def test_request_body_chunks_are_replayed_as_one_message() -> None:
198-
"""SDK-defined: raw ASGI proves chunk overhead is discarded before the body reaches the transport."""
199-
request_messages: Iterator[Message] = iter(
200-
[
201-
{"type": "http.request", "body": b"12", "more_body": True},
202-
{"type": "http.request", "body": b"34", "more_body": True},
203-
{"type": "http.request", "body": b"56", "more_body": False},
204-
]
205-
)
206-
received_messages: list[Message] = []
207-
208-
async def receive() -> Message:
209-
return next(request_messages)
210-
211-
async def app(scope: Scope, receive: Receive, send: Send) -> None:
212-
received_messages.append(await receive())
213-
214-
scope: Scope = {"type": "http", "method": "POST", "path": "/mcp", "headers": []}
215-
middleware = RequestBodyLimitMiddleware(app, max_body_size=8)
216-
217-
await middleware(scope, receive, AsyncMock())
218-
219-
assert received_messages == [{"type": "http.request", "body": b"123456", "more_body": False}]
220-
221-
222-
@pytest.mark.anyio
223-
@pytest.mark.parametrize("method", ["GET", "PUT", "OPTIONS", "HEAD", "DELETE"])
224-
async def test_request_body_limit_applies_to_every_method(method: str) -> None:
225-
"""SDK-defined: the limit is a property of the request body, not of the method that carries it."""
226-
app = AsyncMock()
227-
sent_messages: list[Message] = []
228-
receive = AsyncMock(return_value={"type": "http.request", "body": b"123456789", "more_body": False})
229-
230-
async def send(message: Message) -> None:
231-
sent_messages.append(message)
232-
233-
scope: Scope = {"type": "http", "method": method, "path": "/mcp", "headers": []}
234-
middleware = RequestBodyLimitMiddleware(app, max_body_size=8)
235-
236-
await middleware(scope, receive, send)
237-
238-
assert [message["status"] for message in sent_messages if message["type"] == "http.response.start"] == [413]
239-
app.assert_not_awaited()
240-
241-
242-
@pytest.mark.anyio
243-
async def test_request_body_limit_leaves_non_http_scopes_alone() -> None:
244-
"""SDK-defined: only HTTP requests carry a body to limit; other ASGI scopes go straight to the app."""
245-
app = AsyncMock()
246-
receive = AsyncMock()
247-
send = AsyncMock()
248-
scope: Scope = {"type": "lifespan"}
249-
middleware = RequestBodyLimitMiddleware(app, max_body_size=8)
250-
251-
await middleware(scope, receive, send)
252-
253-
app.assert_awaited_once_with(scope, receive, send)
254-
receive.assert_not_awaited()
255-
256-
257145
def test_request_body_limit_defaults_to_four_mib() -> None:
258146
"""SDK-defined: Streamable HTTP request bodies are limited to 4 MiB by default."""
259147
manager = StreamableHTTPSessionManager(app=Server("test-default-size-limit"))

0 commit comments

Comments
 (0)