test_server_tools.py 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404
  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_uses_explicit_parameters(self) -> None:
  108. signature = inspect.signature(modbus_server.collector_modbus_device_create)
  109. self.assertNotIn("payload", signature.parameters)
  110. with patch("data_collector_mcp.modbus_server.api_create_modbus_device", return_value={"state": 0}) as api_create:
  111. result = modbus_server.collector_modbus_device_create(
  112. project_key="dev-01",
  113. name="modbus_tcp_1",
  114. device_type=1,
  115. ip="127.0.0.1",
  116. port=5502,
  117. slave_id=1,
  118. byte_order=1,
  119. word_order=1,
  120. address_base=0,
  121. group_id=10,
  122. )
  123. self.assertEqual(result, {"state": 0})
  124. api_create.assert_called_once_with(
  125. "dev-01",
  126. {
  127. "name": "modbus_tcp_1",
  128. "device_type": 1,
  129. "ip": "127.0.0.1",
  130. "port": 5502,
  131. "slave_id": 1,
  132. "byte_order": 1,
  133. "word_order": 1,
  134. "address_base": 0,
  135. "serial_port": "",
  136. "timeout": 3,
  137. "is_persistent": True,
  138. "baud_rate": 0,
  139. "data_bit": 0,
  140. "parity": 0,
  141. "stop_bit": 0,
  142. "mode": 0,
  143. "retry_times": 0,
  144. "group_id": 10,
  145. "alarm_interval": 90,
  146. "collect_interval": 5,
  147. },
  148. )
  149. def test_modbus_device_edit_uses_explicit_parameters(self) -> None:
  150. signature = inspect.signature(modbus_server.collector_modbus_device_edit)
  151. self.assertNotIn("payload", signature.parameters)
  152. with patch("data_collector_mcp.modbus_server.api_edit_modbus_device", return_value={"state": 0}) as api_edit:
  153. result = modbus_server.collector_modbus_device_edit(
  154. project_key="dev-01",
  155. ori_id=1,
  156. name="modbus_tcp_edited",
  157. device_type=1,
  158. ip="127.0.0.1",
  159. port=5502,
  160. slave_id=1,
  161. byte_order=2,
  162. word_order=2,
  163. address_offset=1,
  164. device_group_id=10,
  165. )
  166. self.assertEqual(result, {"state": 0})
  167. api_edit.assert_called_once_with(
  168. "dev-01",
  169. {
  170. "ori_id": 1,
  171. "name": "modbus_tcp_edited",
  172. "device_type": 1,
  173. "ip": "127.0.0.1",
  174. "port": 5502,
  175. "slave_id": 1,
  176. "byte_order": 2,
  177. "word_order": 2,
  178. "serial_port": "",
  179. "timeout": 3,
  180. "is_persistent": True,
  181. "baud_rate": 0,
  182. "data_bit": 0,
  183. "parity": 0,
  184. "stop_bit": 0,
  185. "mode": 0,
  186. "address_offset": 1,
  187. "retry_times": 0,
  188. "device_group_id": 10,
  189. "alarm_interval": 90,
  190. "collect_interval": 5,
  191. },
  192. )
  193. def test_modbus_point_edit_uses_explicit_parameters_with_func_code(self) -> None:
  194. signature = inspect.signature(modbus_server.collector_modbus_point_edit)
  195. self.assertNotIn("payload", signature.parameters)
  196. with patch("data_collector_mcp.modbus_server.api_edit_modbus_point", return_value={"state": 0}) as api_edit:
  197. result = modbus_server.collector_modbus_point_edit(
  198. project_key="dev-01",
  199. ori_id=101,
  200. name="holding_register_uint16_edited",
  201. address=10,
  202. data_type="uint16",
  203. func_code=3,
  204. point_id="HR_UINT16_EDITED",
  205. )
  206. self.assertEqual(result, {"state": 0})
  207. api_edit.assert_called_once_with(
  208. "dev-01",
  209. {
  210. "ori_id": 101,
  211. "name": "holding_register_uint16_edited",
  212. "address": 10,
  213. "type": "uint16",
  214. "point_id": "HR_UINT16_EDITED",
  215. "scale_ratio": 1,
  216. "value_offset": 0,
  217. "group_id": 0,
  218. "invalid_values": "",
  219. "valid_range_start": None,
  220. "valid_range_end": None,
  221. "bit": 0,
  222. "describe": "",
  223. "func_code": 3,
  224. },
  225. )
  226. def test_modbus_point_edit_uses_register_type_when_func_code_is_zero(self) -> None:
  227. with patch("data_collector_mcp.modbus_server.api_edit_modbus_point", return_value={"state": 0}) as api_edit:
  228. modbus_server.collector_modbus_point_edit(
  229. project_key="dev-01",
  230. ori_id=101,
  231. name="holding_register_uint16_edited",
  232. address=10,
  233. data_type="uint16",
  234. register_type="holding_register",
  235. )
  236. payload = api_edit.call_args.args[1]
  237. self.assertEqual(payload["register_type"], "holding_register")
  238. self.assertNotIn("func_code", payload)
  239. def test_s7_device_create_uses_explicit_parameters(self) -> None:
  240. signature = inspect.signature(s7_server.collector_s7_device_create)
  241. self.assertNotIn("payload", signature.parameters)
  242. with patch("data_collector_mcp.s7_server.api_create_s7_device", return_value={"state": 0}) as api_create:
  243. result = s7_server.collector_s7_device_create(
  244. project_key="dev-01",
  245. name="s7_1200_1",
  246. ip="127.0.0.1",
  247. rock=0,
  248. slot=1,
  249. tsap_conn_type="OP",
  250. )
  251. self.assertEqual(result, {"state": 0})
  252. api_create.assert_called_once_with(
  253. "dev-01",
  254. {
  255. "name": "s7_1200_1",
  256. "ip": "127.0.0.1",
  257. "rock": 0,
  258. "slot": 1,
  259. "port": 102,
  260. "device_type": 1,
  261. "tsap_conn_type": "OP",
  262. "is_persistent": False,
  263. "device_group_id": 0,
  264. "timeout": 3,
  265. "alarm_interval": 90,
  266. "collect_interval": 5,
  267. },
  268. )
  269. def test_s7_device_edit_uses_explicit_parameters(self) -> None:
  270. signature = inspect.signature(s7_server.collector_s7_device_edit)
  271. self.assertNotIn("payload", signature.parameters)
  272. with patch("data_collector_mcp.s7_server.api_edit_s7_device", return_value={"state": 0}) as api_edit:
  273. result = s7_server.collector_s7_device_edit(
  274. project_key="dev-01",
  275. ori_id=3,
  276. name="s7_edited",
  277. ip="127.0.0.1",
  278. rock=0,
  279. slot=1,
  280. tsap_conn_type="OP",
  281. )
  282. self.assertEqual(result, {"state": 0})
  283. api_edit.assert_called_once_with(
  284. "dev-01",
  285. {
  286. "ori_id": 3,
  287. "name": "s7_edited",
  288. "ip": "127.0.0.1",
  289. "rock": 0,
  290. "slot": 1,
  291. "port": 102,
  292. "device_type": 1,
  293. "tsap_conn_type": "OP",
  294. "is_persistent": False,
  295. "device_group_id": 0,
  296. "timeout": 3,
  297. "alarm_interval": 90,
  298. "collect_interval": 5,
  299. },
  300. )
  301. def test_s7_point_edit_uses_register_area_when_register_type_is_zero(self) -> None:
  302. signature = inspect.signature(s7_server.collector_s7_point_edit)
  303. self.assertNotIn("payload", signature.parameters)
  304. with patch("data_collector_mcp.s7_server.api_edit_s7_point", return_value={"state": 0}) as api_edit:
  305. s7_server.collector_s7_point_edit(
  306. project_key="dev-01",
  307. ori_id=101,
  308. device_id=3,
  309. name="db_real",
  310. address="1.10",
  311. data_type="float32",
  312. register_area="DB",
  313. point_id="DB_REAL",
  314. )
  315. payload = api_edit.call_args.args[1]
  316. self.assertEqual(payload["register_area"], "DB")
  317. self.assertNotIn("register_type", payload)
  318. def test_s7_gateway_tools_forward_to_api(self) -> None:
  319. with patch("data_collector_mcp.s7_server.api_s7_raw_read", return_value={"code": 0}) as raw_read:
  320. self.assertEqual(
  321. s7_server.s7_raw_read(
  322. project_key="dev-01",
  323. ip="192.168.1.10",
  324. rock=0,
  325. slot=1,
  326. read={"area": "DB", "db": 1, "start": 0, "size": 4},
  327. ),
  328. {"code": 0},
  329. )
  330. raw_read.assert_called_once_with(
  331. "dev-01",
  332. ip="192.168.1.10",
  333. rock=0,
  334. slot=1,
  335. read={"area": "DB", "db": 1, "start": 0, "size": 4},
  336. device_type="S7-1200",
  337. port=102,
  338. tsap_conn_type=None,
  339. )
  340. with patch("data_collector_mcp.s7_server.api_s7_point_collect_test", return_value={"code": 0}) as point_test:
  341. self.assertEqual(
  342. s7_server.s7_point_collect_test(
  343. project_key="dev-01",
  344. ip="192.168.1.10",
  345. port=1102,
  346. rock=0,
  347. slot=1,
  348. tsap_conn_type="BASIC",
  349. points=[{"area": "M", "start": 0, "type": "bool"}],
  350. ),
  351. {"code": 0},
  352. )
  353. point_test.assert_called_once_with(
  354. "dev-01",
  355. ip="192.168.1.10",
  356. rock=0,
  357. slot=1,
  358. points=[{"area": "M", "start": 0, "type": "bool"}],
  359. device_type="S7-1200",
  360. port=1102,
  361. tsap_conn_type="BASIC",
  362. )
  363. with patch("data_collector_mcp.s7_server.api_s7_connect_scan", return_value={"code": 0}) as connect_scan:
  364. self.assertEqual(s7_server.s7_connect_scan("dev-01", ip="192.168.1.10"), {"code": 0})
  365. connect_scan.assert_called_once_with("dev-01", ip="192.168.1.10")
  366. if __name__ == "__main__":
  367. unittest.main()