Jeremiah Lowin commited on
Commit
73679eb
·
1 Parent(s): b101b5f

Add integration tests for SSE

Browse files
src/fastmcp/client/client.py CHANGED
@@ -108,9 +108,10 @@ class Client:
108
 
109
  # --- MCP Client Methods ---
110
 
111
- async def ping(self) -> None:
112
  """Send a ping request."""
113
- await self.session.send_ping()
 
114
 
115
  async def progress(
116
  self,
 
108
 
109
  # --- MCP Client Methods ---
110
 
111
+ async def ping(self) -> bool:
112
  """Send a ping request."""
113
+ result = await self.session.send_ping()
114
+ return isinstance(result, mcp.types.EmptyResult)
115
 
116
  async def progress(
117
  self,
src/fastmcp/utilities/tests.py CHANGED
@@ -1,9 +1,20 @@
 
 
1
  import copy
 
 
 
 
2
  from contextlib import contextmanager
3
- from typing import Any
 
 
4
 
5
  from fastmcp.settings import settings
6
 
 
 
 
7
 
8
  @contextmanager
9
  def temporary_settings(**kwargs: Any):
@@ -39,3 +50,64 @@ def temporary_settings(**kwargs: Any):
39
  for attr in kwargs:
40
  if hasattr(settings, attr):
41
  setattr(settings, attr, old_settings[attr])
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
  import copy
4
+ import multiprocessing
5
+ import socket
6
+ import time
7
+ from collections.abc import Callable, Generator
8
  from contextlib import contextmanager
9
+ from typing import TYPE_CHECKING, Any, Literal
10
+
11
+ import uvicorn
12
 
13
  from fastmcp.settings import settings
14
 
15
+ if TYPE_CHECKING:
16
+ from fastmcp.server.server import FastMCP
17
+
18
 
19
  @contextmanager
20
  def temporary_settings(**kwargs: Any):
 
50
  for attr in kwargs:
51
  if hasattr(settings, attr):
52
  setattr(settings, attr, old_settings[attr])
53
+
54
+
55
+ def _run_server(mcp_server: FastMCP, transport: Literal["sse"], port: int) -> None:
56
+ # Some Starlette apps are not pickleable, so we need to create them here based on the indicated transport
57
+ if transport == "sse":
58
+ app = mcp_server.sse_app()
59
+ else:
60
+ raise ValueError(f"Invalid transport: {transport}")
61
+ uvicorn_server = uvicorn.Server(
62
+ config=uvicorn.Config(
63
+ app=app,
64
+ host="127.0.0.1",
65
+ port=port,
66
+ log_level="error",
67
+ )
68
+ )
69
+ uvicorn_server.run()
70
+
71
+
72
+ @contextmanager
73
+ def run_server_in_process(
74
+ server_fn: Callable[[str, int], None],
75
+ ) -> Generator[str, None, None]:
76
+ """
77
+ Context manager that runs a Starlette app in a separate process and returns the
78
+ server URL. When the context manager is exited, the server process is killed.
79
+
80
+ Args:
81
+ app: The Starlette app to run.
82
+
83
+ Returns:
84
+ The server URL.
85
+ """
86
+ host = "127.0.0.1"
87
+ with socket.socket() as s:
88
+ s.bind((host, 0))
89
+ port = s.getsockname()[1]
90
+
91
+ proc = multiprocessing.Process(target=server_fn, args=(host, port), daemon=True)
92
+ proc.start()
93
+
94
+ # Wait for server to be running
95
+ max_attempts = 100
96
+ attempt = 0
97
+ while attempt < max_attempts and proc.is_alive():
98
+ try:
99
+ with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as s:
100
+ s.connect((host, port))
101
+ break
102
+ except ConnectionRefusedError:
103
+ time.sleep(0.01)
104
+ attempt += 1
105
+ else:
106
+ raise RuntimeError(f"Server failed to start after {max_attempts} attempts")
107
+
108
+ yield f"http://{host}:{port}"
109
+
110
+ proc.kill()
111
+ proc.join(timeout=2)
112
+ if proc.is_alive():
113
+ raise RuntimeError("Server process failed to terminate")
tests/client/test_sse.py ADDED
@@ -0,0 +1,92 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ import sys
3
+ from collections.abc import Generator
4
+
5
+ import pytest
6
+ import uvicorn
7
+ from mcp.types import TextResourceContents
8
+
9
+ from fastmcp.client import Client
10
+ from fastmcp.client.transports import SSETransport
11
+ from fastmcp.server.dependencies import get_http_request
12
+ from fastmcp.server.server import FastMCP
13
+ from fastmcp.utilities.tests import run_server_in_process
14
+
15
+
16
+ def fastmcp_server():
17
+ """Fixture that creates a FastMCP server with tools, resources, and prompts."""
18
+ server = FastMCP("TestServer")
19
+
20
+ # Add a tool
21
+ @server.tool()
22
+ def greet(name: str) -> str:
23
+ """Greet someone by name."""
24
+ return f"Hello, {name}!"
25
+
26
+ # Add a second tool
27
+ @server.tool()
28
+ def add(a: int, b: int) -> int:
29
+ """Add two numbers together."""
30
+ return a + b
31
+
32
+ # Add a resource
33
+ @server.resource(uri="data://users")
34
+ async def get_users():
35
+ return ["Alice", "Bob", "Charlie"]
36
+
37
+ # Add a resource template
38
+ @server.resource(uri="data://user/{user_id}")
39
+ async def get_user(user_id: str):
40
+ return {"id": user_id, "name": f"User {user_id}", "active": True}
41
+
42
+ @server.resource(uri="request://headers")
43
+ async def get_headers() -> dict[str, str]:
44
+ request = get_http_request()
45
+
46
+ return dict(request.headers)
47
+
48
+ # Add a prompt
49
+ @server.prompt()
50
+ def welcome(name: str) -> str:
51
+ """Example greeting prompt."""
52
+ return f"Welcome to FastMCP, {name}!"
53
+
54
+ return server
55
+
56
+
57
+ def run_server(host: str, port: int) -> None:
58
+ try:
59
+ app = fastmcp_server().sse_app()
60
+ server = uvicorn.Server(
61
+ config=uvicorn.Config(app=app, host=host, port=port, log_level="error")
62
+ )
63
+ server.run()
64
+ except Exception as e:
65
+ print(f"Server error: {e}")
66
+ sys.exit(1)
67
+ sys.exit(0)
68
+
69
+
70
+ @pytest.fixture(autouse=True, scope="module")
71
+ def sse_server() -> Generator[str, None, None]:
72
+ with run_server_in_process(run_server) as url:
73
+ yield f"{url}/sse"
74
+
75
+
76
+ async def test_ping(sse_server: str):
77
+ """Test pinging the server."""
78
+ async with Client(transport=SSETransport(sse_server)) as client:
79
+ result = await client.ping()
80
+ assert result is True
81
+
82
+
83
+ async def test_http_headers(sse_server: str):
84
+ """Test getting HTTP headers from the server."""
85
+ async with Client(
86
+ transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
87
+ ) as client:
88
+ raw_result = await client.read_resource("request://headers")
89
+ assert isinstance(raw_result[0], TextResourceContents)
90
+ json_result = json.loads(raw_result[0].text)
91
+ assert "x-demo-header" in json_result
92
+ assert json_result["x-demo-header"] == "ABC"
tests/server/test_http_dependencies.py ADDED
@@ -0,0 +1,96 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ import sys
3
+ from collections.abc import Generator
4
+
5
+ import pytest
6
+ import uvicorn
7
+ from mcp.types import TextContent, TextResourceContents
8
+
9
+ from fastmcp.client import Client
10
+ from fastmcp.client.transports import SSETransport
11
+ from fastmcp.server.dependencies import get_http_request
12
+ from fastmcp.server.server import FastMCP
13
+ from fastmcp.utilities.tests import run_server_in_process
14
+
15
+
16
+ def fastmcp_server():
17
+ server = FastMCP()
18
+
19
+ # Add a tool
20
+ @server.tool()
21
+ def get_headers_tool() -> dict[str, str]:
22
+ """Get the HTTP headers from the request."""
23
+ request = get_http_request()
24
+
25
+ return dict(request.headers)
26
+
27
+ @server.resource(uri="request://headers")
28
+ async def get_headers_resource() -> dict[str, str]:
29
+ request = get_http_request()
30
+
31
+ return dict(request.headers)
32
+
33
+ # Add a prompt
34
+ @server.prompt()
35
+ def get_headers_prompt() -> str:
36
+ """Get the HTTP headers from the request."""
37
+ request = get_http_request()
38
+
39
+ return json.dumps(dict(request.headers))
40
+
41
+ return server
42
+
43
+
44
+ def run_server(host: str, port: int) -> None:
45
+ try:
46
+ app = fastmcp_server().sse_app()
47
+ server = uvicorn.Server(
48
+ config=uvicorn.Config(app=app, host=host, port=port, log_level="error")
49
+ )
50
+ server.run()
51
+ except Exception as e:
52
+ print(f"Server error: {e}")
53
+ sys.exit(1)
54
+ sys.exit(0)
55
+
56
+
57
+ @pytest.fixture(autouse=True, scope="module")
58
+ def sse_server() -> Generator[str, None, None]:
59
+ with run_server_in_process(run_server) as url:
60
+ yield f"{url}/sse"
61
+
62
+
63
+ async def test_http_headers_resource(sse_server: str):
64
+ """Test getting HTTP headers from the server."""
65
+ async with Client(
66
+ transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
67
+ ) as client:
68
+ raw_result = await client.read_resource("request://headers")
69
+ assert isinstance(raw_result[0], TextResourceContents)
70
+ json_result = json.loads(raw_result[0].text)
71
+ assert "x-demo-header" in json_result
72
+ assert json_result["x-demo-header"] == "ABC"
73
+
74
+
75
+ async def test_http_headers_tool(sse_server: str):
76
+ """Test getting HTTP headers from the server."""
77
+ async with Client(
78
+ transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
79
+ ) as client:
80
+ result = await client.call_tool("get_headers_tool")
81
+ assert isinstance(result[0], TextContent)
82
+ json_result = json.loads(result[0].text)
83
+ assert "x-demo-header" in json_result
84
+ assert json_result["x-demo-header"] == "ABC"
85
+
86
+
87
+ async def test_http_headers_prompt(sse_server: str):
88
+ """Test getting HTTP headers from the server."""
89
+ async with Client(
90
+ transport=SSETransport(sse_server, headers={"X-DEMO-HEADER": "ABC"})
91
+ ) as client:
92
+ result = await client.get_prompt("get_headers_prompt")
93
+ assert isinstance(result.messages[0].content, TextContent)
94
+ json_result = json.loads(result.messages[0].content.text)
95
+ assert "x-demo-header" in json_result
96
+ assert json_result["x-demo-header"] == "ABC"