mirror of
https://github.com/PrefectHQ/fastmcp.git
synced 2026-08-22 21:44:18 +02:00
Add integration tests for SSE
This commit is contained in:
parent
e5748fffb0
commit
9c936168cd
4 changed files with 264 additions and 3 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
92
tests/client/test_sse.py
Normal file
92
tests/client/test_sse.py
Normal file
|
|
@ -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"
|
||||
96
tests/server/test_http_dependencies.py
Normal file
96
tests/server/test_http_dependencies.py
Normal file
|
|
@ -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"
|
||||
Loading…
Add table
Add a link
Reference in a new issue