Skip to content

Commit d838cd3

Browse files
committed
[v1.x] 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 (listed in its __all__). The middleware's own unit tests move with it, plus the mid-stream disconnect replay case main already carries; no behaviour change.
1 parent a79c0f3 commit d838cd3

9 files changed

Lines changed: 207 additions & 164 deletions

File tree

src/mcp/server/auth/routes.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818
from mcp.server.auth.provider import OAuthAuthorizationServerProvider
1919
from mcp.server.auth.settings import ClientRegistrationOptions, RevocationOptions
2020
from mcp.server.streamable_http import MCP_PROTOCOL_VERSION_HEADER
21-
from mcp.server.streamable_http_manager import DEFAULT_MAX_REQUEST_BODY_SIZE, RequestBodyLimitMiddleware
21+
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE, RequestBodyLimitMiddleware
2222
from mcp.shared.auth import OAuthMetadata
2323

2424

src/mcp/server/fastmcp/server.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -62,8 +62,8 @@
6262
from mcp.server.sse import SseServerTransport
6363
from mcp.server.stdio import stdio_server
6464
from mcp.server.streamable_http import EventStore
65-
from mcp.server.streamable_http_manager import DEFAULT_MAX_REQUEST_BODY_SIZE, StreamableHTTPSessionManager
66-
from mcp.server.transport_security import TransportSecuritySettings
65+
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
66+
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE, TransportSecuritySettings
6767
from mcp.shared.context import LifespanContextT, RequestContext, RequestT
6868
from mcp.types import Annotations, AnyFunction, ContentBlock, GetPromptResult, Icon, ToolAnnotations
6969
from mcp.types import Prompt as MCPPrompt

src/mcp/server/sse.py

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -53,8 +53,9 @@ async def handle_sse(request):
5353

5454
import mcp.types as types
5555
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser, AuthorizationContext, authorization_context
56-
from mcp.server.streamable_http_manager import DEFAULT_MAX_REQUEST_BODY_SIZE, RequestBodyLimitMiddleware
5756
from mcp.server.transport_security import (
57+
DEFAULT_MAX_REQUEST_BODY_SIZE,
58+
RequestBodyLimitMiddleware,
5859
TransportSecurityMiddleware,
5960
TransportSecuritySettings,
6061
)

src/mcp/server/streamable_http_manager.py

Lines changed: 9 additions & 68 deletions
Original file line numberDiff line numberDiff line change
@@ -4,17 +4,15 @@
44

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

1211
import anyio
1312
from anyio.abc import TaskStatus
14-
from starlette.datastructures import Headers
1513
from starlette.requests import Request
1614
from starlette.responses import Response
17-
from starlette.types import ASGIApp, Message, Receive, Scope, Send
15+
from starlette.types import Receive, Scope, Send
1816

1917
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser, AuthorizationContext, authorization_context
2018
from mcp.server.lowlevel.server import Server as MCPServer
@@ -23,13 +21,16 @@
2321
EventStore,
2422
StreamableHTTPServerTransport,
2523
)
26-
from mcp.server.transport_security import TransportSecuritySettings
24+
from mcp.server.transport_security import (
25+
DEFAULT_MAX_REQUEST_BODY_SIZE,
26+
RequestBodyLimitMiddleware,
27+
TransportSecuritySettings,
28+
)
2729
from mcp.types import INVALID_REQUEST, ErrorData, JSONRPCError
2830

29-
logger = logging.getLogger(__name__)
31+
__all__ = ["DEFAULT_MAX_REQUEST_BODY_SIZE", "RequestBodyLimitMiddleware", "StreamableHTTPSessionManager"]
3032

31-
DEFAULT_MAX_REQUEST_BODY_SIZE: Final = 4 * 1024 * 1024
32-
"""Default maximum HTTP request body size in bytes (4 MiB)."""
33+
logger = logging.getLogger(__name__)
3334

3435

