Apply the request body limit to the SSE and OAuth endpoints (#3336)
This commit is contained in:
+1
-1
@@ -755,7 +755,7 @@ Transport-specific parameters have been moved off the `MCPServer` constructor an
|
||||
- `sse_path`, `message_path` - SSE transport paths, on `run(transport="sse", ...)` and `sse_app()`
|
||||
- `streamable_http_path` - StreamableHTTP endpoint path, on `run(transport="streamable-http", ...)` and `streamable_http_app()`
|
||||
- `json_response`, `stateless_http` - StreamableHTTP behavior, same two places; each also removes a server-to-client channel, see [Server-initiated sampling, elicitation, and roots raise `NoBackChannelError`](#server-initiated-sampling-elicitation-and-roots-raise-nobackchannelerror)
|
||||
- `max_request_body_size` - StreamableHTTP request-body limit, same two places
|
||||
- `max_request_body_size` - HTTP request-body limit, on `run()` for both HTTP transports and on both app methods
|
||||
- `event_store`, `retry_interval` - StreamableHTTP event handling, same two places
|
||||
- `transport_security` - DNS rebinding protection, on `run()` for both HTTP transports and on both app methods
|
||||
|
||||
|
||||
+1
-1
@@ -67,7 +67,7 @@ Each transport has its own keyword arguments, all on `run()`:
|
||||
* `streamable_http_path`: where the MCP endpoint lives. Default `/mcp`.
|
||||
* `json_response=True`: answer each POST with a single JSON body instead of an SSE stream. That body has room for the response and nothing else, so a tool that calls back into the client mid-request (`ctx.elicit()`, sampling) raises `NoBackChannelError` on this leg, and notifications tied to the in-flight call (progress from `ctx.report_progress()`, per-call log messages) are dropped; the standalone `GET` stream still carries unrelated ones.
|
||||
* `stateless_http=True`: a fresh transport per request, no session tracking.
|
||||
* `max_request_body_size`: largest accepted POST body in bytes. Defaults to 4 MiB; larger requests
|
||||
* `max_request_body_size`: largest accepted request body in bytes. Defaults to 4 MiB; larger requests
|
||||
receive HTTP 413 before parsing or session creation. Raise it only when legitimate MCP messages
|
||||
exceed that size.
|
||||
* `event_store`, `retry_interval`, `transport_security`: resumability and DNS-rebinding protection. They can wait, until you deploy somewhere other than localhost; **[Deploy & scale](deploy.md)** covers `transport_security`.
|
||||
|
||||
@@ -17,6 +17,7 @@ from mcp.server.auth.handlers.token import TokenHandler
|
||||
from mcp.server.auth.middleware.client_auth import ClientAuthenticator
|
||||
from mcp.server.auth.provider import OAuthAuthorizationServerProvider
|
||||
from mcp.server.auth.settings import ClientRegistrationOptions, RevocationOptions
|
||||
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE, RequestBodyLimitMiddleware
|
||||
from mcp.shared.auth import JWT_BEARER_GRANT_TYPE, OAuthMetadata, ProtectedResourceMetadata
|
||||
from mcp.shared.inbound import MCP_PROTOCOL_VERSION_HEADER
|
||||
|
||||
@@ -51,17 +52,24 @@ REVOCATION_PATH = "/revoke"
|
||||
ID_JAG_GRANT_PROFILE = "urn:ietf:params:oauth:grant-profile:id-jag"
|
||||
|
||||
|
||||
def cors_middleware(
|
||||
handler: Callable[[Request], Response | Awaitable[Response]],
|
||||
allow_methods: list[str],
|
||||
) -> ASGIApp:
|
||||
cors_app = CORSMiddleware(
|
||||
app=request_response(handler),
|
||||
def _cors(app: ASGIApp, allow_methods: list[str]) -> ASGIApp:
|
||||
return CORSMiddleware(
|
||||
app=app,
|
||||
allow_origins="*",
|
||||
allow_methods=allow_methods,
|
||||
allow_headers=[MCP_PROTOCOL_VERSION_HEADER],
|
||||
)
|
||||
return cors_app
|
||||
|
||||
|
||||
def _body_limited(app: ASGIApp) -> ASGIApp:
|
||||
return RequestBodyLimitMiddleware(app, DEFAULT_MAX_REQUEST_BODY_SIZE)
|
||||
|
||||
|
||||
def cors_middleware(
|
||||
handler: Callable[[Request], Response | Awaitable[Response]],
|
||||
allow_methods: list[str],
|
||||
) -> ASGIApp:
|
||||
return _cors(request_response(handler), allow_methods)
|
||||
|
||||
|
||||
def create_auth_routes(
|
||||
@@ -84,11 +92,13 @@ def create_auth_routes(
|
||||
supports_identity_assertion=identity_assertion_enabled,
|
||||
)
|
||||
client_authenticator = ClientAuthenticator(provider)
|
||||
token_handler = TokenHandler(provider, client_authenticator, identity_assertion_enabled=identity_assertion_enabled)
|
||||
|
||||
# Create routes
|
||||
# Allow CORS requests for endpoints meant to be hit by the OAuth client
|
||||
# (with the client secret). This is intended to support things like MCP Inspector,
|
||||
# where the client runs in a web browser.
|
||||
# where the client runs in a web browser. CORS is the outermost wrapper so that
|
||||
# responses produced by inner layers (such as a 413) still carry CORS headers.
|
||||
routes = [
|
||||
Route(
|
||||
"/.well-known/oauth-authorization-server",
|
||||
@@ -102,17 +112,12 @@ def create_auth_routes(
|
||||
AUTHORIZATION_PATH,
|
||||
# do not allow CORS for authorization endpoint;
|
||||
# clients should just redirect to this
|
||||
endpoint=AuthorizationHandler(provider).handle,
|
||||
endpoint=_body_limited(request_response(AuthorizationHandler(provider).handle)),
|
||||
methods=["GET", "POST"],
|
||||
),
|
||||
Route(
|
||||
TOKEN_PATH,
|
||||
endpoint=cors_middleware(
|
||||
TokenHandler(
|
||||
provider, client_authenticator, identity_assertion_enabled=identity_assertion_enabled
|
||||
).handle,
|
||||
["POST", "OPTIONS"],
|
||||
),
|
||||
endpoint=_cors(_body_limited(request_response(token_handler.handle)), ["POST", "OPTIONS"]),
|
||||
methods=["POST", "OPTIONS"],
|
||||
),
|
||||
]
|
||||
@@ -125,10 +130,7 @@ def create_auth_routes(
|
||||
routes.append(
|
||||
Route(
|
||||
REGISTRATION_PATH,
|
||||
endpoint=cors_middleware(
|
||||
registration_handler.handle,
|
||||
["POST", "OPTIONS"],
|
||||
),
|
||||
endpoint=_cors(_body_limited(request_response(registration_handler.handle)), ["POST", "OPTIONS"]),
|
||||
methods=["POST", "OPTIONS"],
|
||||
)
|
||||
)
|
||||
@@ -138,10 +140,7 @@ def create_auth_routes(
|
||||
routes.append(
|
||||
Route(
|
||||
REVOCATION_PATH,
|
||||
endpoint=cors_middleware(
|
||||
revocation_handler.handle,
|
||||
["POST", "OPTIONS"],
|
||||
),
|
||||
endpoint=_cors(_body_limited(request_response(revocation_handler.handle)), ["POST", "OPTIONS"]),
|
||||
methods=["POST", "OPTIONS"],
|
||||
)
|
||||
)
|
||||
|
||||
@@ -65,12 +65,8 @@ from mcp.server.context import HandlerResult, ServerMiddleware, ServerRequestCon
|
||||
from mcp.server.models import InitializationOptions
|
||||
from mcp.server.runner import serve_dual_era_loop
|
||||
from mcp.server.streamable_http import EventStore
|
||||
from mcp.server.streamable_http_manager import (
|
||||
DEFAULT_MAX_REQUEST_BODY_SIZE,
|
||||
StreamableHTTPASGIApp,
|
||||
StreamableHTTPSessionManager,
|
||||
)
|
||||
from mcp.server.transport_security import TransportSecuritySettings
|
||||
from mcp.server.streamable_http_manager import StreamableHTTPASGIApp, StreamableHTTPSessionManager
|
||||
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE, TransportSecuritySettings
|
||||
from mcp.shared._stream_protocols import ReadStream, WriteStream
|
||||
from mcp.shared.exceptions import MCPDeprecationWarning
|
||||
from mcp.shared.message import SessionMessage
|
||||
|
||||
@@ -87,9 +87,9 @@ from mcp.server.request_state import RequestStateBoundary, RequestStateSecurity
|
||||
from mcp.server.sse import SseServerTransport
|
||||
from mcp.server.stdio import stdio_server
|
||||
from mcp.server.streamable_http import EventStore
|
||||
from mcp.server.streamable_http_manager import DEFAULT_MAX_REQUEST_BODY_SIZE, StreamableHTTPSessionManager
|
||||
from mcp.server.streamable_http_manager import StreamableHTTPSessionManager
|
||||
from mcp.server.subscriptions import InMemorySubscriptionBus, ListenHandler, SubscriptionBus
|
||||
from mcp.server.transport_security import TransportSecuritySettings
|
||||
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE, TransportSecuritySettings
|
||||
from mcp.shared.exceptions import MCPError
|
||||
from mcp.shared.uri_template import UriTemplate
|
||||
|
||||
@@ -365,6 +365,7 @@ class MCPServer(Generic[LifespanResultT]):
|
||||
port: int = ...,
|
||||
sse_path: str = ...,
|
||||
message_path: str = ...,
|
||||
max_request_body_size: int = ...,
|
||||
transport_security: TransportSecuritySettings | None = ...,
|
||||
) -> None: ...
|
||||
|
||||
@@ -1031,6 +1032,7 @@ class MCPServer(Generic[LifespanResultT]):
|
||||
port: int = 8000,
|
||||
sse_path: str = "/sse",
|
||||
message_path: str = "/messages/",
|
||||
max_request_body_size: int = DEFAULT_MAX_REQUEST_BODY_SIZE,
|
||||
transport_security: TransportSecuritySettings | None = None,
|
||||
) -> None:
|
||||
"""Run the server using SSE transport."""
|
||||
@@ -1039,6 +1041,7 @@ class MCPServer(Generic[LifespanResultT]):
|
||||
starlette_app = self.sse_app(
|
||||
sse_path=sse_path,
|
||||
message_path=message_path,
|
||||
max_request_body_size=max_request_body_size,
|
||||
transport_security=transport_security,
|
||||
host=host,
|
||||
)
|
||||
@@ -1093,6 +1096,7 @@ class MCPServer(Generic[LifespanResultT]):
|
||||
*,
|
||||
sse_path: str = "/sse",
|
||||
message_path: str = "/messages/",
|
||||
max_request_body_size: int = DEFAULT_MAX_REQUEST_BODY_SIZE,
|
||||
transport_security: TransportSecuritySettings | None = None,
|
||||
host: str = "127.0.0.1",
|
||||
) -> Starlette:
|
||||
@@ -1105,7 +1109,9 @@ class MCPServer(Generic[LifespanResultT]):
|
||||
allowed_origins=["http://127.0.0.1:*", "http://localhost:*", "http://[::1]:*"],
|
||||
)
|
||||
|
||||
sse = SseServerTransport(message_path, security_settings=transport_security)
|
||||
sse = SseServerTransport(
|
||||
message_path, security_settings=transport_security, max_request_body_size=max_request_body_size
|
||||
)
|
||||
|
||||
async def handle_sse(scope: Scope, receive: Receive, send: Send): # pragma: no cover
|
||||
# Add client ID from auth context into request context if available
|
||||
|
||||
+26
-1
@@ -52,6 +52,8 @@ from starlette.types import Receive, Scope, Send
|
||||
|
||||
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser, AuthorizationContext, authorization_context
|
||||
from mcp.server.transport_security import (
|
||||
DEFAULT_MAX_REQUEST_BODY_SIZE,
|
||||
RequestBodyLimitMiddleware,
|
||||
TransportSecurityMiddleware,
|
||||
TransportSecuritySettings,
|
||||
)
|
||||
@@ -79,7 +81,12 @@ class SseServerTransport:
|
||||
_session_owners: dict[UUID, AuthorizationContext]
|
||||
_security: TransportSecurityMiddleware
|
||||
|
||||
def __init__(self, endpoint: str, security_settings: TransportSecuritySettings | None = None) -> None:
|
||||
def __init__(
|
||||
self,
|
||||
endpoint: str,
|
||||
security_settings: TransportSecuritySettings | None = None,
|
||||
max_request_body_size: int = DEFAULT_MAX_REQUEST_BODY_SIZE,
|
||||
) -> None:
|
||||
"""Creates a new SSE server transport, which will direct the client to POST
|
||||
messages to the relative path given.
|
||||
|
||||
@@ -87,6 +94,9 @@ class SseServerTransport:
|
||||
endpoint: A relative path where messages should be posted
|
||||
(e.g., "/messages/").
|
||||
security_settings: Optional security settings for DNS rebinding protection.
|
||||
max_request_body_size: Maximum size in bytes for POSTed message bodies. Requests that
|
||||
declare or stream a larger body receive HTTP 413. Defaults to 4 MiB, matching
|
||||
`StreamableHTTPSessionManager`.
|
||||
|
||||
Note:
|
||||
We use relative paths instead of full URLs for several reasons:
|
||||
@@ -103,6 +113,9 @@ class SseServerTransport:
|
||||
|
||||
super().__init__()
|
||||
|
||||
if max_request_body_size <= 0:
|
||||
raise ValueError("max_request_body_size must be a positive number of bytes")
|
||||
|
||||
# Validate that endpoint is a relative path and not a full URL
|
||||
if "://" in endpoint or endpoint.startswith("//") or "?" in endpoint or "#" in endpoint:
|
||||
raise ValueError(
|
||||
@@ -118,6 +131,7 @@ class SseServerTransport:
|
||||
self._read_stream_writers = {}
|
||||
self._session_owners = {}
|
||||
self._security = TransportSecurityMiddleware(security_settings)
|
||||
self._post_message_app = RequestBodyLimitMiddleware(self._handle_post_message, max_request_body_size)
|
||||
logger.debug(f"SseServerTransport initialized with endpoint: {endpoint}")
|
||||
|
||||
@asynccontextmanager
|
||||
@@ -203,6 +217,17 @@ class SseServerTransport:
|
||||
self._session_owners.pop(session_id, None)
|
||||
|
||||
async def handle_post_message(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
"""ASGI application for the message endpoint.
|
||||
|
||||
Only POST is accepted (other methods get 405), and bodies larger than
|
||||
`max_request_body_size` are answered with 413 before the message is handled.
|
||||
"""
|
||||
if scope["method"] != "POST":
|
||||
response = Response(status_code=405, headers={"Allow": "POST"})
|
||||
return await response(scope, receive, send)
|
||||
await self._post_message_app(scope, receive, send)
|
||||
|
||||
async def _handle_post_message(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
logger.debug("Handling POST message")
|
||||
request = Request(scope, receive)
|
||||
|
||||
|
||||
@@ -4,25 +4,25 @@ from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import logging
|
||||
from collections import deque
|
||||
from collections.abc import AsyncIterator
|
||||
from typing import TYPE_CHECKING, Any, Final
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from uuid import uuid4
|
||||
|
||||
import anyio
|
||||
from anyio.abc import TaskStatus
|
||||
from mcp_types import DEFAULT_NEGOTIATED_VERSION, INVALID_REQUEST, ErrorData, JSONRPCError
|
||||
from mcp_types.version import HANDSHAKE_PROTOCOL_VERSIONS
|
||||
from starlette.datastructures import Headers
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
from starlette.types import ASGIApp, Message, Receive, Scope, Send
|
||||
from starlette.types import Receive, Scope, Send
|
||||
|
||||
from mcp.server._streamable_http_modern import handle_modern_request
|
||||
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser, AuthorizationContext, authorization_context
|
||||
from mcp.server.connection import Connection
|
||||
from mcp.server.runner import serve_connection, serve_loop
|
||||
from mcp.server.streamable_http import MCP_SESSION_ID_HEADER, EventStore, StreamableHTTPServerTransport
|
||||
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE as DEFAULT_MAX_REQUEST_BODY_SIZE
|
||||
from mcp.server.transport_security import RequestBodyLimitMiddleware as RequestBodyLimitMiddleware
|
||||
from mcp.server.transport_security import TransportSecuritySettings
|
||||
from mcp.shared._compat import resync_tracer
|
||||
from mcp.shared.inbound import MCP_PROTOCOL_VERSION_HEADER
|
||||
@@ -34,9 +34,6 @@ if TYPE_CHECKING:
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_MAX_REQUEST_BODY_SIZE: Final = 4 * 1024 * 1024
|
||||
"""Default maximum Streamable HTTP request body size in bytes (4 MiB)."""
|
||||
|
||||
|
||||
class StreamableHTTPSessionManager:
|
||||
"""Manages StreamableHTTP sessions with optional resumability via event store.
|
||||
@@ -70,7 +67,7 @@ class StreamableHTTPSessionManager:
|
||||
retry_interval is also configured, ensure the idle timeout comfortably exceeds the retry interval to
|
||||
avoid reaping sessions during normal SSE polling gaps. Default is None (no timeout). A value of 1800
|
||||
(30 minutes) is recommended for most deployments.
|
||||
max_request_body_size: Maximum size in bytes for Streamable HTTP POST request bodies. Requests that
|
||||
max_request_body_size: Maximum size in bytes for Streamable HTTP request bodies. Requests that
|
||||
exceed this limit receive a 413 response before parsing or session creation. Defaults to 4 MiB.
|
||||
"""
|
||||
|
||||
@@ -371,66 +368,6 @@ class StreamableHTTPSessionManager:
|
||||
await response(scope, receive, send)
|
||||
|
||||
|
||||
class RequestBodyLimitMiddleware:
|
||||
"""Reject oversized HTTP request bodies before invoking an ASGI application."""
|
||||
|
||||
def __init__(self, app: ASGIApp, max_body_size: int) -> None:
|
||||
self.app = app
|
||||
self.max_body_size = max_body_size
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["type"] != "http" or scope["method"] != "POST":
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
headers = Headers(scope=scope)
|
||||
content_length = headers.get("content-length")
|
||||
if content_length is not None:
|
||||
try:
|
||||
declared_size = int(content_length)
|
||||
except ValueError:
|
||||
pass
|
||||
else:
|
||||
if declared_size > self.max_body_size:
|
||||
response = Response("Request body too large", status_code=413)
|
||||
return await response(scope, receive, send)
|
||||
|
||||
received_body = bytearray()
|
||||
received_request = False
|
||||
body_complete = False
|
||||
trailing_message: Message | None = None
|
||||
while True:
|
||||
message = await receive()
|
||||
if message["type"] != "http.request":
|
||||
trailing_message = message
|
||||
break
|
||||
|
||||
received_request = True
|
||||
body = message.get("body", b"")
|
||||
if len(received_body) + len(body) > self.max_body_size:
|
||||
response = Response("Request body too large", status_code=413)
|
||||
return await response(scope, receive, send)
|
||||
received_body.extend(body)
|
||||
if not message.get("more_body", False):
|
||||
body_complete = True
|
||||
break
|
||||
|
||||
cached_messages: deque[Message] = deque()
|
||||
if received_request:
|
||||
cached_messages.append(
|
||||
{"type": "http.request", "body": bytes(received_body), "more_body": not body_complete}
|
||||
)
|
||||
if trailing_message is not None:
|
||||
cached_messages.append(trailing_message)
|
||||
|
||||
async def replay() -> Message:
|
||||
if cached_messages:
|
||||
return cached_messages.popleft()
|
||||
return await receive()
|
||||
|
||||
await self.app(scope, replay, send)
|
||||
|
||||
|
||||
class StreamableHTTPASGIApp:
|
||||
"""ASGI application for Streamable HTTP server transport."""
|
||||
|
||||
|
||||
@@ -1,13 +1,20 @@
|
||||
"""DNS rebinding protection for MCP server transports."""
|
||||
"""Request checks shared by the HTTP server transports: Host/Origin header validation and body size limits."""
|
||||
|
||||
import logging
|
||||
from collections import deque
|
||||
from typing import Final
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from starlette.datastructures import Headers
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
from starlette.types import ASGIApp, Message, Receive, Scope, Send
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_MAX_REQUEST_BODY_SIZE: Final = 4 * 1024 * 1024
|
||||
"""Default maximum HTTP request body size in bytes (4 MiB)."""
|
||||
|
||||
|
||||
# TODO(Marcelo): We should flatten these settings. To be fair, I don't think we should even have this middleware.
|
||||
class TransportSecuritySettings(BaseModel):
|
||||
@@ -114,3 +121,63 @@ class TransportSecurityMiddleware:
|
||||
return Response("Invalid Origin header", status_code=403)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
class RequestBodyLimitMiddleware:
|
||||
"""Reject oversized HTTP request bodies before invoking an ASGI application."""
|
||||
|
||||
def __init__(self, app: ASGIApp, max_body_size: int) -> None:
|
||||
self.app = app
|
||||
self.max_body_size = max_body_size
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["type"] != "http":
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
headers = Headers(scope=scope)
|
||||
content_length = headers.get("content-length")
|
||||
if content_length is not None:
|
||||
try:
|
||||
declared_size = int(content_length)
|
||||
except ValueError:
|
||||
pass
|
||||
else:
|
||||
if declared_size > self.max_body_size:
|
||||
response = Response("Request body too large", status_code=413)
|
||||
return await response(scope, receive, send)
|
||||
|
||||
received_body = bytearray()
|
||||
received_request = False
|
||||
body_complete = False
|
||||
trailing_message: Message | None = None
|
||||
while True:
|
||||
message = await receive()
|
||||
if message["type"] != "http.request":
|
||||
trailing_message = message
|
||||
break
|
||||
|
||||
received_request = True
|
||||
body = message.get("body", b"")
|
||||
if len(received_body) + len(body) > self.max_body_size:
|
||||
response = Response("Request body too large", status_code=413)
|
||||
return await response(scope, receive, send)
|
||||
received_body.extend(body)
|
||||
if not message.get("more_body", False):
|
||||
body_complete = True
|
||||
break
|
||||
|
||||
cached_messages: deque[Message] = deque()
|
||||
if received_request:
|
||||
cached_messages.append(
|
||||
{"type": "http.request", "body": bytes(received_body), "more_body": not body_complete}
|
||||
)
|
||||
if trailing_message is not None:
|
||||
cached_messages.append(trailing_message)
|
||||
|
||||
async def replay() -> Message:
|
||||
if cached_messages:
|
||||
return cached_messages.popleft()
|
||||
return await receive()
|
||||
|
||||
await self.app(scope, replay, send)
|
||||
|
||||
@@ -16,6 +16,7 @@ from starlette.applications import Starlette
|
||||
from mcp.server.auth.provider import AuthorizeError, RegistrationError, TokenError
|
||||
from mcp.server.auth.routes import create_auth_routes
|
||||
from mcp.server.auth.settings import ClientRegistrationOptions, RevocationOptions
|
||||
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE
|
||||
from tests.server.mcpserver.auth.test_auth_integration import MockOAuthProvider
|
||||
|
||||
|
||||
@@ -288,3 +289,61 @@ async def test_token_error_handling_refresh_token(
|
||||
data = refresh_response.json()
|
||||
assert data["error"] == "invalid_scope"
|
||||
assert data["error_description"] == "The requested scope is invalid"
|
||||
|
||||
|
||||
_FORM = "application/x-www-form-urlencoded"
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
@pytest.mark.parametrize(
|
||||
("method", "path", "content_type"),
|
||||
[
|
||||
("POST", "/token", _FORM),
|
||||
("POST", "/revoke", _FORM),
|
||||
("POST", "/register", "application/json"),
|
||||
("POST", "/authorize", _FORM),
|
||||
# The other methods these routes accept reach the same body-reading handlers.
|
||||
("OPTIONS", "/token", _FORM),
|
||||
("OPTIONS", "/revoke", _FORM),
|
||||
("OPTIONS", "/register", "application/json"),
|
||||
("HEAD", "/authorize", _FORM),
|
||||
],
|
||||
)
|
||||
async def test_oversized_request_body_returns_413(
|
||||
client: httpx2.AsyncClient, method: str, path: str, content_type: str
|
||||
):
|
||||
"""Each endpoint that reads a request body rejects one over 4 MiB before parsing it, whatever the method."""
|
||||
response = await client.request(
|
||||
method, path, content=b"x" * (DEFAULT_MAX_REQUEST_BODY_SIZE + 1), headers={"Content-Type": content_type}
|
||||
)
|
||||
assert response.status_code == 413
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_request_body_within_the_limit_is_still_parsed(client: httpx2.AsyncClient):
|
||||
"""A small body is passed through to the handler intact: the form is parsed and its fields validated."""
|
||||
response = await client.post("/token", data={"grant_type": "authorization_code"})
|
||||
assert response.status_code == 401
|
||||
assert response.json() == {"error": "invalid_client", "error_description": "Missing client_id"}
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_cors_preflight_is_still_answered(client: httpx2.AsyncClient):
|
||||
"""A CORS preflight to a body-limited endpoint is answered by the CORS layer as before."""
|
||||
response = await client.options(
|
||||
"/token", headers={"Origin": "https://client.example.com", "Access-Control-Request-Method": "POST"}
|
||||
)
|
||||
assert response.status_code == 200
|
||||
assert response.headers["access-control-allow-origin"] == "*"
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_oversized_cross_origin_request_gets_413_with_cors_headers(client: httpx2.AsyncClient):
|
||||
"""The 413 is produced inside the CORS layer, so a browser client can still read it."""
|
||||
response = await client.post(
|
||||
"/token",
|
||||
content=b"x" * (DEFAULT_MAX_REQUEST_BODY_SIZE + 1),
|
||||
headers={"Content-Type": _FORM, "Origin": "https://client.example.com"},
|
||||
)
|
||||
assert response.status_code == 413
|
||||
assert response.headers["access-control-allow-origin"] == "*"
|
||||
|
||||
@@ -5,6 +5,7 @@ from typing import Any
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import anyio
|
||||
import httpx2
|
||||
import pytest
|
||||
from inline_snapshot import snapshot
|
||||
from mcp_types import (
|
||||
@@ -1812,6 +1813,19 @@ def test_streamable_http_no_redirect() -> None:
|
||||
assert streamable_routes[0].path == "/mcp", "Streamable route path should be /mcp"
|
||||
|
||||
|
||||
async def test_sse_app_applies_the_configured_request_body_limit() -> None:
|
||||
"""`sse_app(max_request_body_size=...)` rejects larger POSTs to the message endpoint with HTTP 413."""
|
||||
app = MCPServer("test").sse_app(max_request_body_size=8, host="0.0.0.0")
|
||||
transport = httpx2.ASGITransport(app=app)
|
||||
async with httpx2.AsyncClient(transport=transport, base_url="http://localhost") as http:
|
||||
response = await http.post(
|
||||
"/messages/?session_id=12345678123456781234567812345678",
|
||||
content=b"123456789",
|
||||
headers={"Content-Type": "application/json"},
|
||||
)
|
||||
assert response.status_code == 413
|
||||
|
||||
|
||||
async def test_report_progress_delegates_to_session_report_progress():
|
||||
"""Context.report_progress delegates to ServerSession.report_progress unconditionally.
|
||||
|
||||
|
||||
@@ -18,7 +18,7 @@ from mcp.server import Server
|
||||
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser
|
||||
from mcp.server.auth.provider import AccessToken
|
||||
from mcp.server.sse import SseServerTransport
|
||||
from mcp.server.transport_security import TransportSecuritySettings
|
||||
from mcp.server.transport_security import DEFAULT_MAX_REQUEST_BODY_SIZE, TransportSecuritySettings
|
||||
from mcp.shared._stream_protocols import WriteStream
|
||||
from mcp.shared.message import SessionMessage
|
||||
from tests.interaction.transports import StreamingASGITransport
|
||||
@@ -204,9 +204,18 @@ def _authenticated_user(client_id: str, subject: str | None = None, issuer: str
|
||||
|
||||
|
||||
def _sse_scope(
|
||||
method: str, path: str, user: AuthenticatedUser | None, *, query_string: bytes = b"", body: bytes = b""
|
||||
method: str,
|
||||
path: str,
|
||||
user: AuthenticatedUser | None,
|
||||
*,
|
||||
query_string: bytes = b"",
|
||||
body: bytes | list[bytes] = b"",
|
||||
) -> tuple[Scope, Receive, Send, list[Message]]:
|
||||
"""Build an ASGI scope/receive/send triple for a request to the SSE transport."""
|
||||
"""Build an ASGI scope/receive/send triple for a request to the SSE transport.
|
||||
|
||||
`body` may be a list of chunks to deliver the request body over several `http.request` messages;
|
||||
no Content-Length header is set either way.
|
||||
"""
|
||||
scope: Scope = {
|
||||
"type": "http",
|
||||
"method": method,
|
||||
@@ -218,9 +227,11 @@ def _sse_scope(
|
||||
if user is not None:
|
||||
scope["user"] = user
|
||||
sent: list[Message] = []
|
||||
chunks = list(body) if isinstance(body, list) else [body]
|
||||
|
||||
async def receive() -> Message:
|
||||
return {"type": "http.request", "body": body, "more_body": False}
|
||||
chunk = chunks.pop(0)
|
||||
return {"type": "http.request", "body": chunk, "more_body": bool(chunks)}
|
||||
|
||||
async def send(message: Message) -> None:
|
||||
sent.append(message)
|
||||
@@ -233,6 +244,10 @@ def _response_status(sent: list[Message]) -> int:
|
||||
return response_start["status"]
|
||||
|
||||
|
||||
def _response_body(sent: list[Message]) -> bytes:
|
||||
return b"".join(msg.get("body", b"") for msg in sent if msg["type"] == "http.response.body")
|
||||
|
||||
|
||||
async def _post_message(transport: SseServerTransport, session_id: str, user: AuthenticatedUser | None) -> int:
|
||||
"""POST a message to an SSE session as `user` and return the response status."""
|
||||
body = b'{"jsonrpc": "2.0", "id": 1, "method": "ping", "params": null}'
|
||||
@@ -368,6 +383,80 @@ async def test_sse_post_with_a_disallowed_host_is_rejected_before_session_lookup
|
||||
assert _response_status(sent) == 421
|
||||
|
||||
|
||||
# A well-formed session ID that no live session owns.
|
||||
_UNKNOWN_SESSION = b"session_id=12345678123456781234567812345678"
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_sse_post_body_over_the_limit_returns_413():
|
||||
"""A POST body larger than max_request_body_size is answered with 413 before any session handling."""
|
||||
transport = SseServerTransport("/messages/", max_request_body_size=8)
|
||||
scope, receive, send, sent = _sse_scope(
|
||||
"POST", "/messages/", None, query_string=_UNKNOWN_SESSION, body=b"123456789"
|
||||
)
|
||||
|
||||
await transport.handle_post_message(scope, receive, send)
|
||||
assert _response_status(sent) == 413
|
||||
assert _response_body(sent) == b"Request body too large"
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_sse_post_body_limit_defaults_to_four_mib():
|
||||
"""Without an explicit limit, a body one byte over 4 MiB (and no Content-Length) is answered with 413."""
|
||||
transport = SseServerTransport("/messages/")
|
||||
body = b"x" * (DEFAULT_MAX_REQUEST_BODY_SIZE + 1)
|
||||
scope, receive, send, sent = _sse_scope("POST", "/messages/", None, query_string=_UNKNOWN_SESSION, body=body)
|
||||
|
||||
await transport.handle_post_message(scope, receive, send)
|
||||
assert _response_status(sent) == 413
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_sse_post_streamed_body_over_the_limit_returns_413():
|
||||
"""The limit counts bytes across body chunks, not just a declared Content-Length."""
|
||||
transport = SseServerTransport("/messages/", max_request_body_size=8)
|
||||
scope, receive, send, sent = _sse_scope(
|
||||
"POST", "/messages/", None, query_string=_UNKNOWN_SESSION, body=[b"1234", b"56789"]
|
||||
)
|
||||
|
||||
await transport.handle_post_message(scope, receive, send)
|
||||
assert _response_status(sent) == 413
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_sse_post_within_the_limit_reaches_session_lookup():
|
||||
"""A body within the limit is passed on intact: an unknown session still gets its 404."""
|
||||
transport = SseServerTransport("/messages/", max_request_body_size=64)
|
||||
scope, receive, send, sent = _sse_scope(
|
||||
"POST", "/messages/", None, query_string=_UNKNOWN_SESSION, body=[b'{"jsonrpc": ', b'"2.0"}']
|
||||
)
|
||||
|
||||
await transport.handle_post_message(scope, receive, send)
|
||||
assert _response_status(sent) == 404
|
||||
assert _response_body(sent) == b"Could not find session"
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
@pytest.mark.parametrize("method", ["GET", "PUT"])
|
||||
async def test_sse_message_endpoint_answers_405_to_non_post(method: str):
|
||||
"""The message endpoint only accepts POST; other methods get 405 with an Allow header."""
|
||||
transport = SseServerTransport("/messages/")
|
||||
scope, receive, send, sent = _sse_scope(method, "/messages/", None, query_string=_UNKNOWN_SESSION, body=b"{}")
|
||||
|
||||
await transport.handle_post_message(scope, receive, send)
|
||||
assert _response_status(sent) == 405
|
||||
response_start = next(msg for msg in sent if msg["type"] == "http.response.start")
|
||||
assert (b"allow", b"POST") in response_start["headers"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("max_request_body_size", [0, -1])
|
||||
def test_sse_transport_rejects_a_non_positive_body_limit(max_request_body_size: int):
|
||||
"""The body limit must be a positive number of bytes, matching StreamableHTTPSessionManager."""
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
SseServerTransport("/messages/", max_request_body_size=max_request_body_size)
|
||||
assert str(exc_info.value) == "max_request_body_size must be a positive number of bytes"
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_sse_round_trip_delivers_posted_messages_and_streams_responses():
|
||||
"""A POSTed JSON-RPC message reaches the server's read stream, and a message
|
||||
|
||||
@@ -10,7 +10,7 @@ import anyio
|
||||
import httpx2
|
||||
import pytest
|
||||
from mcp_types import INVALID_REQUEST, ListToolsResult, PaginatedRequestParams
|
||||
from starlette.types import Message, Receive, Scope, Send
|
||||
from starlette.types import Message, Scope
|
||||
|
||||
from mcp import Client
|
||||
from mcp.client.streamable_http import streamable_http_client
|
||||
@@ -18,11 +18,7 @@ from mcp.server import Server, ServerRequestContext, streamable_http_manager
|
||||
from mcp.server.auth.middleware.bearer_auth import AuthenticatedUser
|
||||
from mcp.server.auth.provider import AccessToken
|
||||
from mcp.server.streamable_http import MCP_SESSION_ID_HEADER, StreamableHTTPServerTransport
|
||||
from mcp.server.streamable_http_manager import (
|
||||
DEFAULT_MAX_REQUEST_BODY_SIZE,
|
||||
RequestBodyLimitMiddleware,
|
||||
StreamableHTTPSessionManager,
|
||||
)
|
||||
from mcp.server.streamable_http_manager import DEFAULT_MAX_REQUEST_BODY_SIZE, StreamableHTTPSessionManager
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
@@ -146,79 +142,6 @@ async def test_oversized_streamed_body_is_rejected_before_session_creation(
|
||||
assert response_start["status"] == 413
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_client_disconnect_while_streaming_request_body_is_replayed() -> None:
|
||||
"""SDK-defined: raw ASGI is required to prove a disconnect before body completion reaches the transport."""
|
||||
disconnect: Message = {"type": "http.disconnect"}
|
||||
request_messages: Iterator[Message] = iter(
|
||||
[{"type": "http.request", "body": b"1234", "more_body": True}, disconnect]
|
||||
)
|
||||
received_messages: list[Message] = []
|
||||
|
||||
async def receive() -> Message:
|
||||
return next(request_messages)
|
||||
|
||||
async def app(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
received_messages.append(await receive())
|
||||
received_messages.append(await receive())
|
||||
|
||||
scope: Scope = {"type": "http", "method": "POST", "path": "/mcp", "headers": []}
|
||||
middleware = RequestBodyLimitMiddleware(app, max_body_size=8)
|
||||
|
||||
await middleware(scope, receive, AsyncMock())
|
||||
|
||||
assert received_messages == [
|
||||
{"type": "http.request", "body": b"1234", "more_body": True},
|
||||
disconnect,
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_client_disconnect_before_request_body_is_replayed() -> None:
|
||||
"""SDK-defined: raw ASGI proves a disconnect before the first body message reaches the transport."""
|
||||
disconnect: Message = {"type": "http.disconnect"}
|
||||
received_messages: list[Message] = []
|
||||
|
||||
async def receive() -> Message:
|
||||
return disconnect
|
||||
|
||||
async def app(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
received_messages.append(await receive())
|
||||
|
||||
scope: Scope = {"type": "http", "method": "POST", "path": "/mcp", "headers": []}
|
||||
middleware = RequestBodyLimitMiddleware(app, max_body_size=8)
|
||||
|
||||
await middleware(scope, receive, AsyncMock())
|
||||
|
||||
assert received_messages == [disconnect]
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_request_body_chunks_are_replayed_as_one_message() -> None:
|
||||
"""SDK-defined: raw ASGI proves chunk overhead is discarded before the body reaches the transport."""
|
||||
request_messages: Iterator[Message] = iter(
|
||||
[
|
||||
{"type": "http.request", "body": b"12", "more_body": True},
|
||||
{"type": "http.request", "body": b"34", "more_body": True},
|
||||
{"type": "http.request", "body": b"56", "more_body": False},
|
||||
]
|
||||
)
|
||||
received_messages: list[Message] = []
|
||||
|
||||
async def receive() -> Message:
|
||||
return next(request_messages)
|
||||
|
||||
async def app(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
received_messages.append(await receive())
|
||||
|
||||
scope: Scope = {"type": "http", "method": "POST", "path": "/mcp", "headers": []}
|
||||
middleware = RequestBodyLimitMiddleware(app, max_body_size=8)
|
||||
|
||||
await middleware(scope, receive, AsyncMock())
|
||||
|
||||
assert received_messages == [{"type": "http.request", "body": b"123456", "more_body": False}]
|
||||
|
||||
|
||||
def test_request_body_limit_defaults_to_four_mib() -> None:
|
||||
"""SDK-defined: Streamable HTTP request bodies are limited to 4 MiB by default."""
|
||||
manager = StreamableHTTPSessionManager(app=Server("test-default-size-limit"))
|
||||
|
||||
@@ -1,9 +1,17 @@
|
||||
"""Tests for the transport-security request validation middleware."""
|
||||
"""Tests for the request checks shared by the HTTP server transports."""
|
||||
|
||||
from collections.abc import Iterator
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
from starlette.requests import Request
|
||||
from starlette.types import Message, Receive, Scope, Send
|
||||
|
||||
from mcp.server.transport_security import TransportSecurityMiddleware, TransportSecuritySettings
|
||||
from mcp.server.transport_security import (
|
||||
RequestBodyLimitMiddleware,
|
||||
TransportSecurityMiddleware,
|
||||
TransportSecuritySettings,
|
||||
)
|
||||
|
||||
|
||||
def _request(host: str | None, origin: str | None, content_type: str | None = "application/json") -> Request:
|
||||
@@ -86,3 +94,111 @@ async def test_validate_request_ignores_content_type_on_get() -> None:
|
||||
middleware = TransportSecurityMiddleware(SETTINGS)
|
||||
response = await middleware.validate_request(_request("good.example", None, content_type="text/plain"))
|
||||
assert response is None
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_client_disconnect_while_streaming_request_body_is_replayed() -> None:
|
||||
"""SDK-defined: raw ASGI is required to prove a disconnect before body completion reaches the transport."""
|
||||
disconnect: Message = {"type": "http.disconnect"}
|
||||
request_messages: Iterator[Message] = iter(
|
||||
[{"type": "http.request", "body": b"1234", "more_body": True}, disconnect]
|
||||
)
|
||||
received_messages: list[Message] = []
|
||||
|
||||
async def receive() -> Message:
|
||||
return next(request_messages)
|
||||
|
||||
async def app(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
received_messages.append(await receive())
|
||||
received_messages.append(await receive())
|
||||
|
||||
scope: Scope = {"type": "http", "method": "POST", "path": "/mcp", "headers": []}
|
||||
middleware = RequestBodyLimitMiddleware(app, max_body_size=8)
|
||||
|
||||
await middleware(scope, receive, AsyncMock())
|
||||
|
||||
assert received_messages == [
|
||||
{"type": "http.request", "body": b"1234", "more_body": True},
|
||||
disconnect,
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_client_disconnect_before_request_body_is_replayed() -> None:
|
||||
"""SDK-defined: raw ASGI proves a disconnect before the first body message reaches the transport."""
|
||||
disconnect: Message = {"type": "http.disconnect"}
|
||||
received_messages: list[Message] = []
|
||||
|
||||
async def receive() -> Message:
|
||||
return disconnect
|
||||
|
||||
async def app(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
received_messages.append(await receive())
|
||||
|
||||
scope: Scope = {"type": "http", "method": "POST", "path": "/mcp", "headers": []}
|
||||
middleware = RequestBodyLimitMiddleware(app, max_body_size=8)
|
||||
|
||||
await middleware(scope, receive, AsyncMock())
|
||||
|
||||
assert received_messages == [disconnect]
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_request_body_chunks_are_replayed_as_one_message() -> None:
|
||||
"""SDK-defined: raw ASGI proves chunk overhead is discarded before the body reaches the transport."""
|
||||
request_messages: Iterator[Message] = iter(
|
||||
[
|
||||
{"type": "http.request", "body": b"12", "more_body": True},
|
||||
{"type": "http.request", "body": b"34", "more_body": True},
|
||||
{"type": "http.request", "body": b"56", "more_body": False},
|
||||
]
|
||||
)
|
||||
received_messages: list[Message] = []
|
||||
|
||||
async def receive() -> Message:
|
||||
return next(request_messages)
|
||||
|
||||
async def app(scope: Scope, receive: Receive, send: Send) -> None:
|
||||
received_messages.append(await receive())
|
||||
|
||||
scope: Scope = {"type": "http", "method": "POST", "path": "/mcp", "headers": []}
|
||||
middleware = RequestBodyLimitMiddleware(app, max_body_size=8)
|
||||
|
||||
await middleware(scope, receive, AsyncMock())
|
||||
|
||||
assert received_messages == [{"type": "http.request", "body": b"123456", "more_body": False}]
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
@pytest.mark.parametrize("method", ["GET", "PUT", "OPTIONS", "HEAD", "DELETE"])
|
||||
async def test_request_body_limit_applies_to_every_method(method: str) -> None:
|
||||
"""SDK-defined: the limit is a property of the request body, not of the method that carries it."""
|
||||
app = AsyncMock()
|
||||
sent_messages: list[Message] = []
|
||||
receive = AsyncMock(return_value={"type": "http.request", "body": b"123456789", "more_body": False})
|
||||
|
||||
async def send(message: Message) -> None:
|
||||
sent_messages.append(message)
|
||||
|
||||
scope: Scope = {"type": "http", "method": method, "path": "/mcp", "headers": []}
|
||||
middleware = RequestBodyLimitMiddleware(app, max_body_size=8)
|
||||
|
||||
await middleware(scope, receive, send)
|
||||
|
||||
assert [message["status"] for message in sent_messages if message["type"] == "http.response.start"] == [413]
|
||||
app.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.anyio
|
||||
async def test_request_body_limit_leaves_non_http_scopes_alone() -> None:
|
||||
"""SDK-defined: only HTTP requests carry a body to limit; other ASGI scopes go straight to the app."""
|
||||
app = AsyncMock()
|
||||
receive = AsyncMock()
|
||||
send = AsyncMock()
|
||||
scope: Scope = {"type": "lifespan"}
|
||||
middleware = RequestBodyLimitMiddleware(app, max_body_size=8)
|
||||
|
||||
await middleware(scope, receive, send)
|
||||
|
||||
app.assert_awaited_once_with(scope, receive, send)
|
||||
receive.assert_not_awaited()
|
||||
|
||||
Reference in New Issue
Block a user