From 9c936168cd99894b5d52fe5c27a43dd39b42ef64 Mon Sep 17 00:00:00 2001 From: Jeremiah Lowin <153965+jlowin@users.noreply.github.com> Date: Wed, 7 May 2025 11:04:47 -0400 Subject: [PATCH] Add integration tests for SSE --- src/fastmcp/client/client.py | 5 +- src/fastmcp/utilities/tests.py | 74 +++++++++++++++++++- tests/client/test_sse.py | 92 ++++++++++++++++++++++++ tests/server/test_http_dependencies.py | 96 ++++++++++++++++++++++++++ 4 files changed, 264 insertions(+), 3 deletions(-) create mode 100644 tests/client/test_sse.py create mode 100644 tests/server/test_http_dependencies.py diff --git a/src/fastmcp/client/client.py b/src/fastmcp/client/client.py index 32536e13c..000b99c00 100644 --- a/src/fastmcp/client/client.py +++ b/src/fastmcp/client/client.py @@ -108,9 +108,10 @@ class Client: # --- MCP Client Methods --- - async def ping(self) -> None: + async def ping(self) -> bool: """Send a ping request.""" - await self.session.send_ping() + result = await self.session.send_ping() + return isinstance(result, mcp.types.EmptyResult) async def progress( self, diff --git a/src/fastmcp/utilities/tests.py b/src/fastmcp/utilities/tests.py index 773845e4d..fe259c0ce 100644 --- a/src/fastmcp/utilities/tests.py +++ b/src/fastmcp/utilities/tests.py @@ -1,9 +1,20 @@ +from __future__ import annotations + import copy +import multiprocessing +import socket +import time +from collections.abc import Callable, Generator from contextlib import contextmanager -from typing import Any +from typing import TYPE_CHECKING, Any, Literal + +import uvicorn from fastmcp.settings import settings +if TYPE_CHECKING: + from fastmcp.server.server import FastMCP + @contextmanager def temporary_settings(**kwargs: Any): @@ -39,3 +50,64 @@ def temporary_settings(**kwargs: Any): for attr in kwargs: if hasattr(settings, attr): setattr(settings, attr, old_settings[attr]) + + +def _run_server(mcp_server: FastMCP, transport: Literal["sse"], port: int) -> None: + # Some Starlette apps are not pickleable, so we need to create them here based on the indicated transport + if transport == "sse": + app = mcp_server.sse_app() + else: + raise ValueError(f"Invalid transport: {transport}") + uvicorn_server = uvicorn.Server( + config=uvicorn.Config( + app=app, + host="127.0.0.1", + port=port, + log_level="error", + ) + ) + uvicorn_server.run() + + +@contextmanager +def run_server_in_process( + server_fn: Callable[[str, int], None], +) -> Generator[str, None, None]: + """ + Context manager that runs a Starlette app in a separate process and returns the + server URL. When the context manager is exited, the server process is killed. + + Args: + app: The Starlette app to run. + + Returns: + The server URL. + """ + host = "127.0.0.1" + with socket.socket() as s: + s.bind((host, 0)) + port = s.getsockname()[1] + + proc = multiprocessing.Process(target=server_fn, args=(host, port), daemon=True) + proc.start() + + # Wait for server to be running + max_attempts = 100 + attempt = 0 + while attempt < max_attempts and proc.is_alive(): + try: + with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s: + s.connect((host, port)) + break + except ConnectionRefusedError: + time.sleep(0.01) + attempt += 1 + else: + raise RuntimeError(f"Server failed to start after {max_attempts} attempts") + + yield f"http://{host}:{port}" + + proc.kill() + proc.join(timeout=2) + if proc.is_alive(): + raise RuntimeError("Server process failed to terminate") diff --git a/tests/client/test_sse.py b/tests/client/test_sse.py new file mode 100644 index 000000000..58ba87c79 --- /dev/null +++ b/tests/client/test_sse.py @@ -0,0 +1,92 @@ +import json +import sys +from collections.abc import Generator + +import pytest +import uvicorn +from mcp.types import TextResourceContents + +from fastmcp.client import Client +from fastmcp.client.transports import SSETransport +from fastmcp.server.dependencies import get_http_request +from fastmcp.server.server import FastMCP +from fastmcp.utilities.tests import run_server_in_process + + +def fastmcp_server(): + """Fixture that creates a FastMCP server with tools, resources, and prompts.""" + server = FastMCP("TestServer") + + # Add a tool + @server.tool() + def greet(name: str) -> str: + """Greet someone by name.""" + return f"Hello, {name}!" + + # Add a second tool + @server.tool() + def add(a: int, b: int) -> int: + """Add two numbers together.""" + return a + b + + # Add a resource + @server.resource(uri="data://users") + async def get_users(): + return ["Alice", "Bob", "Charlie"] + + # Add a resource template + @server.resource(uri="data://user/{user_id}") + async def get_user(user_id: str): + return {"id": user_id, "name": f"User {user_id}", "active": True} + + @server.resource(uri="request://headers") + async def get_headers() -> dict[str, str]: + request = get_http_request() + + return dict(request.headers) + + # Add a prompt + @server.prompt() + def welcome(name: str) -> str: + """Example greeting prompt.""" + return f"Welcome to FastMCP, {name}!" + + return server + + +def run_server(host: str, port: int) -> None: + try: + app = fastmcp_server().sse_app() + server = uvicorn.Server( + config=uvicorn.Config(app=app, host=host, port=port, log_level="error") + ) + server.run() + except Exception as e: + print(f"Server error: {e}") + sys.exit(1) + sys.exit(0) + + +@pytest.fixture(autouse=True, scope="module") +def sse_server() -> Generator[str, None, None]: + with run_server_in_process(run_server) as url: + yield f"{url}/sse" + + +async def test_ping(sse_server: str): + """Test pinging the server.""" + async with Client(transport=SSETransport(sse_server)) as client: + result = await client.ping() + assert result is True + + +async def test_http_headers(sse_server: str): + """Test getting HTTP headers from the server.""" + async with Client( + transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"}) + ) as client: + raw_result = await client.read_resource("request://headers") + assert isinstance(raw_result[0], TextResourceContents) + json_result = json.loads(raw_result[0].text) + assert "x-demo-header" in json_result + assert json_result["x-demo-header"] == "ABC" diff --git a/tests/server/test_http_dependencies.py b/tests/server/test_http_dependencies.py new file mode 100644 index 000000000..0e9c61ba5 --- /dev/null +++ b/tests/server/test_http_dependencies.py @@ -0,0 +1,96 @@ +import json +import sys +from collections.abc import Generator + +import pytest +import uvicorn +from mcp.types import TextContent, TextResourceContents + +from fastmcp.client import Client +from fastmcp.client.transports import SSETransport +from fastmcp.server.dependencies import get_http_request +from fastmcp.server.server import FastMCP +from fastmcp.utilities.tests import run_server_in_process + + +def fastmcp_server(): + server = FastMCP() + + # Add a tool + @server.tool() + def get_headers_tool() -> dict[str, str]: + """Get the HTTP headers from the request.""" + request = get_http_request() + + return dict(request.headers) + + @server.resource(uri="request://headers") + async def get_headers_resource() -> dict[str, str]: + request = get_http_request() + + return dict(request.headers) + + # Add a prompt + @server.prompt() + def get_headers_prompt() -> str: + """Get the HTTP headers from the request.""" + request = get_http_request() + + return json.dumps(dict(request.headers)) + + return server + + +def run_server(host: str, port: int) -> None: + try: + app = fastmcp_server().sse_app() + server = uvicorn.Server( + config=uvicorn.Config(app=app, host=host, port=port, log_level="error") + ) + server.run() + except Exception as e: + print(f"Server error: {e}") + sys.exit(1) + sys.exit(0) + + +@pytest.fixture(autouse=True, scope="module") +def sse_server() -> Generator[str, None, None]: + with run_server_in_process(run_server) as url: + yield f"{url}/sse" + + +async def test_http_headers_resource(sse_server: str): + """Test getting HTTP headers from the server.""" + async with Client( + transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"}) + ) as client: + raw_result = await client.read_resource("request://headers") + assert isinstance(raw_result[0], TextResourceContents) + json_result = json.loads(raw_result[0].text) + assert "x-demo-header" in json_result + assert json_result["x-demo-header"] == "ABC" + + +async def test_http_headers_tool(sse_server: str): + """Test getting HTTP headers from the server.""" + async with Client( + transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"}) + ) as client: + result = await client.call_tool("get_headers_tool") + assert isinstance(result[0], TextContent) + json_result = json.loads(result[0].text) + assert "x-demo-header" in json_result + assert json_result["x-demo-header"] == "ABC" + + +async def test_http_headers_prompt(sse_server: str): + """Test getting HTTP headers from the server.""" + async with Client( + transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"}) + ) as client: + result = await client.get_prompt("get_headers_prompt") + assert isinstance(result.messages[0].content, TextContent) + json_result = json.loads(result.messages[0].content.text) + assert "x-demo-header" in json_result + assert json_result["x-demo-header"] == "ABC"