3536
class StreamableHTTPSessionManager:
@@ -361,63 +362,3 @@ async def run_server(*, task_status: TaskStatus[None] = anyio.TASK_STATUS_IGNORE
361362
body.model_dump_json(by_alias=True, exclude_none=True), status_code=404, media_type="application/json"
362363
)
363364
await response(scope, receive, send)
364-
365-
366-
class RequestBodyLimitMiddleware:
367-
"""Reject oversized HTTP request bodies before invoking an ASGI application."""
368-
369-
def __init__(self, app: ASGIApp, max_body_size: int) -> None:
370-
self.app = app
371-
self.max_body_size = max_body_size
372-
373-
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
374-
if scope["type"] != "http":
375-
await self.app(scope, receive, send)
376-
return
377-
378-
headers = Headers(scope=scope)
379-
content_length = headers.get("content-length")
380-
if content_length is not None:
381-
try:
382-
declared_size = int(content_length)
383-
except ValueError:
384-
pass
385-
else:
386-
if declared_size > self.max_body_size:
387-
response = Response("Request body too large", status_code=413)
388-
return await response(scope, receive, send)
389-
390-
received_body = bytearray()
391-
received_request = False
392-
body_complete = False
393-
trailing_message: Message | None = None
394-
while True:
395-
message = await receive()
396-
if message["type"] != "http.request":
397-
trailing_message = message
398-
break
399-
400-
received_request = True
401-
body = message.get("body", b"")
402-
if len(received_body) + len(body) > self.max_body_size:
403-
response = Response("Request body too large", status_code=413)
404-
return await response(scope, receive, send)
405-
received_body.extend(body)
406-
if not message.get("more_body", False):
407-
body_complete = True
408-
break
409-
410-
cached_messages: deque[Message] = deque()
411-
if received_request:
412-
cached_messages.append(
413-
{"type": "http.request", "body": bytes(received_body), "more_body": not body_complete}
414-
)
415-
if trailing_message is not None:
416-
cached_messages.append(trailing_message)
417-
418-
async def replay() -> Message:
419-
if cached_messages:
420-
return cached_messages.popleft()
421-
return await receive()
422-
423-
await self.app(scope, replay, send)

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 HTTPConnection
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
class TransportSecuritySettings(BaseModel):
1320
"""Settings for MCP transport security features.
@@ -125,3 +132,63 @@ async def validate_request(self, request: HTTPConnection, is_post: bool = False)
125132
return Response("Invalid Origin header", status_code=403) # pragma: no cover
126133

127134
return None # pragma: no cover
135+
136+
137+
class RequestBodyLimitMiddleware:
138+
"""Reject oversized HTTP request bodies before invoking an ASGI application."""
139+
140+
def __init__(self, app: ASGIApp, max_body_size: int) -> None:
141+
self.app = app
142+
self.max_body_size = max_body_size
143+
144+
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
145+
if scope["type"] != "http":
146+
await self.app(scope, receive, send)
147+
return
148+
149+
headers = Headers(scope=scope)
150+
content_length = headers.get("content-length")
151+
if content_length is not None:
152+
try:
153+
declared_size = int(content_length)
154+
except ValueError:
155+
pass
156+
else:
157+
if declared_size > self.max_body_size:
158+
response = Response("Request body too large", status_code=413)
159+
return await response(scope, receive, send)
160+
161+
received_body = bytearray()
162+
received_request = False
163+
body_complete = False
164+
trailing_message: Message | None = None
165+
while True:
166+
message = await receive()
167+
if message["type"] != "http.request":
168+
trailing_message = message
169+
break
170+
171+
received_request = True
172+
body = message.get("body", b"")
173+
if len(received_body) + len(body) > self.max_body_size:
174+
response = Response("Request body too large", status_code=413)
175+
return await response(scope, receive, send)
176+
received_body.extend(body)
177+
if not message.get("more_body", False):
178+
body_complete = True
179+
break
180+
181+
cached_messages: deque[Message] = deque()
182+
if received_request:
183+
cached_messages.append(
184+
{"type": "http.request", "body": bytes(received_body), "more_body": not body_complete}
185+
)
186+
if trailing_message is not None:
187+
cached_messages.append(trailing_message)
188+
189+
async def replay() -> Message:
190+
if cached_messages:
191+
return cached_messages.popleft()
192+
return await receive()
193+
194+
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
@@ -14,7 +14,7 @@
1414

1515
from mcp.server.auth.provider import AuthorizeError, RegistrationError, TokenError
1616
from mcp.server.auth.routes import create_auth_routes
17-
from mcp.server.streamable_http_manager import DEFAULT_MAX_REQUEST_BODY_SIZE
17+
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE
1818

1919
# TODO(Marcelo): This TYPE_CHECKING shouldn't be here, but pytest doesn't seem to get the module correctly.
2020
if TYPE_CHECKING:

tests/server/test_sse_security.py

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -21,8 +21,7 @@
2121
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser
2222
from mcp.server.auth.provider import AccessToken
2323
from mcp.server.sse import SseServerTransport
24-
from mcp.server.streamable_http_manager import DEFAULT_MAX_REQUEST_BODY_SIZE
25-
from mcp.server.transport_security import TransportSecuritySettings
24+
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE, TransportSecuritySettings
2625
from mcp.types import Tool
2726
from tests.test_helpers import wait_for_server
2827

tests/server/test_streamable_http_manager.py

Lines changed: 1 addition & 88 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77

88
import anyio
99
import pytest
10-
from starlette.types import Message, Receive, Scope, Send
10+
from starlette.types import Message, Scope
1111

1212
from mcp.server import streamable_http_manager
1313
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser
@@ -16,7 +16,6 @@
1616
from mcp.server.streamable_http import MCP_SESSION_ID_HEADER, StreamableHTTPServerTransport
1717
from mcp.server.streamable_http_manager import (
1818
DEFAULT_MAX_REQUEST_BODY_SIZE,
19-
RequestBodyLimitMiddleware,
2019
StreamableHTTPSessionManager,
2120
)
2221
from mcp.types import INVALID_REQUEST
@@ -143,92 +142,6 @@ async def send(message: Message) -> None:
143142
assert response_start["status"] == 413
144143

145144

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

0 commit comments

Comments
 (0)