test_server_tools.py 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385
  1. from __future__ import annotations
  2. import asyncio
  3. import inspect
  4. import unittest
  5. from unittest.mock import patch
  6. from data_collector_mcp import app, common_server, modbus_server, s7_server
  7. class ServerToolTests(unittest.TestCase):
  8. def test_app_imports_tool_modules_for_registration(self) -> None:
  9. self.assertIs(app.common_server, common_server)
  10. self.assertIs(app.modbus_server, modbus_server)
  11. self.assertIs(app.s7_server, s7_server)
  12. def test_app_registers_expected_mcp_tools(self) -> None:
  13. tools = asyncio.run(app.mcp.list_tools())
  14. tool_names = {tool.name for tool in tools}
  15. self.assertEqual(
  16. tool_names,
  17. {
  18. "project.list",
  19. "collector.device_list",
  20. "collector.device_connect",
  21. "collector.device_disconnect",
  22. "collector.device_points",
  23. "modbus.point_collect_test",
  24. "s7.raw_read",
  25. "s7.point_collect_test",
  26. "s7.connect_scan",
  27. "collector.modbus_device_create",
  28. "collector.modbus_device_edit",
  29. "collector.modbus_point_create",
  30. "collector.modbus_point_edit",
  31. "collector.s7_device_create",
  32. "collector.s7_device_edit",
  33. "collector.s7_point_create",
  34. "collector.s7_point_edit",
  35. },
  36. )
  37. def test_project_list_filters_enabled_projects_and_sorts(self) -> None:
  38. with patch(
  39. "data_collector_mcp.common_server.load_projects_config",
  40. return_value=[
  41. {
  42. "project_key": "z-prod",
  43. "project_name": "Prod",
  44. "base_url": "http://gateway.prod",
  45. "data_collector_base_url": "http://collector.prod",
  46. "enabled": True,
  47. },
  48. {
  49. "project_key": "disabled",
  50. "project_name": "Disabled",
  51. "base_url": "http://gateway.disabled",
  52. "data_collector_base_url": "http://collector.disabled",
  53. "enabled": False,
  54. },
  55. {
  56. "project_key": "a-dev",
  57. "project_name": "Dev",
  58. "base_url": "http://gateway.dev",
  59. "data_collector_base_url": "http://collector.dev",
  60. "enabled": True,
  61. },
  62. ],
  63. ):
  64. result = common_server.project_list()
  65. self.assertEqual(
  66. result,
  67. {
  68. "projects": [
  69. {
  70. "project_key": "a-dev",
  71. "project_name": "Dev",
  72. "base_url": "http://gateway.dev",
  73. "data_collector_base_url": "http://collector.dev",
  74. },
  75. {
  76. "project_key": "z-prod",
  77. "project_name": "Prod",
  78. "base_url": "http://gateway.prod",
  79. "data_collector_base_url": "http://collector.prod",
  80. },
  81. ],
  82. "total": 2,
  83. },
  84. )
  85. def test_common_device_tools_forward_to_api(self) -> None:
  86. with patch("data_collector_mcp.common_server.api_list_devices", return_value={"state": 0}) as list_devices:
  87. self.assertEqual(common_server.collector_device_list("dev-01", num_points=True), {"state": 0})
  88. list_devices.assert_called_once_with("dev-01", num_points=True)
  89. with patch("data_collector_mcp.common_server.api_connect_device", return_value={"state": 0}) as connect:
  90. self.assertEqual(
  91. common_server.collector_device_connect("dev-01", device_id=1, device_type="s7"),
  92. {"state": 0},
  93. )
  94. connect.assert_called_once_with("dev-01", device_id=1, device_type="s7")
  95. with patch("data_collector_mcp.common_server.api_disconnect_device", return_value={"state": 0}) as disconnect:
  96. self.assertEqual(
  97. common_server.collector_device_disconnect("dev-01", device_id=1, device_type="modbus"),
  98. {"state": 0},
  99. )
  100. disconnect.assert_called_once_with("dev-01", device_id=1, device_type="modbus")
  101. with patch("data_collector_mcp.common_server.api_list_device_points", return_value={"state": 0}) as points:
  102. self.assertEqual(
  103. common_server.collector_device_points("dev-01", device_id=1, device_type="s7", group_id=2),
  104. {"state": 0},
  105. )
  106. points.assert_called_once_with("dev-01", device_id=1, device_type="s7", group_id=2)
  107. def test_modbus_device_create_accepts_batch_devices(self) -> None:
  108. signature = inspect.signature(modbus_server.collector_modbus_device_create)
  109. self.assertIn("devices", signature.parameters)
  110. devices = [
  111. {
  112. "name": "modbus_tcp_1",
  113. "device_type": 1,
  114. "ip": "127.0.0.1",
  115. "port": 5502,
  116. "slave_id": 1,
  117. "byte_order": 1,
  118. "word_order": 1,
  119. "address_base": 0,
  120. "group_id": 10,
  121. }
  122. ]
  123. with patch("data_collector_mcp.modbus_server.api_create_modbus_devices", return_value={"state": 0}) as api_create:
  124. result = modbus_server.collector_modbus_device_create(
  125. project_key="dev-01",
  126. devices=devices,
  127. )
  128. self.assertEqual(result, {"state": 0})
  129. api_create.assert_called_once_with("dev-01", devices)
  130. def test_modbus_point_create_accepts_batch_points(self) -> None:
  131. points = [{"device_id": 1, "name": "temperature", "address": 10, "type": "uint16", "func_code": 3}]
  132. with patch("data_collector_mcp.modbus_server.api_create_modbus_points", return_value={"state": 0}) as api_create:
  133. result = modbus_server.collector_modbus_point_create(project_key="dev-01", points=points)
  134. self.assertEqual(result, {"state": 0})
  135. api_create.assert_called_once_with("dev-01", points)
  136. def test_modbus_device_edit_uses_explicit_parameters(self) -> None:
  137. signature = inspect.signature(modbus_server.collector_modbus_device_edit)
  138. self.assertNotIn("payload", signature.parameters)
  139. with patch("data_collector_mcp.modbus_server.api_edit_modbus_device", return_value={"state": 0}) as api_edit:
  140. result = modbus_server.collector_modbus_device_edit(
  141. project_key="dev-01",
  142. ori_id=1,
  143. name="modbus_tcp_edited",
  144. device_type=1,
  145. ip="127.0.0.1",
  146. port=5502,
  147. slave_id=1,
  148. byte_order=2,
  149. word_order=2,
  150. address_offset=1,
  151. device_group_id=10,
  152. )
  153. self.assertEqual(result, {"state": 0})
  154. api_edit.assert_called_once_with(
  155. "dev-01",
  156. {
  157. "ori_id": 1,
  158. "name": "modbus_tcp_edited",
  159. "device_type": 1,
  160. "ip": "127.0.0.1",
  161. "port": 5502,
  162. "slave_id": 1,
  163. "byte_order": 2,
  164. "word_order": 2,
  165. "serial_port": "",
  166. "timeout": 3,
  167. "is_persistent": False,
  168. "baud_rate": 0,
  169. "data_bit": 0,
  170. "parity": 0,
  171. "stop_bit": 0,
  172. "mode": 0,
  173. "address_offset": 1,
  174. "retry_times": 0,
  175. "device_group_id": 10,
  176. "alarm_interval": 90,
  177. "collect_interval": 5,
  178. },
  179. )
  180. def test_modbus_point_edit_uses_explicit_parameters_with_func_code(self) -> None:
  181. signature = inspect.signature(modbus_server.collector_modbus_point_edit)
  182. self.assertNotIn("payload", signature.parameters)
  183. with patch("data_collector_mcp.modbus_server.api_edit_modbus_point", return_value={"state": 0}) as api_edit:
  184. result = modbus_server.collector_modbus_point_edit(
  185. project_key="dev-01",
  186. ori_id=101,
  187. name="holding_register_uint16_edited",
  188. address=10,
  189. data_type="uint16",
  190. func_code=3,
  191. point_id="HR_UINT16_EDITED",
  192. )
  193. self.assertEqual(result, {"state": 0})
  194. api_edit.assert_called_once_with(
  195. "dev-01",
  196. {
  197. "ori_id": 101,
  198. "name": "holding_register_uint16_edited",
  199. "address": 10,
  200. "type": "uint16",
  201. "point_id": "HR_UINT16_EDITED",
  202. "scale_ratio": 1,
  203. "value_offset": 0,
  204. "group_id": 0,
  205. "invalid_values": "",
  206. "valid_range_start": None,
  207. "valid_range_end": None,
  208. "bit": 0,
  209. "describe": "",
  210. "func_code": 3,
  211. },
  212. )
  213. def test_modbus_point_edit_uses_register_type_when_func_code_is_zero(self) -> None:
  214. with patch("data_collector_mcp.modbus_server.api_edit_modbus_point", return_value={"state": 0}) as api_edit:
  215. modbus_server.collector_modbus_point_edit(
  216. project_key="dev-01",
  217. ori_id=101,
  218. name="holding_register_uint16_edited",
  219. address=10,
  220. data_type="uint16",
  221. register_type="holding_register",
  222. )
  223. payload = api_edit.call_args.args[1]
  224. self.assertEqual(payload["register_type"], "holding_register")
  225. self.assertNotIn("func_code", payload)
  226. def test_s7_device_create_accepts_batch_devices(self) -> None:
  227. signature = inspect.signature(s7_server.collector_s7_device_create)
  228. self.assertIn("devices", signature.parameters)
  229. devices = [{"name": "s7_1200_1", "ip": "127.0.0.1", "rock": 0, "slot": 1, "tsap_conn_type": "OP"}]
  230. with patch("data_collector_mcp.s7_server.api_create_s7_devices", return_value={"state": 0}) as api_create:
  231. result = s7_server.collector_s7_device_create(
  232. project_key="dev-01",
  233. devices=devices,
  234. )
  235. self.assertEqual(result, {"state": 0})
  236. api_create.assert_called_once_with(
  237. "dev-01",
  238. [{"name": "s7_1200_1", "ip": "127.0.0.1", "rock": 0, "slot": 1, "tsap_conn_type": "OP"}],
  239. )
  240. def test_s7_point_create_accepts_batch_points(self) -> None:
  241. points = [{"device_id": 3, "name": "db_real", "address": "1.10", "type": "REAL", "register_area": "DB"}]
  242. with patch("data_collector_mcp.s7_server.api_create_s7_points", return_value={"state": 0}) as api_create:
  243. result = s7_server.collector_s7_point_create(project_key="dev-01", points=points)
  244. self.assertEqual(result, {"state": 0})
  245. api_create.assert_called_once_with("dev-01", points)
  246. def test_s7_device_edit_uses_explicit_parameters(self) -> None:
  247. signature = inspect.signature(s7_server.collector_s7_device_edit)
  248. self.assertNotIn("payload", signature.parameters)
  249. with patch("data_collector_mcp.s7_server.api_edit_s7_device", return_value={"state": 0}) as api_edit:
  250. result = s7_server.collector_s7_device_edit(
  251. project_key="dev-01",
  252. ori_id=3,
  253. name="s7_edited",
  254. ip="127.0.0.1",
  255. rock=0,
  256. slot=1,
  257. tsap_conn_type="OP",
  258. )
  259. self.assertEqual(result, {"state": 0})
  260. api_edit.assert_called_once_with(
  261. "dev-01",
  262. {
  263. "ori_id": 3,
  264. "name": "s7_edited",
  265. "ip": "127.0.0.1",
  266. "rock": 0,
  267. "slot": 1,
  268. "port": 102,
  269. "device_type": 1,
  270. "tsap_conn_type": "OP",
  271. "is_persistent": False,
  272. "device_group_id": 0,
  273. "timeout": 3,
  274. "alarm_interval": 90,
  275. "collect_interval": 5,
  276. },
  277. )
  278. def test_s7_point_edit_uses_register_area_when_register_type_is_zero(self) -> None:
  279. signature = inspect.signature(s7_server.collector_s7_point_edit)
  280. self.assertNotIn("payload", signature.parameters)
  281. with patch("data_collector_mcp.s7_server.api_edit_s7_point", return_value={"state": 0}) as api_edit:
  282. s7_server.collector_s7_point_edit(
  283. project_key="dev-01",
  284. ori_id=101,
  285. device_id=3,
  286. name="db_real",
  287. address="1.10",
  288. data_type="float32",
  289. register_area="DB",
  290. point_id="DB_REAL",
  291. )
  292. payload = api_edit.call_args.args[1]
  293. self.assertEqual(payload["register_area"], "DB")
  294. self.assertNotIn("register_type", payload)
  295. def test_s7_gateway_tools_forward_to_api(self) -> None:
  296. with patch("data_collector_mcp.s7_server.api_s7_raw_read", return_value={"code": 0}) as raw_read:
  297. self.assertEqual(
  298. s7_server.s7_raw_read(
  299. project_key="dev-01",
  300. ip="192.168.1.10",
  301. rock=0,
  302. slot=1,
  303. read={"area": "DB", "db": 1, "start": 0, "size": 4},
  304. ),
  305. {"code": 0},
  306. )
  307. raw_read.assert_called_once_with(
  308. "dev-01",
  309. ip="192.168.1.10",
  310. rock=0,
  311. slot=1,
  312. read={"area": "DB", "db": 1, "start": 0, "size": 4},
  313. device_type="S7-1200",
  314. port=102,
  315. tsap_conn_type=None,
  316. )
  317. with patch("data_collector_mcp.s7_server.api_s7_point_collect_test", return_value={"code": 0}) as point_test:
  318. self.assertEqual(
  319. s7_server.s7_point_collect_test(
  320. project_key="dev-01",
  321. ip="192.168.1.10",
  322. port=1102,
  323. rock=0,
  324. slot=1,
  325. tsap_conn_type="BASIC",
  326. points=[{"area": "M", "start": 0, "type": "bool"}],
  327. ),
  328. {"code": 0},
  329. )
  330. point_test.assert_called_once_with(
  331. "dev-01",
  332. ip="192.168.1.10",
  333. rock=0,
  334. slot=1,
  335. points=[{"area": "M", "start": 0, "type": "bool"}],
  336. device_type="S7-1200",
  337. port=1102,
  338. tsap_conn_type="BASIC",
  339. )
  340. with patch("data_collector_mcp.s7_server.api_s7_connect_scan", return_value={"code": 0}) as connect_scan:
  341. self.assertEqual(s7_server.s7_connect_scan("dev-01", ip="192.168.1.10"), {"code": 0})
  342. connect_scan.assert_called_once_with("dev-01", ip="192.168.1.10")
  343. if __name__ == "__main__":
  344. unittest.main()