import json import multiprocessing import socket import time from collections.abc import AsyncGenerator, Generator from typing import Any from unittest.mock import AsyncMock, MagicMock, Mock, patch from urllib.parse import urlparse import anyio import httpx import pytest import uvicorn from httpx_sse import ServerSentEvent from inline_snapshot import snapshot from starlette.applications import Starlette from starlette.requests import Request from starlette.responses import Response from starlette.routing import Mount, Route import mcp.client.sse import mcp.types as types from mcp.client.session import ClientSession from mcp.client.sse import _extract_session_id_from_endpoint, sse_client from mcp.server import Server from mcp.server.sse import SseServerTransport from mcp.server.transport_security import TransportSecuritySettings from mcp.shared.exceptions import MCPError from mcp.types import ( EmptyResult, Implementation, InitializeResult, JSONRPCResponse, ReadResourceResult, ServerCapabilities, TextContent, TextResourceContents, Tool, ) from tests.test_helpers import wait_for_server SERVER_NAME = "test_server_for_SSE" @pytest.fixture def server_port() -> int: with socket.socket() as s: s.bind(("127.0.0.1", 0)) return s.getsockname()[1] @pytest.fixture def server_url(server_port: int) -> str: return f"http://127.0.0.1:{server_port}" # Test server implementation class ServerTest(Server): # pragma: no cover def __init__(self): super().__init__(SERVER_NAME) @self.read_resource() async def handle_read_resource(uri: str) -> str | bytes: parsed = urlparse(uri) if parsed.scheme == "foobar": return f"Read {parsed.netloc}" if parsed.scheme == "slow": # Simulate a slow resource await anyio.sleep(2.0) return f"Slow response from {parsed.netloc}" raise MCPError(code=404, message="OOPS! no resource with that URI was found") @self.list_tools() async def handle_list_tools() -> list[Tool]: return [ Tool( name="test_tool", description="A test tool", input_schema={"type": "object", "properties": {}}, ) ] @self.call_tool() async def handle_call_tool(name: str, args: dict[str, Any]) -> list[TextContent]: return [TextContent(type="text", text=f"Called {name}")] # Test fixtures def make_server_app() -> Starlette: # pragma: no cover """Create test Starlette app with SSE transport""" # Configure security with allowed hosts/origins for testing security_settings = TransportSecuritySettings( allowed_hosts=["127.0.0.1:*", "localhost:*"], allowed_origins=["http://127.0.0.1:*", "http://localhost:*"] ) sse = SseServerTransport("/messages/", security_settings=security_settings) server = ServerTest() async def handle_sse(request: Request) -> Response: async with sse.connect_sse(request.scope, request.receive, request._send) as streams: await server.run(streams[0], streams[1], server.create_initialization_options()) return Response() app = Starlette( routes=[ Route("/sse", endpoint=handle_sse), Mount("/messages/", app=sse.handle_post_message), ] ) return app def run_server(server_port: int) -> None: # pragma: no cover app = make_server_app() server = uvicorn.Server(config=uvicorn.Config(app=app, host="127.0.0.1", port=server_port, log_level="error")) print(f"starting server on {server_port}") server.run() # Give server time to start while not server.started: print("waiting for server to start") time.sleep(0.5) @pytest.fixture() def server(server_port: int) -> Generator[None, None, None]: proc = multiprocessing.Process(target=run_server, kwargs={"server_port": server_port}, daemon=True) print("starting process") proc.start() # Wait for server to be running print("waiting for server to start") wait_for_server(server_port) yield print("killing server") # Signal the server to stop proc.kill() proc.join(timeout=2) if proc.is_alive(): # pragma: no cover print("server process failed to terminate") @pytest.fixture() async def http_client(server: None, server_url: str) -> AsyncGenerator[httpx.AsyncClient, None]: """Create test client""" async with httpx.AsyncClient(base_url=server_url) as client: yield client # Tests @pytest.mark.anyio async def test_raw_sse_connection(http_client: httpx.AsyncClient) -> None: """Test the SSE connection establishment simply with an HTTP client.""" async with anyio.create_task_group(): async def connection_test() -> None: async with http_client.stream("GET", "/sse") as response: assert response.status_code == 200 assert response.headers["content-type"] == "text/event-stream; charset=utf-8" line_number = 0 async for line in response.aiter_lines(): # pragma: no branch if line_number == 0: assert line == "event: endpoint" elif line_number == 1: assert line.startswith("data: /messages/?session_id=") else: return line_number += 1 # Add timeout to prevent test from hanging if it fails with anyio.fail_after(3): await connection_test() @pytest.mark.anyio async def test_sse_client_basic_connection(server: None, server_url: str) -> None: async with sse_client(server_url + "/sse") as streams: async with ClientSession(*streams) as session: # Test initialization result = await session.initialize() assert isinstance(result, InitializeResult) assert result.server_info.name == SERVER_NAME # Test ping ping_result = await session.send_ping() assert isinstance(ping_result, EmptyResult) @pytest.mark.anyio async def test_sse_client_on_session_created(server: None, server_url: str) -> None: captured_session_id: str | None = None def on_session_created(session_id: str) -> None: nonlocal captured_session_id captured_session_id = session_id async with sse_client(server_url + "/sse", on_session_created=on_session_created) as streams: async with ClientSession(*streams) as session: result = await session.initialize() assert isinstance(result, InitializeResult) assert captured_session_id is not None assert len(captured_session_id) > 0 @pytest.mark.parametrize( "endpoint_url,expected", [ ("/messages?sessionId=abc123", "abc123"), ("/messages?session_id=def456", "def456"), ("/messages?sessionId=abc&session_id=def", "abc"), ("/messages?other=value", None), ("/messages", None), ("", None), ], ) def test_extract_session_id_from_endpoint(endpoint_url: str, expected: str | None) -> None: assert _extract_session_id_from_endpoint(endpoint_url) == expected @pytest.mark.anyio async def test_sse_client_on_session_created_not_called_when_no_session_id( server: None, server_url: str, monkeypatch: pytest.MonkeyPatch ) -> None: callback_mock = Mock() def mock_extract(url: str) -> None: return None monkeypatch.setattr(mcp.client.sse, "_extract_session_id_from_endpoint", mock_extract) async with sse_client(server_url + "/sse", on_session_created=callback_mock) as streams: async with ClientSession(*streams) as session: result = await session.initialize() assert isinstance(result, InitializeResult) callback_mock.assert_not_called() @pytest.fixture async def initialized_sse_client_session(server: None, server_url: str) -> AsyncGenerator[ClientSession, None]: async with sse_client(server_url + "/sse", sse_read_timeout=0.5) as streams: async with ClientSession(*streams) as session: await session.initialize() yield session @pytest.mark.anyio async def test_sse_client_happy_request_and_response( initialized_sse_client_session: ClientSession, ) -> None: session = initialized_sse_client_session response = await session.read_resource(uri="foobar://should-work") assert len(response.contents) == 1 assert isinstance(response.contents[0], TextResourceContents) assert response.contents[0].text == "Read should-work" @pytest.mark.anyio async def test_sse_client_exception_handling( initialized_sse_client_session: ClientSession, ) -> None: session = initialized_sse_client_session with pytest.raises(MCPError, match="OOPS! no resource with that URI was found"): await session.read_resource(uri="xxx://will-not-work") @pytest.mark.anyio @pytest.mark.skip("this test highlights a possible bug in SSE read timeout exception handling") async def test_sse_client_timeout( # pragma: no cover initialized_sse_client_session: ClientSession, ) -> None: session = initialized_sse_client_session # sanity check that normal, fast responses are working response = await session.read_resource(uri="foobar://1") assert isinstance(response, ReadResourceResult) with anyio.move_on_after(3): with pytest.raises(MCPError, match="Read timed out"): response = await session.read_resource(uri="slow://2") # we should receive an error here return pytest.fail("the client should have timed out and returned an error already") def run_mounted_server(server_port: int) -> None: # pragma: no cover app = make_server_app() main_app = Starlette(routes=[Mount("/mounted_app", app=app)]) server = uvicorn.Server(config=uvicorn.Config(app=main_app, host="127.0.0.1", port=server_port, log_level="error")) print(f"starting server on {server_port}") server.run() # Give server time to start while not server.started: print("waiting for server to start") time.sleep(0.5) @pytest.fixture() def mounted_server(server_port: int) -> Generator[None, None, None]: proc = multiprocessing.Process(target=run_mounted_server, kwargs={"server_port": server_port}, daemon=True) print("starting process") proc.start() # Wait for server to be running print("waiting for server to start") wait_for_server(server_port) yield print("killing server") # Signal the server to stop proc.kill() proc.join(timeout=2) if proc.is_alive(): # pragma: no cover print("server process failed to terminate") @pytest.mark.anyio async def test_sse_client_basic_connection_mounted_app(mounted_server: None, server_url: str) -> None: async with sse_client(server_url + "/mounted_app/sse") as streams: async with ClientSession(*streams) as session: # Test initialization result = await session.initialize() assert isinstance(result, InitializeResult) assert result.server_info.name == SERVER_NAME # Test ping ping_result = await session.send_ping() assert isinstance(ping_result, EmptyResult) # Test server with request context that returns headers in the response class RequestContextServer(Server[object, Request]): # pragma: no cover def __init__(self): super().__init__("request_context_server") @self.call_tool() async def handle_call_tool(name: str, args: dict[str, Any]) -> list[TextContent]: headers_info = {} context = self.request_context if context.request: headers_info = dict(context.request.headers) if name == "echo_headers": return [TextContent(type="text", text=json.dumps(headers_info))] elif name == "echo_context": context_data = { "request_id": args.get("request_id"), "headers": headers_info, } return [TextContent(type="text", text=json.dumps(context_data))] return [TextContent(type="text", text=f"Called {name}")] @self.list_tools() async def handle_list_tools() -> list[Tool]: return [ Tool( name="echo_headers", description="Echoes request headers", input_schema={"type": "object", "properties": {}}, ), Tool( name="echo_context", description="Echoes request context", input_schema={ "type": "object", "properties": {"request_id": {"type": "string"}}, "required": ["request_id"], }, ), ] def run_context_server(server_port: int) -> None: # pragma: no cover """Run a server that captures request context""" # Configure security with allowed hosts/origins for testing security_settings = TransportSecuritySettings( allowed_hosts=["127.0.0.1:*", "localhost:*"], allowed_origins=["http://127.0.0.1:*", "http://localhost:*"] ) sse = SseServerTransport("/messages/", security_settings=security_settings) context_server = RequestContextServer() async def handle_sse(request: Request) -> Response: async with sse.connect_sse(request.scope, request.receive, request._send) as streams: await context_server.run(streams[0], streams[1], context_server.create_initialization_options()) return Response() app = Starlette( routes=[ Route("/sse", endpoint=handle_sse), Mount("/messages/", app=sse.handle_post_message), ] ) server = uvicorn.Server(config=uvicorn.Config(app=app, host="127.0.0.1", port=server_port, log_level="error")) print(f"starting context server on {server_port}") server.run() @pytest.fixture() def context_server(server_port: int) -> Generator[None, None, None]: """Fixture that provides a server with request context capture""" proc = multiprocessing.Process(target=run_context_server, kwargs={"server_port": server_port}, daemon=True) print("starting context server process") proc.start() # Wait for server to be running print("waiting for context server to start") wait_for_server(server_port) yield print("killing context server") proc.kill() proc.join(timeout=2) if proc.is_alive(): # pragma: no cover print("context server process failed to terminate") @pytest.mark.anyio async def test_request_context_propagation(context_server: None, server_url: str) -> None: """Test that request context is properly propagated through SSE transport.""" # Test with custom headers custom_headers = { "Authorization": "Bearer test-token", "X-Custom-Header": "test-value", "X-Trace-Id": "trace-123", } async with sse_client(server_url + "/sse", headers=custom_headers) as ( read_stream, write_stream, ): async with ClientSession(read_stream, write_stream) as session: # Initialize the session result = await session.initialize() assert isinstance(result, InitializeResult) # Call the tool that echoes headers back tool_result = await session.call_tool("echo_headers", {}) # Parse the JSON response assert len(tool_result.content) == 1 headers_data = json.loads(tool_result.content[0].text if tool_result.content[0].type == "text" else "{}") # Verify headers were propagated assert headers_data.get("authorization") == "Bearer test-token" assert headers_data.get("x-custom-header") == "test-value" assert headers_data.get("x-trace-id") == "trace-123" @pytest.mark.anyio async def test_request_context_isolation(context_server: None, server_url: str) -> None: """Test that request contexts are isolated between different SSE clients.""" contexts: list[dict[str, Any]] = [] # Create multiple clients with different headers for i in range(3): headers = {"X-Request-Id": f"request-{i}", "X-Custom-Value": f"value-{i}"} async with sse_client(server_url + "/sse", headers=headers) as ( read_stream, write_stream, ): async with ClientSession(read_stream, write_stream) as session: await session.initialize() # Call the tool that echoes context tool_result = await session.call_tool("echo_context", {"request_id": f"request-{i}"}) assert len(tool_result.content) == 1 context_data = json.loads( tool_result.content[0].text if tool_result.content[0].type == "text" else "{}" ) contexts.append(context_data) # Verify each request had its own context assert len(contexts) == 3 for i, ctx in enumerate(contexts): assert ctx["request_id"] == f"request-{i}" assert ctx["headers"].get("x-request-id") == f"request-{i}" assert ctx["headers"].get("x-custom-value") == f"value-{i}" def test_sse_message_id_coercion(): """Previously, the `RequestId` would coerce a string that looked like an integer into an integer. See for more details. As per the JSON-RPC 2.0 specification, the id in the response object needs to be the same type as the id in the request object. In other words, we can't perform the coercion. See for more details. """ json_message = '{"jsonrpc": "2.0", "id": "123", "method": "ping", "params": null}' msg = types.JSONRPCRequest.model_validate_json(json_message) assert msg == snapshot(types.JSONRPCRequest(method="ping", jsonrpc="2.0", id="123")) json_message = '{"jsonrpc": "2.0", "id": 123, "method": "ping", "params": null}' msg = types.JSONRPCRequest.model_validate_json(json_message) assert msg == snapshot(types.JSONRPCRequest(method="ping", jsonrpc="2.0", id=123)) @pytest.mark.parametrize( "endpoint, expected_result", [ # Valid endpoints - should normalize and work ("/messages/", "/messages/"), ("messages/", "/messages/"), ("/", "/"), # Invalid endpoints - should raise ValueError ("http://example.com/messages/", ValueError), ("//example.com/messages/", ValueError), ("ftp://example.com/messages/", ValueError), ("/messages/?param=value", ValueError), ("/messages/#fragment", ValueError), ], ) def test_sse_server_transport_endpoint_validation(endpoint: str, expected_result: str | type[Exception]): """Test that SseServerTransport properly validates and normalizes endpoints.""" if isinstance(expected_result, type): # Test invalid endpoints that should raise an exception with pytest.raises(expected_result, match="is not a relative path.*expecting a relative path"): SseServerTransport(endpoint) else: # Test valid endpoints that should normalize correctly sse = SseServerTransport(endpoint) assert sse._endpoint == expected_result assert sse._endpoint.startswith("/") # ResourceWarning filter: When mocking aconnect_sse, the sse_client's internal task # group doesn't receive proper cancellation signals, so the sse_reader task's finally # block (which closes read_stream_writer) doesn't execute. This is a test artifact - # the actual code path (`if not sse.data: continue`) IS exercised and works correctly. # Production code with real SSE connections cleans up properly. @pytest.mark.filterwarnings("ignore::ResourceWarning") @pytest.mark.anyio async def test_sse_client_handles_empty_keepalive_pings() -> None: """Test that SSE client properly handles empty data lines (keep-alive pings). Per the MCP spec (Streamable HTTP transport): "The server SHOULD immediately send an SSE event consisting of an event ID and an empty data field in order to prime the client to reconnect." This test mocks the SSE event stream to include empty "message" events and verifies the client skips them without crashing. """ # Build a proper JSON-RPC response using types (not hardcoded strings) init_result = InitializeResult( protocol_version="2024-11-05", capabilities=ServerCapabilities(), server_info=Implementation(name="test", version="1.0"), ) response = JSONRPCResponse( jsonrpc="2.0", id=1, result=init_result.model_dump(by_alias=True, exclude_none=True), ) response_json = response.model_dump_json(by_alias=True, exclude_none=True) # Create mock SSE events using httpx_sse's ServerSentEvent async def mock_aiter_sse() -> AsyncGenerator[ServerSentEvent, None]: # First: endpoint event yield ServerSentEvent(event="endpoint", data="/messages/?session_id=abc123") # Empty data keep-alive ping - this is what we're testing yield ServerSentEvent(event="message", data="") # Real JSON-RPC response yield ServerSentEvent(event="message", data=response_json) mock_event_source = MagicMock() mock_event_source.aiter_sse.return_value = mock_aiter_sse() mock_event_source.response = MagicMock() mock_event_source.response.raise_for_status = MagicMock() mock_aconnect_sse = MagicMock() mock_aconnect_sse.__aenter__ = AsyncMock(return_value=mock_event_source) mock_aconnect_sse.__aexit__ = AsyncMock(return_value=None) mock_client = MagicMock() mock_client.__aenter__ = AsyncMock(return_value=mock_client) mock_client.__aexit__ = AsyncMock(return_value=None) mock_client.post = AsyncMock(return_value=MagicMock(status_code=200, raise_for_status=MagicMock())) with ( patch("mcp.client.sse.create_mcp_http_client", return_value=mock_client), patch("mcp.client.sse.aconnect_sse", return_value=mock_aconnect_sse), ): async with sse_client("http://test/sse") as (read_stream, _): # Read the message - should skip the empty one and get the real response msg = await read_stream.receive() # If we get here without error, the empty message was skipped successfully assert not isinstance(msg, Exception) assert isinstance(msg.message, types.JSONRPCResponse) assert msg.message.id == 1