From 6df3d734265eb49bf758a5e5eb420937337184e9 Mon Sep 17 00:00:00 2001 From: Max <224885523+maxisbey@users.noreply.github.com> Date: Tue, 23 Jun 2026 14:57:45 +0100 Subject: [PATCH 1/3] [v1.x] Buffer per-request StreamableHTTP streams; store priming event before dispatch (#2948) --- src/mcp/server/streamable_http.py | 185 +++++++++++--------- tests/server/test_streamable_http_router.py | 116 ++++++++++++ tests/shared/test_streamable_http.py | 69 ++------ 3 files changed, 229 insertions(+), 141 deletions(-) create mode 100644 tests/server/test_streamable_http_router.py diff --git a/src/mcp/server/streamable_http.py b/src/mcp/server/streamable_http.py index c241e831a9..8e8d902ccc 100644 --- a/src/mcp/server/streamable_http.py +++ b/src/mcp/server/streamable_http.py @@ -14,8 +14,9 @@ from collections.abc import AsyncGenerator, Awaitable, Callable from contextlib import asynccontextmanager from dataclasses import dataclass +from functools import partial from http import HTTPStatus -from typing import Any +from typing import Any, Final import anyio from anyio.streams.memory import MemoryObjectReceiveStream, MemoryObjectSendStream @@ -60,6 +61,11 @@ # Special key for the standalone GET stream GET_STREAM_KEY = "_GET_stream" +# Buffer for the per-request `_request_streams` so the serial `message_router` +# can deposit a response and move on instead of head-of-line blocking the +# whole session on a lazily-started `sse_writer`. See #1764. +REQUEST_STREAM_BUFFER_SIZE: Final = 16 + # Session ID validation pattern (visible ASCII characters ranging from 0x21 to 0x7E) # Pattern ensures entire string contains only valid characters by using ^ and $ anchors SESSION_ID_PATTERN = re.compile(r"^[\x21-\x7E]+$") @@ -67,6 +73,8 @@ # Type aliases StreamId = str EventId = str +# An SSE event-dict as accepted by sse-starlette (`event`, `data`, `id`, `retry`). +SSEEvent = dict[str, Any] @dataclass @@ -178,7 +186,7 @@ def __init__( MemoryObjectReceiveStream[EventMessage], ], ] = {} - self._sse_stream_writers: dict[RequestId, MemoryObjectSendStream[dict[str, str]]] = {} + self._sse_stream_writers: dict[RequestId, MemoryObjectSendStream[SSEEvent]] = {} self._terminated = False # Idle timeout cancel scope; managed by the session manager. self.idle_scope: anyio.CancelScope | None = None @@ -267,31 +275,48 @@ async def close_standalone_stream_callback() -> None: return SessionMessage(message, metadata=metadata) - async def _maybe_send_priming_event( - self, - request_id: RequestId, - sse_stream_writer: MemoryObjectSendStream[dict[str, Any]], - protocol_version: str, - ) -> None: - """Send priming event for SSE resumability if event_store is configured. + async def _mint_priming_event(self, stream_id: StreamId, protocol_version: str) -> SSEEvent | None: + """Store the priming cursor for `stream_id` and return its SSE wire form. - Only sends priming events to clients with protocol version >= 2025-11-25, - which includes the fix for handling empty SSE data. Older clients would - crash trying to parse empty data as JSON. + Called before the request is dispatched so the priming row precedes + anything `message_router` can store for this stream. Returns `None` + when no event store is configured or the client predates 2025-11-25 + (older clients cannot parse the empty-data event). """ if not self._event_store: - return - # Priming events have empty data which older clients cannot handle. + return None if protocol_version < "2025-11-25": - return - priming_event_id = await self._event_store.store_event( - str(request_id), # Convert RequestId to StreamId (str) - None, # Priming event has no payload - ) - priming_event: dict[str, str | int] = {"id": priming_event_id, "data": ""} + return None + priming_event_id = await self._event_store.store_event(stream_id, None) + priming_event: SSEEvent = {"id": priming_event_id, "data": ""} if self._retry_interval is not None: priming_event["retry"] = self._retry_interval - await sse_stream_writer.send(priming_event) + return priming_event + + async def _run_sse_writer( # pragma: no cover + self, + request_id: RequestId, + sse_stream_writer: MemoryObjectSendStream[SSEEvent], + request_stream_reader: MemoryObjectReceiveStream[EventMessage], + priming_event: SSEEvent | None, + ) -> None: + """Forward `_request_streams[request_id]` onto the SSE wire for one POST.""" + try: + async with sse_stream_writer, request_stream_reader: + if priming_event is not None: + await sse_stream_writer.send(priming_event) + async for event_message in request_stream_reader: + await sse_stream_writer.send(self._create_event_data(event_message)) + if isinstance(event_message.message.root, JSONRPCResponse | JSONRPCError): + break + except anyio.ClosedResourceError: + logger.debug("SSE stream closed by close_sse_stream()") + except Exception: + logger.exception("Error in SSE writer") + finally: + logger.debug("Closing SSE writer") + self._sse_stream_writers.pop(request_id, None) + await self._clean_up_memory_streams(request_id) def _create_error_response( self, @@ -348,7 +373,7 @@ def _get_session_id(self, request: Request) -> str | None: # pragma: no cover """Extract the session ID from request headers.""" return request.headers.get(MCP_SESSION_ID_HEADER) - def _create_event_data(self, event_message: EventMessage) -> dict[str, str]: # pragma: no cover + def _create_event_data(self, event_message: EventMessage) -> SSEEvent: # pragma: no cover """Create event data dictionary from an EventMessage.""" event_data = { "event": "message", @@ -530,13 +555,13 @@ async def _handle_post_request(self, scope: Scope, request: Request, receive: Re else request.headers.get(MCP_PROTOCOL_VERSION_HEADER, DEFAULT_NEGOTIATED_VERSION) ) - # Extract the request ID outside the try block for proper scope - request_id = str(message.root.id) # pragma: no cover - # Register this stream for the request ID - self._request_streams[request_id] = anyio.create_memory_object_stream[EventMessage](0) # pragma: no cover - request_stream_reader = self._request_streams[request_id][1] # pragma: no cover + request_id = str(message.root.id) if self.is_json_response_enabled: # pragma: no cover + self._request_streams[request_id] = anyio.create_memory_object_stream[EventMessage]( + REQUEST_STREAM_BUFFER_SIZE + ) + request_stream_reader = self._request_streams[request_id][1] # Process the message metadata = ServerMessageMetadata(request_context=request) session_message = SessionMessage(message, metadata=metadata) @@ -580,44 +605,19 @@ async def _handle_post_request(self, scope: Scope, request: Request, receive: Re finally: await self._clean_up_memory_streams(request_id) else: # pragma: no cover - # Create SSE stream - sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[dict[str, str]](0) + # Mint the priming event before any per-request state exists: + # `EventStore.store_event` is user code and may raise, in which + # case the outer handler returns a 500 with nothing to clean up. + # Still strictly precedes dispatch, so storage order == wire order. + priming_event = await self._mint_priming_event(request_id, protocol_version) - # Store writer reference so close_sse_stream() can close it + sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[SSEEvent](0) self._sse_stream_writers[request_id] = sse_stream_writer + self._request_streams[request_id] = anyio.create_memory_object_stream[EventMessage]( + REQUEST_STREAM_BUFFER_SIZE + ) + request_stream_reader = self._request_streams[request_id][1] - async def sse_writer(): - # Get the request ID from the incoming request message - try: - async with sse_stream_writer, request_stream_reader: - # Send priming event for SSE resumability - await self._maybe_send_priming_event(request_id, sse_stream_writer, protocol_version) - - # Process messages from the request-specific stream - async for event_message in request_stream_reader: - # Build the event data - event_data = self._create_event_data(event_message) - await sse_stream_writer.send(event_data) - - # If response, remove from pending streams and close - if isinstance( - event_message.message.root, - JSONRPCResponse | JSONRPCError, - ): - break - except anyio.ClosedResourceError: - # Expected when close_sse_stream() is called - logger.debug("SSE stream closed by close_sse_stream()") - except Exception: - logger.exception("Error in SSE writer") - finally: - logger.debug("Closing SSE writer") - self._sse_stream_writers.pop(request_id, None) - await self._clean_up_memory_streams(request_id) - - # Create and start EventSourceResponse - # SSE stream mode (original behavior) - # Set up headers headers = { "Cache-Control": "no-cache, no-transform", "Connection": "keep-alive", @@ -626,7 +626,9 @@ async def sse_writer(): } response = EventSourceResponse( content=sse_stream_reader, - data_sender_callable=sse_writer, + data_sender_callable=partial( + self._run_sse_writer, request_id, sse_stream_writer, request_stream_reader, priming_event + ), headers=headers, ) @@ -644,16 +646,15 @@ async def sse_writer(): await sse_stream_reader.aclose() await self._clean_up_memory_streams(request_id) - except Exception as err: # pragma: no cover + except Exception as err: logger.exception("Error handling POST request") response = self._create_error_response( - f"Error handling POST request: {err}", + "Error handling POST request", HTTPStatus.INTERNAL_SERVER_ERROR, INTERNAL_ERROR, ) await response(scope, receive, send) - if writer: - await writer.send(Exception(err)) + await writer.send(Exception(err)) return async def _handle_get_request(self, request: Request, send: Send) -> None: # pragma: no cover @@ -706,13 +707,15 @@ async def _handle_get_request(self, request: Request, send: Send) -> None: # pr return # Create SSE stream - sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[dict[str, str]](0) + sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[SSEEvent](0) async def standalone_sse_writer(): try: # Create a standalone message stream for server-initiated messages - self._request_streams[GET_STREAM_KEY] = anyio.create_memory_object_stream[EventMessage](0) + self._request_streams[GET_STREAM_KEY] = anyio.create_memory_object_stream[EventMessage]( + REQUEST_STREAM_BUFFER_SIZE + ) standalone_stream_reader = self._request_streams[GET_STREAM_KEY][1] async with sse_stream_writer, standalone_stream_reader: @@ -903,7 +906,7 @@ async def _replay_events(self, last_event_id: str, request: Request, send: Send) replay_protocol_version = request.headers.get(MCP_PROTOCOL_VERSION_HEADER, DEFAULT_NEGOTIATED_VERSION) # Create SSE stream for replay - sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[dict[str, str]](0) + sse_stream_writer, sse_stream_reader = anyio.create_memory_object_stream[SSEEvent](0) async def replay_sender(): try: @@ -918,22 +921,32 @@ async def send_event(event_message: EventMessage) -> None: # If stream ID not in mapping, create it if stream_id and stream_id not in self._request_streams: - # Register SSE writer so close_sse_stream() can close it - self._sse_stream_writers[stream_id] = sse_stream_writer - - # Send priming event for this new connection - await self._maybe_send_priming_event(stream_id, sse_stream_writer, replay_protocol_version) - - # Create new request streams for this connection - self._request_streams[stream_id] = anyio.create_memory_object_stream[EventMessage](0) - msg_reader = self._request_streams[stream_id][1] - - # Forward messages to SSE - async with msg_reader: - async for event_message in msg_reader: - event_data = self._create_event_data(event_message) - - await sse_stream_writer.send(event_data) + try: + # Register SSE writer so close_sse_stream() can close it + self._sse_stream_writers[stream_id] = sse_stream_writer + + # Prime the resumed connection so the client sees the stream + # is re-registered. The replay→live-tail ordering window here + # is pre-existing and tracked separately. + priming_event = await self._mint_priming_event(stream_id, replay_protocol_version) + if priming_event is not None: + await sse_stream_writer.send(priming_event) + + # Create new request streams for this connection + self._request_streams[stream_id] = anyio.create_memory_object_stream[EventMessage]( + REQUEST_STREAM_BUFFER_SIZE + ) + msg_reader = self._request_streams[stream_id][1] + + # Forward messages to SSE + async with msg_reader: + async for event_message in msg_reader: + event_data = self._create_event_data(event_message) + + await sse_stream_writer.send(event_data) + finally: + self._sse_stream_writers.pop(stream_id, None) + await self._clean_up_memory_streams(stream_id) except anyio.ClosedResourceError: # Expected when close_sse_stream() is called logger.debug("Replay SSE stream closed by close_sse_stream()") diff --git a/tests/server/test_streamable_http_router.py b/tests/server/test_streamable_http_router.py new file mode 100644 index 0000000000..e78c17e91f --- /dev/null +++ b/tests/server/test_streamable_http_router.py @@ -0,0 +1,116 @@ +"""Regression coverage for the StreamableHTTP per-session response router.""" + +import anyio +import pytest +from starlette.types import Message, Scope + +from mcp.server.streamable_http import ( + REQUEST_STREAM_BUFFER_SIZE, + EventCallback, + EventId, + EventMessage, + EventStore, + StreamableHTTPServerTransport, + StreamId, +) +from mcp.shared.message import SessionMessage +from mcp.types import JSONRPCMessage, JSONRPCResponse + + +class _PrimingFailingStore(EventStore): + async def store_event(self, stream_id: StreamId, message: JSONRPCMessage | None) -> EventId: + raise RuntimeError("backend unavailable") + + async def replay_events_after(self, last_event_id: EventId, send_callback: EventCallback) -> StreamId | None: + raise NotImplementedError + + +@pytest.mark.anyio +async def test_router_unconsumed_request_stream_does_not_block_siblings() -> None: + """A response whose `sse_writer` is not yet receiving must not park the router (#1764). + + Drives the routing layer directly (the production race does not reproduce + on loopback), so this pins the router semantics, not the call sites. + """ + transport = StreamableHTTPServerTransport(mcp_session_id="sid", is_json_response_enabled=False) + streams = transport._request_streams + async with transport.connect() as (_read_stream, write_stream): + # Model two concurrent POSTs at the point _handle_post_request has + # registered the per-request stream but A's sse_writer has not yet + # reached its first receive(). + streams["A"] = anyio.create_memory_object_stream[EventMessage](REQUEST_STREAM_BUFFER_SIZE) + streams["B"] = anyio.create_memory_object_stream[EventMessage](REQUEST_STREAM_BUFFER_SIZE) + a_send, a_recv = streams["A"] + b_reader = streams["B"][1] + b_received = anyio.Event() + + async def consume_b() -> None: + async with b_reader: + await b_reader.receive() + b_received.set() + + async def server_writes() -> None: + await write_stream.send(SessionMessage(JSONRPCMessage(JSONRPCResponse(jsonrpc="2.0", id="A", result={})))) + await write_stream.send(SessionMessage(JSONRPCMessage(JSONRPCResponse(jsonrpc="2.0", id="B", result={})))) + + async with anyio.create_task_group() as tg: + tg.start_soon(consume_b) + tg.start_soon(server_writes) + with anyio.fail_after(5): + await b_received.wait() + # A's response was buffered for its (late) consumer, not dropped. + assert a_send.statistics().current_buffer_used == 1 + await a_recv.aclose() + await a_send.aclose() + + +@pytest.mark.anyio +async def test_priming_store_failure_leaves_no_per_request_state() -> None: + """`EventStore.store_event` raising on the priming row must not leak per-request entries.""" + transport = StreamableHTTPServerTransport( + mcp_session_id=None, + is_json_response_enabled=False, + event_store=_PrimingFailingStore(), + ) + + body = b'{"jsonrpc":"2.0","id":"req-1","method":"tools/list","params":{}}' + scope: Scope = { + "type": "http", + "method": "POST", + "path": "/", + "query_string": b"", + "headers": [ + (b"accept", b"application/json, text/event-stream"), + (b"content-type", b"application/json"), + (b"mcp-protocol-version", b"2025-11-25"), + ], + } + body_sent = False + + async def receive() -> Message: + nonlocal body_sent + if not body_sent: + body_sent = True + return {"type": "http.request", "body": body, "more_body": False} + raise NotImplementedError + + sent: list[Message] = [] + + async def asgi_send(message: Message) -> None: + sent.append(message) + + async with transport.connect() as (read_stream, _write_stream): + async with anyio.create_task_group() as tg: + tg.start_soon(transport.handle_request, scope, receive, asgi_send) + with anyio.fail_after(5): + forwarded = await read_stream.receive() + assert isinstance(forwarded, Exception) + # handle_request has returned; connect()'s finally (which clears + # _request_streams unconditionally) has not yet run. + assert transport._request_streams == {} + assert transport._sse_stream_writers == {} + + assert sent[0]["type"] == "http.response.start" + assert sent[0]["status"] == 500 + body = b"".join(m.get("body", b"") for m in sent if m["type"] == "http.response.body") + assert b"backend unavailable" not in body diff --git a/tests/shared/test_streamable_http.py b/tests/shared/test_streamable_http.py index 731dd20dd3..631c96e86c 100644 --- a/tests/shared/test_streamable_http.py +++ b/tests/shared/test_streamable_http.py @@ -1786,81 +1786,40 @@ async def test_handle_sse_event_skips_empty_data(): @pytest.mark.anyio -async def test_priming_event_not_sent_for_old_protocol_version(): - """Test that _maybe_send_priming_event skips for old protocol versions (backwards compat).""" - # Create a transport with an event store +async def test_priming_event_not_minted_for_old_protocol_version(): + """`_mint_priming_event` returns None for pre-2025-11-25 clients (backwards compat).""" transport = StreamableHTTPServerTransport( "/mcp", event_store=SimpleEventStore(), ) - # Create a mock stream writer - write_stream, read_stream = anyio.create_memory_object_stream[dict[str, Any]](1) - - try: - # Call _maybe_send_priming_event with OLD protocol version - should NOT send - await transport._maybe_send_priming_event("test-request-id", write_stream, "2025-06-18") - - # Nothing should have been written to the stream - assert write_stream.statistics().current_buffer_used == 0 - - # Now test with NEW protocol version - should send - await transport._maybe_send_priming_event("test-request-id-2", write_stream, "2025-11-25") - - # Should have written a priming event - assert write_stream.statistics().current_buffer_used == 1 - finally: - await write_stream.aclose() - await read_stream.aclose() + assert await transport._mint_priming_event("test-request-id", "2025-06-18") is None + event = await transport._mint_priming_event("test-request-id-2", "2025-11-25") + assert event is not None + assert event["data"] == "" + assert "retry" not in event @pytest.mark.anyio -async def test_priming_event_not_sent_without_event_store(): - """Test that _maybe_send_priming_event returns early when no event_store is configured.""" - # Create a transport WITHOUT an event store +async def test_priming_event_not_minted_without_event_store(): + """`_mint_priming_event` returns None when no event store is configured.""" transport = StreamableHTTPServerTransport("/mcp") - # Create a mock stream writer - write_stream, read_stream = anyio.create_memory_object_stream[dict[str, Any]](1) - - try: - # Call _maybe_send_priming_event - should return early without sending - await transport._maybe_send_priming_event("test-request-id", write_stream, "2025-11-25") - - # Nothing should have been written to the stream - assert write_stream.statistics().current_buffer_used == 0 - finally: - await write_stream.aclose() - await read_stream.aclose() + assert await transport._mint_priming_event("test-request-id", "2025-11-25") is None @pytest.mark.anyio async def test_priming_event_includes_retry_interval(): - """Test that _maybe_send_priming_event includes retry field when retry_interval is set.""" - # Create a transport with an event store AND retry_interval + """`_mint_priming_event` carries the configured `retry` field.""" transport = StreamableHTTPServerTransport( "/mcp", event_store=SimpleEventStore(), retry_interval=5000, ) - # Create a mock stream writer - write_stream, read_stream = anyio.create_memory_object_stream[dict[str, Any]](1) - - try: - # Call _maybe_send_priming_event with new protocol version - await transport._maybe_send_priming_event("test-request-id", write_stream, "2025-11-25") - - # Should have written a priming event with retry field - assert write_stream.statistics().current_buffer_used == 1 - - # Read the event and verify it has retry field - event = await read_stream.receive() - assert "retry" in event - assert event["retry"] == 5000 - finally: - await write_stream.aclose() - await read_stream.aclose() + event = await transport._mint_priming_event("test-request-id", "2025-11-25") + assert event is not None + assert event["retry"] == 5000 @pytest.mark.anyio From 47204674fb26185c2cf45f065831f27b8e5d5c65 Mon Sep 17 00:00:00 2001 From: Max <224885523+maxisbey@users.noreply.github.com> Date: Thu, 25 Jun 2026 23:48:24 +0200 Subject: [PATCH 2/3] [v1.x] Set Development Status classifier to Production/Stable (#2976) --- pyproject.toml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pyproject.toml b/pyproject.toml index 769ddcf709..8611f5c4b6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -12,7 +12,7 @@ maintainers = [ keywords = ["git", "mcp", "llm", "automation"] license = { text = "MIT" } classifiers = [ - "Development Status :: 4 - Beta", + "Development Status :: 5 - Production/Stable", "Intended Audience :: Developers", "License :: OSI Approved :: MIT License", "Programming Language :: Python :: 3", From 777b8d06710c140e3606b0d4598e2aa48546c266 Mon Sep 17 00:00:00 2001 From: Max <224885523+maxisbey@users.noreply.github.com> Date: Fri, 26 Jun 2026 13:31:33 +0200 Subject: [PATCH 3/3] [v1.x] Support TransportSecuritySettings in the WebSocket server transport (#2992) --- src/mcp/server/transport_security.py | 4 +- src/mcp/server/websocket.py | 21 ++- tests/server/test_websocket_security.py | 172 ++++++++++++++++++++++++ 3 files changed, 194 insertions(+), 3 deletions(-) create mode 100644 tests/server/test_websocket_security.py diff --git a/src/mcp/server/transport_security.py b/src/mcp/server/transport_security.py index ee1e4505a7..5022a1a2fe 100644 --- a/src/mcp/server/transport_security.py +++ b/src/mcp/server/transport_security.py @@ -3,7 +3,7 @@ import logging from pydantic import BaseModel, Field -from starlette.requests import Request +from starlette.requests import HTTPConnection from starlette.responses import Response logger = logging.getLogger(__name__) @@ -99,7 +99,7 @@ def _validate_content_type(self, content_type: str | None) -> bool: # pragma: n return True - async def validate_request(self, request: Request, is_post: bool = False) -> Response | None: + async def validate_request(self, request: HTTPConnection, is_post: bool = False) -> Response | None: """Validate request headers for DNS rebinding protection. Returns None if validation passes, or an error Response if validation fails. diff --git a/src/mcp/server/websocket.py b/src/mcp/server/websocket.py index 2b21604a76..d3526f2ad6 100644 --- a/src/mcp/server/websocket.py +++ b/src/mcp/server/websocket.py @@ -9,6 +9,7 @@ from typing_extensions import deprecated import mcp.types as types +from mcp.server.transport_security import TransportSecurityMiddleware, TransportSecuritySettings from mcp.shared.message import SessionMessage logger = logging.getLogger(__name__) @@ -19,16 +20,34 @@ " the MCP specification; use the streamable HTTP transport instead." ) @asynccontextmanager -async def websocket_server(scope: Scope, receive: Receive, send: Send): +async def websocket_server( + scope: Scope, + receive: Receive, + send: Send, + security_settings: TransportSecuritySettings | None = None, +): """ WebSocket server transport for MCP. This is an ASGI application, suitable to be used with a framework like Starlette and a server like Hypercorn. + Set `security_settings` to enable Host/Origin header validation before the + handshake is accepted (same settings type as the SSE and Streamable HTTP + transports). When validation fails this raises `ValueError` after rejecting + the handshake. + Deprecated: this transport will be removed in mcp 2.0. WebSocket was never part of the MCP specification; use the streamable HTTP transport instead. """ websocket = WebSocket(scope, receive, send) + + security = TransportSecurityMiddleware(security_settings) + error_response = await security.validate_request(websocket, is_post=False) + if error_response is not None: + # Reject the handshake; the ASGI server maps a pre-accept close to HTTP 403. + await websocket.close() + raise ValueError("Request validation failed") + await websocket.accept(subprotocol="mcp") read_stream: MemoryObjectReceiveStream[SessionMessage | Exception] diff --git a/tests/server/test_websocket_security.py b/tests/server/test_websocket_security.py new file mode 100644 index 0000000000..35f778080d --- /dev/null +++ b/tests/server/test_websocket_security.py @@ -0,0 +1,172 @@ +"""Tests for WebSocket server request validation.""" + +# pyright: reportDeprecated=false + +import logging +import multiprocessing +import socket +import warnings + +import pytest +import uvicorn +from starlette.applications import Starlette +from starlette.routing import WebSocketRoute +from starlette.types import Message, Scope +from starlette.websockets import WebSocket +from websockets.asyncio.client import connect +from websockets.exceptions import InvalidStatus +from websockets.typing import Subprotocol + +from mcp.server import Server +from mcp.server.transport_security import TransportSecuritySettings +from mcp.server.websocket import websocket_server +from tests.test_helpers import wait_for_server + +logger = logging.getLogger(__name__) +SERVER_NAME = "test_ws_security_server" + +# This suite intentionally exercises the deprecated WebSocket transport. +pytestmark = pytest.mark.filterwarnings( + "ignore:The WebSocket (client|server) transport is deprecated:DeprecationWarning" +) + + +@pytest.fixture +def server_port() -> int: + with socket.socket() as s: + s.bind(("127.0.0.1", 0)) + return s.getsockname()[1] + + +def run_server_with_settings(port: int, security_settings: TransportSecuritySettings | None = None): # pragma: no cover + """Run a WebSocket MCP server with the given security settings.""" + warnings.filterwarnings("ignore", category=DeprecationWarning) + server = Server(SERVER_NAME) + + async def handle_ws(websocket: WebSocket) -> None: + try: + async with websocket_server( + websocket.scope, websocket.receive, websocket.send, security_settings=security_settings + ) as streams: + await server.run(streams[0], streams[1], server.create_initialization_options()) + except ValueError as exc: + logger.debug(f"WebSocket connection failed validation: {exc}") + + app = Starlette(routes=[WebSocketRoute("/ws", endpoint=handle_ws)]) + uvicorn.run(app, host="127.0.0.1", port=port, log_level="error") + + +def start_server_process(port: int, security_settings: TransportSecuritySettings | None = None): + """Start the server in a subprocess and wait until it accepts connections.""" + process = multiprocessing.Process(target=run_server_with_settings, args=(port, security_settings)) + process.start() + wait_for_server(port) + return process + + +@pytest.mark.anyio +async def test_ws_security_default_settings(server_port: int) -> None: + """With no security settings the WebSocket transport accepts any Origin (matches SSE/StreamableHTTP default).""" + process = start_server_process(server_port) + try: + async with connect( + f"ws://127.0.0.1:{server_port}/ws", + subprotocols=[Subprotocol("mcp")], + additional_headers={"Origin": "http://evil.com"}, + ) as ws: + assert ws.subprotocol == "mcp" + finally: + process.terminate() + process.join() + + +@pytest.mark.anyio +async def test_ws_security_invalid_origin_header(server_port: int) -> None: + """An Origin not in allowed_origins is rejected before the handshake completes.""" + settings = TransportSecuritySettings( + enable_dns_rebinding_protection=True, allowed_hosts=["127.0.0.1:*"], allowed_origins=["http://localhost:*"] + ) + process = start_server_process(server_port, settings) + try: + with pytest.raises(InvalidStatus) as exc_info: + async with connect( + f"ws://127.0.0.1:{server_port}/ws", + subprotocols=[Subprotocol("mcp")], + additional_headers={"Origin": "http://evil.com"}, + ): + pytest.fail("handshake should have been rejected") # pragma: no cover + assert exc_info.value.response.status_code == 403 + finally: + process.terminate() + process.join() + + +@pytest.mark.anyio +async def test_ws_security_invalid_host_header(server_port: int) -> None: + """A Host not in allowed_hosts is rejected before the handshake completes.""" + settings = TransportSecuritySettings(enable_dns_rebinding_protection=True, allowed_hosts=["example.com"]) + process = start_server_process(server_port, settings) + try: + with pytest.raises(InvalidStatus) as exc_info: + async with connect(f"ws://127.0.0.1:{server_port}/ws", subprotocols=[Subprotocol("mcp")]): + pytest.fail("handshake should have been rejected") # pragma: no cover + assert exc_info.value.response.status_code == 403 + finally: + process.terminate() + process.join() + + +@pytest.mark.anyio +async def test_ws_security_allowed_origin(server_port: int) -> None: + """An Origin matching allowed_origins is accepted.""" + settings = TransportSecuritySettings( + enable_dns_rebinding_protection=True, allowed_hosts=["127.0.0.1:*"], allowed_origins=["http://localhost:*"] + ) + process = start_server_process(server_port, settings) + try: + async with connect( + f"ws://127.0.0.1:{server_port}/ws", + subprotocols=[Subprotocol("mcp")], + additional_headers={"Origin": "http://localhost:8080"}, + ) as ws: + assert ws.subprotocol == "mcp" + finally: + process.terminate() + process.join() + + +@pytest.mark.anyio +async def test_ws_security_disabled(server_port: int) -> None: + """Explicitly disabling protection accepts any Origin.""" + settings = TransportSecuritySettings(enable_dns_rebinding_protection=False) + process = start_server_process(server_port, settings) + try: + async with connect( + f"ws://127.0.0.1:{server_port}/ws", + subprotocols=[Subprotocol("mcp")], + additional_headers={"Origin": "http://evil.com"}, + ) as ws: + assert ws.subprotocol == "mcp" + finally: + process.terminate() + process.join() + + +@pytest.mark.anyio +async def test_ws_security_rejects_before_accept() -> None: + """A failing validation closes the connection before the handshake is accepted.""" + settings = TransportSecuritySettings(enable_dns_rebinding_protection=True, allowed_hosts=["example.com"]) + sent: list[Message] = [] + + async def receive() -> Message: + raise NotImplementedError + + async def send(message: Message) -> None: + sent.append(message) + + scope: Scope = {"type": "websocket", "headers": [(b"host", b"evil.com")]} + with pytest.raises(ValueError, match="Request validation failed"): + async with websocket_server(scope, receive, send, security_settings=settings): + pytest.fail("should not yield streams") # pragma: no cover + + assert [m["type"] for m in sent] == ["websocket.close"]