You can not select more than 25 topics Topics must start with a chinese character,a letter or number, can include dashes ('-') and can be up to 35 characters long.

test_mcp_tools.py 9.9 kB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271
  1. from unittest.mock import AsyncMock, MagicMock
  2. import pytest
  3. from autogen_core import CancellationToken
  4. from autogen_ext.tools.mcp import (
  5. SseMcpToolAdapter,
  6. SseServerParams,
  7. StdioMcpToolAdapter,
  8. StdioServerParams,
  9. )
  10. from json_schema_to_pydantic import create_model
  11. from mcp import ClientSession, Tool
  12. @pytest.fixture
  13. def sample_tool() -> Tool:
  14. return Tool(
  15. name="test_tool",
  16. description="A test tool",
  17. inputSchema={
  18. "type": "object",
  19. "properties": {"test_param": {"type": "string"}},
  20. "required": ["test_param"],
  21. },
  22. )
  23. @pytest.fixture
  24. def sample_server_params() -> StdioServerParams:
  25. return StdioServerParams(command="echo", args=["test"])
  26. @pytest.fixture
  27. def sample_sse_tool() -> Tool:
  28. return Tool(
  29. name="test_sse_tool",
  30. description="A test SSE tool",
  31. inputSchema={
  32. "type": "object",
  33. "properties": {"test_param": {"type": "string"}},
  34. "required": ["test_param"],
  35. },
  36. )
  37. @pytest.fixture
  38. def mock_sse_session() -> AsyncMock:
  39. session = AsyncMock(spec=ClientSession)
  40. session.initialize = AsyncMock()
  41. session.call_tool = AsyncMock()
  42. session.list_tools = AsyncMock()
  43. return session
  44. @pytest.fixture
  45. def mock_session() -> AsyncMock:
  46. session = AsyncMock(spec=ClientSession)
  47. session.initialize = AsyncMock()
  48. session.call_tool = AsyncMock()
  49. session.list_tools = AsyncMock()
  50. return session
  51. @pytest.fixture
  52. def mock_tool_response() -> MagicMock:
  53. response = MagicMock()
  54. response.isError = False
  55. response.content = {"result": "test_output"}
  56. return response
  57. @pytest.fixture
  58. def cancellation_token() -> CancellationToken:
  59. return CancellationToken()
  60. def test_adapter_config_serialization(sample_tool: Tool, sample_server_params: StdioServerParams) -> None:
  61. """Test that adapter can be saved to and loaded from config."""
  62. original_adapter = StdioMcpToolAdapter(server_params=sample_server_params, tool=sample_tool)
  63. config = original_adapter.dump_component()
  64. loaded_adapter = StdioMcpToolAdapter.load_component(config)
  65. # Test that the loaded adapter has the same properties
  66. assert loaded_adapter.name == "test_tool"
  67. assert loaded_adapter.description == "A test tool"
  68. # Verify schema structure
  69. schema = loaded_adapter.schema
  70. assert "parameters" in schema, "Schema must have parameters"
  71. params_schema = schema["parameters"]
  72. assert isinstance(params_schema, dict), "Parameters must be a dict"
  73. assert "type" in params_schema, "Parameters must have type"
  74. assert "required" in params_schema, "Parameters must have required fields"
  75. assert "properties" in params_schema, "Parameters must have properties"
  76. # Compare schema content
  77. assert params_schema["type"] == sample_tool.inputSchema["type"]
  78. assert params_schema["required"] == sample_tool.inputSchema["required"]
  79. assert (
  80. params_schema["properties"]["test_param"]["type"] == sample_tool.inputSchema["properties"]["test_param"]["type"]
  81. )
  82. @pytest.mark.asyncio
  83. async def test_mcp_tool_execution(
  84. sample_tool: Tool,
  85. sample_server_params: StdioServerParams,
  86. mock_session: AsyncMock,
  87. mock_tool_response: MagicMock,
  88. cancellation_token: CancellationToken,
  89. monkeypatch: pytest.MonkeyPatch,
  90. ) -> None:
  91. """Test that adapter properly executes tools through ClientSession."""
  92. mock_context = AsyncMock()
  93. mock_context.__aenter__.return_value = mock_session
  94. monkeypatch.setattr(
  95. "autogen_ext.tools.mcp._base.create_mcp_server_session",
  96. lambda *args, **kwargs: mock_context, # type: ignore
  97. )
  98. mock_session.call_tool.return_value = mock_tool_response
  99. adapter = StdioMcpToolAdapter(server_params=sample_server_params, tool=sample_tool)
  100. result = await adapter.run(
  101. args=create_model(sample_tool.inputSchema)(**{"test_param": "test"}),
  102. cancellation_token=cancellation_token,
  103. )
  104. assert result == mock_tool_response.content
  105. mock_session.initialize.assert_called_once()
  106. mock_session.call_tool.assert_called_once()
  107. @pytest.mark.asyncio
  108. async def test_adapter_from_server_params(
  109. sample_tool: Tool,
  110. sample_server_params: StdioServerParams,
  111. mock_session: AsyncMock,
  112. monkeypatch: pytest.MonkeyPatch,
  113. ) -> None:
  114. """Test that adapter can be created from server parameters."""
  115. mock_context = AsyncMock()
  116. mock_context.__aenter__.return_value = mock_session
  117. monkeypatch.setattr(
  118. "autogen_ext.tools.mcp._base.create_mcp_server_session",
  119. lambda *args, **kwargs: mock_context, # type: ignore
  120. )
  121. mock_session.list_tools.return_value.tools = [sample_tool]
  122. adapter = await StdioMcpToolAdapter.from_server_params(sample_server_params, "test_tool")
  123. assert isinstance(adapter, StdioMcpToolAdapter)
  124. assert adapter.name == "test_tool"
  125. assert adapter.description == "A test tool"
  126. # Verify schema structure
  127. schema = adapter.schema
  128. assert "parameters" in schema, "Schema must have parameters"
  129. params_schema = schema["parameters"]
  130. assert isinstance(params_schema, dict), "Parameters must be a dict"
  131. assert "type" in params_schema, "Parameters must have type"
  132. assert "required" in params_schema, "Parameters must have required fields"
  133. assert "properties" in params_schema, "Parameters must have properties"
  134. # Compare schema content
  135. assert params_schema["type"] == sample_tool.inputSchema["type"]
  136. assert params_schema["required"] == sample_tool.inputSchema["required"]
  137. assert (
  138. params_schema["properties"]["test_param"]["type"] == sample_tool.inputSchema["properties"]["test_param"]["type"]
  139. )
  140. @pytest.mark.asyncio
  141. async def test_sse_adapter_config_serialization(sample_sse_tool: Tool) -> None:
  142. """Test that SSE adapter can be saved to and loaded from config."""
  143. params = SseServerParams(url="http://test-url")
  144. original_adapter = SseMcpToolAdapter(server_params=params, tool=sample_sse_tool)
  145. config = original_adapter.dump_component()
  146. loaded_adapter = SseMcpToolAdapter.load_component(config)
  147. # Test that the loaded adapter has the same properties
  148. assert loaded_adapter.name == "test_sse_tool"
  149. assert loaded_adapter.description == "A test SSE tool"
  150. # Verify schema structure
  151. schema = loaded_adapter.schema
  152. assert "parameters" in schema, "Schema must have parameters"
  153. params_schema = schema["parameters"]
  154. assert isinstance(params_schema, dict), "Parameters must be a dict"
  155. assert "type" in params_schema, "Parameters must have type"
  156. assert "required" in params_schema, "Parameters must have required fields"
  157. assert "properties" in params_schema, "Parameters must have properties"
  158. # Compare schema content
  159. assert params_schema["type"] == sample_sse_tool.inputSchema["type"]
  160. assert params_schema["required"] == sample_sse_tool.inputSchema["required"]
  161. assert (
  162. params_schema["properties"]["test_param"]["type"]
  163. == sample_sse_tool.inputSchema["properties"]["test_param"]["type"]
  164. )
  165. @pytest.mark.asyncio
  166. async def test_sse_tool_execution(
  167. sample_sse_tool: Tool,
  168. mock_sse_session: AsyncMock,
  169. monkeypatch: pytest.MonkeyPatch,
  170. ) -> None:
  171. """Test that SSE adapter properly executes tools through ClientSession."""
  172. params = SseServerParams(url="http://test-url")
  173. mock_context = AsyncMock()
  174. mock_context.__aenter__.return_value = mock_sse_session
  175. mock_sse_session.call_tool.return_value = MagicMock(isError=False, content={"result": "test_output"})
  176. monkeypatch.setattr(
  177. "autogen_ext.tools.mcp._base.create_mcp_server_session",
  178. lambda *args, **kwargs: mock_context, # type: ignore
  179. )
  180. adapter = SseMcpToolAdapter(server_params=params, tool=sample_sse_tool)
  181. result = await adapter.run(
  182. args=create_model(sample_sse_tool.inputSchema)(**{"test_param": "test"}),
  183. cancellation_token=CancellationToken(),
  184. )
  185. assert result == mock_sse_session.call_tool.return_value.content
  186. mock_sse_session.initialize.assert_called_once()
  187. mock_sse_session.call_tool.assert_called_once()
  188. @pytest.mark.asyncio
  189. async def test_sse_adapter_from_server_params(
  190. sample_sse_tool: Tool,
  191. mock_sse_session: AsyncMock,
  192. monkeypatch: pytest.MonkeyPatch,
  193. ) -> None:
  194. """Test that SSE adapter can be created from server parameters."""
  195. params = SseServerParams(url="http://test-url")
  196. mock_context = AsyncMock()
  197. mock_context.__aenter__.return_value = mock_sse_session
  198. monkeypatch.setattr(
  199. "autogen_ext.tools.mcp._base.create_mcp_server_session",
  200. lambda *args, **kwargs: mock_context, # type: ignore
  201. )
  202. mock_sse_session.list_tools.return_value.tools = [sample_sse_tool]
  203. adapter = await SseMcpToolAdapter.from_server_params(params, "test_sse_tool")
  204. assert isinstance(adapter, SseMcpToolAdapter)
  205. assert adapter.name == "test_sse_tool"
  206. assert adapter.description == "A test SSE tool"
  207. # Verify schema structure
  208. schema = adapter.schema
  209. assert "parameters" in schema, "Schema must have parameters"
  210. params_schema = schema["parameters"]
  211. assert isinstance(params_schema, dict), "Parameters must be a dict"
  212. assert "type" in params_schema, "Parameters must have type"
  213. assert "required" in params_schema, "Parameters must have required fields"
  214. assert "properties" in params_schema, "Parameters must have properties"
  215. # Compare schema content
  216. assert params_schema["type"] == sample_sse_tool.inputSchema["type"]
  217. assert params_schema["required"] == sample_sse_tool.inputSchema["required"]
  218. assert (
  219. params_schema["properties"]["test_param"]["type"]
  220. == sample_sse_tool.inputSchema["properties"]["test_param"]["type"]
  221. )