test_server_tools.py 20 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517
  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, bacnet_server, 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. self.assertIs(app.bacnet_server, bacnet_server)
  13. def test_app_registers_expected_mcp_tools(self) -> None:
  14. tools = asyncio.run(app.mcp.list_tools())
  15. tool_names = {tool.name for tool in tools}
  16. self.assertEqual(
  17. tool_names,
  18. {
  19. "project.list",
  20. "collector.device_list",
  21. "collector.device_connect",
  22. "collector.device_disconnect",
  23. "collector.device_points",
  24. "modbus.point_collect_test",
  25. "bacnet.point_collect_test",
  26. "bacnet.point_search",
  27. "bacnet.bbmd_whois",
  28. "s7.raw_read",
  29. "s7.point_collect_test",
  30. "s7.connect_scan",
  31. "collector.modbus_device_create",
  32. "collector.modbus_device_edit",
  33. "collector.modbus_point_create",
  34. "collector.modbus_point_edit",
  35. "collector.s7_device_create",
  36. "collector.s7_device_edit",
  37. "collector.s7_point_create",
  38. "collector.s7_point_edit",
  39. "collector.bacnet_device_create",
  40. "collector.bacnet_device_edit",
  41. "collector.bacnet_point_create",
  42. "collector.bacnet_point_edit",
  43. },
  44. )
  45. def test_project_list_filters_enabled_projects_and_sorts(self) -> None:
  46. with patch(
  47. "data_collector_mcp.common_server.load_projects_config",
  48. return_value=[
  49. {
  50. "project_key": "z-prod",
  51. "project_name": "Prod",
  52. "base_url": "http://gateway.prod",
  53. "data_collector_base_url": "http://collector.prod",
  54. "enabled": True,
  55. },
  56. {
  57. "project_key": "disabled",
  58. "project_name": "Disabled",
  59. "base_url": "http://gateway.disabled",
  60. "data_collector_base_url": "http://collector.disabled",
  61. "enabled": False,
  62. },
  63. {
  64. "project_key": "a-dev",
  65. "project_name": "Dev",
  66. "base_url": "http://gateway.dev",
  67. "data_collector_base_url": "http://collector.dev",
  68. "enabled": True,
  69. },
  70. ],
  71. ):
  72. result = common_server.project_list()
  73. self.assertEqual(
  74. result,
  75. {
  76. "projects": [
  77. {
  78. "project_key": "a-dev",
  79. "project_name": "Dev",
  80. "base_url": "http://gateway.dev",
  81. "data_collector_base_url": "http://collector.dev",
  82. },
  83. {
  84. "project_key": "z-prod",
  85. "project_name": "Prod",
  86. "base_url": "http://gateway.prod",
  87. "data_collector_base_url": "http://collector.prod",
  88. },
  89. ],
  90. "total": 2,
  91. },
  92. )
  93. def test_common_device_tools_forward_to_api(self) -> None:
  94. with patch("data_collector_mcp.common_server.api_list_devices", return_value={"state": 0}) as list_devices:
  95. self.assertEqual(common_server.collector_device_list("dev-01", num_points=True), {"state": 0})
  96. list_devices.assert_called_once_with("dev-01", num_points=True)
  97. with patch("data_collector_mcp.common_server.api_connect_device", return_value={"state": 0}) as connect:
  98. self.assertEqual(
  99. common_server.collector_device_connect("dev-01", device_id=1, device_type="s7"),
  100. {"state": 0},
  101. )
  102. connect.assert_called_once_with("dev-01", device_id=1, device_type="s7")
  103. with patch("data_collector_mcp.common_server.api_disconnect_device", return_value={"state": 0}) as disconnect:
  104. self.assertEqual(
  105. common_server.collector_device_disconnect("dev-01", device_id=1, device_type="modbus"),
  106. {"state": 0},
  107. )
  108. disconnect.assert_called_once_with("dev-01", device_id=1, device_type="modbus")
  109. with patch("data_collector_mcp.common_server.api_list_device_points", return_value={"state": 0}) as points:
  110. self.assertEqual(
  111. common_server.collector_device_points("dev-01", device_id=1, device_type="s7", group_id=2),
  112. {"state": 0},
  113. )
  114. points.assert_called_once_with("dev-01", device_id=1, device_type="s7", group_id=2)
  115. def test_modbus_device_create_accepts_batch_devices(self) -> None:
  116. signature = inspect.signature(modbus_server.collector_modbus_device_create)
  117. self.assertIn("devices", signature.parameters)
  118. devices = [
  119. {
  120. "name": "modbus_tcp_1",
  121. "device_type": 1,
  122. "ip": "127.0.0.1",
  123. "port": 5502,
  124. "slave_id": 1,
  125. "byte_order": 1,
  126. "word_order": 1,
  127. "address_base": 0,
  128. "group_id": 10,
  129. }
  130. ]
  131. with patch("data_collector_mcp.modbus_server.api_create_modbus_devices", return_value={"state": 0}) as api_create:
  132. result = modbus_server.collector_modbus_device_create(
  133. project_key="dev-01",
  134. devices=devices,
  135. )
  136. self.assertEqual(result, {"state": 0})
  137. api_create.assert_called_once_with("dev-01", devices)
  138. def test_modbus_point_create_accepts_batch_points(self) -> None:
  139. points = [{"device_id": 1, "name": "temperature", "address": 10, "type": "uint16", "func_code": 3}]
  140. with patch("data_collector_mcp.modbus_server.api_create_modbus_points", return_value={"state": 0}) as api_create:
  141. result = modbus_server.collector_modbus_point_create(project_key="dev-01", points=points)
  142. self.assertEqual(result, {"state": 0})
  143. api_create.assert_called_once_with("dev-01", points)
  144. def test_modbus_device_edit_uses_explicit_parameters(self) -> None:
  145. signature = inspect.signature(modbus_server.collector_modbus_device_edit)
  146. self.assertNotIn("payload", signature.parameters)
  147. with patch("data_collector_mcp.modbus_server.api_edit_modbus_device", return_value={"state": 0}) as api_edit:
  148. result = modbus_server.collector_modbus_device_edit(
  149. project_key="dev-01",
  150. ori_id=1,
  151. name="modbus_tcp_edited",
  152. device_type=1,
  153. ip="127.0.0.1",
  154. port=5502,
  155. slave_id=1,
  156. byte_order=2,
  157. word_order=2,
  158. address_offset=1,
  159. device_group_id=10,
  160. )
  161. self.assertEqual(result, {"state": 0})
  162. api_edit.assert_called_once_with(
  163. "dev-01",
  164. {
  165. "ori_id": 1,
  166. "name": "modbus_tcp_edited",
  167. "device_type": 1,
  168. "ip": "127.0.0.1",
  169. "port": 5502,
  170. "slave_id": 1,
  171. "byte_order": 2,
  172. "word_order": 2,
  173. "serial_port": "",
  174. "timeout": 3,
  175. "is_persistent": False,
  176. "baud_rate": 0,
  177. "data_bit": 0,
  178. "parity": 0,
  179. "stop_bit": 0,
  180. "mode": 0,
  181. "address_offset": 1,
  182. "retry_times": 0,
  183. "device_group_id": 10,
  184. "alarm_interval": 90,
  185. "collect_interval": 5,
  186. },
  187. )
  188. def test_modbus_point_edit_uses_explicit_parameters_with_func_code(self) -> None:
  189. signature = inspect.signature(modbus_server.collector_modbus_point_edit)
  190. self.assertNotIn("payload", signature.parameters)
  191. with patch("data_collector_mcp.modbus_server.api_edit_modbus_point", return_value={"state": 0}) as api_edit:
  192. result = modbus_server.collector_modbus_point_edit(
  193. project_key="dev-01",
  194. ori_id=101,
  195. name="holding_register_uint16_edited",
  196. address=10,
  197. data_type="uint16",
  198. func_code=3,
  199. point_id="HR_UINT16_EDITED",
  200. )
  201. self.assertEqual(result, {"state": 0})
  202. api_edit.assert_called_once_with(
  203. "dev-01",
  204. {
  205. "ori_id": 101,
  206. "name": "holding_register_uint16_edited",
  207. "address": 10,
  208. "type": "uint16",
  209. "point_id": "HR_UINT16_EDITED",
  210. "scale_ratio": 1,
  211. "value_offset": 0,
  212. "group_id": 0,
  213. "invalid_values": "",
  214. "valid_range_start": None,
  215. "valid_range_end": None,
  216. "bit": 0,
  217. "describe": "",
  218. "func_code": 3,
  219. },
  220. )
  221. def test_modbus_point_edit_uses_register_type_when_func_code_is_zero(self) -> None:
  222. with patch("data_collector_mcp.modbus_server.api_edit_modbus_point", return_value={"state": 0}) as api_edit:
  223. modbus_server.collector_modbus_point_edit(
  224. project_key="dev-01",
  225. ori_id=101,
  226. name="holding_register_uint16_edited",
  227. address=10,
  228. data_type="uint16",
  229. register_type="holding_register",
  230. )
  231. payload = api_edit.call_args.args[1]
  232. self.assertEqual(payload["register_type"], "holding_register")
  233. self.assertNotIn("func_code", payload)
  234. def test_s7_device_create_accepts_batch_devices(self) -> None:
  235. signature = inspect.signature(s7_server.collector_s7_device_create)
  236. self.assertIn("devices", signature.parameters)
  237. devices = [{"name": "s7_1200_1", "ip": "127.0.0.1", "rock": 0, "slot": 1, "tsap_conn_type": "OP"}]
  238. with patch("data_collector_mcp.s7_server.api_create_s7_devices", return_value={"state": 0}) as api_create:
  239. result = s7_server.collector_s7_device_create(
  240. project_key="dev-01",
  241. devices=devices,
  242. )
  243. self.assertEqual(result, {"state": 0})
  244. api_create.assert_called_once_with(
  245. "dev-01",
  246. [{"name": "s7_1200_1", "ip": "127.0.0.1", "rock": 0, "slot": 1, "tsap_conn_type": "OP"}],
  247. )
  248. def test_s7_point_create_accepts_batch_points(self) -> None:
  249. points = [{"device_id": 3, "name": "db_real", "address": "1.10", "type": "REAL", "register_area": "DB"}]
  250. with patch("data_collector_mcp.s7_server.api_create_s7_points", return_value={"state": 0}) as api_create:
  251. result = s7_server.collector_s7_point_create(project_key="dev-01", points=points)
  252. self.assertEqual(result, {"state": 0})
  253. api_create.assert_called_once_with("dev-01", points)
  254. def test_s7_device_edit_uses_explicit_parameters(self) -> None:
  255. signature = inspect.signature(s7_server.collector_s7_device_edit)
  256. self.assertNotIn("payload", signature.parameters)
  257. with patch("data_collector_mcp.s7_server.api_edit_s7_device", return_value={"state": 0}) as api_edit:
  258. result = s7_server.collector_s7_device_edit(
  259. project_key="dev-01",
  260. ori_id=3,
  261. name="s7_edited",
  262. ip="127.0.0.1",
  263. rock=0,
  264. slot=1,
  265. tsap_conn_type="OP",
  266. )
  267. self.assertEqual(result, {"state": 0})
  268. api_edit.assert_called_once_with(
  269. "dev-01",
  270. {
  271. "ori_id": 3,
  272. "name": "s7_edited",
  273. "ip": "127.0.0.1",
  274. "rock": 0,
  275. "slot": 1,
  276. "port": 102,
  277. "device_type": 1,
  278. "tsap_conn_type": "OP",
  279. "is_persistent": False,
  280. "device_group_id": 0,
  281. "timeout": 3,
  282. "alarm_interval": 90,
  283. "collect_interval": 5,
  284. },
  285. )
  286. def test_s7_point_edit_uses_register_area_when_register_type_is_zero(self) -> None:
  287. signature = inspect.signature(s7_server.collector_s7_point_edit)
  288. self.assertNotIn("payload", signature.parameters)
  289. with patch("data_collector_mcp.s7_server.api_edit_s7_point", return_value={"state": 0}) as api_edit:
  290. s7_server.collector_s7_point_edit(
  291. project_key="dev-01",
  292. ori_id=101,
  293. device_id=3,
  294. name="db_real",
  295. address="1.10",
  296. data_type="float32",
  297. register_area="DB",
  298. point_id="DB_REAL",
  299. )
  300. payload = api_edit.call_args.args[1]
  301. self.assertEqual(payload["register_area"], "DB")
  302. self.assertNotIn("register_type", payload)
  303. def test_s7_gateway_tools_forward_to_api(self) -> None:
  304. with patch("data_collector_mcp.s7_server.api_s7_raw_read", return_value={"code": 0}) as raw_read:
  305. self.assertEqual(
  306. s7_server.s7_raw_read(
  307. project_key="dev-01",
  308. ip="192.168.1.10",
  309. rock=0,
  310. slot=1,
  311. read={"area": "DB", "db": 1, "start": 0, "size": 4},
  312. ),
  313. {"code": 0},
  314. )
  315. raw_read.assert_called_once_with(
  316. "dev-01",
  317. ip="192.168.1.10",
  318. rock=0,
  319. slot=1,
  320. read={"area": "DB", "db": 1, "start": 0, "size": 4},
  321. device_type="S7-1200",
  322. port=102,
  323. tsap_conn_type=None,
  324. )
  325. with patch("data_collector_mcp.s7_server.api_s7_point_collect_test", return_value={"code": 0}) as point_test:
  326. self.assertEqual(
  327. s7_server.s7_point_collect_test(
  328. project_key="dev-01",
  329. ip="192.168.1.10",
  330. port=1102,
  331. rock=0,
  332. slot=1,
  333. tsap_conn_type="BASIC",
  334. points=[{"area": "M", "start": 0, "type": "bool"}],
  335. ),
  336. {"code": 0},
  337. )
  338. point_test.assert_called_once_with(
  339. "dev-01",
  340. ip="192.168.1.10",
  341. rock=0,
  342. slot=1,
  343. points=[{"area": "M", "start": 0, "type": "bool"}],
  344. device_type="S7-1200",
  345. port=1102,
  346. tsap_conn_type="BASIC",
  347. )
  348. with patch("data_collector_mcp.s7_server.api_s7_connect_scan", return_value={"code": 0}) as connect_scan:
  349. self.assertEqual(s7_server.s7_connect_scan("dev-01", ip="192.168.1.10"), {"code": 0})
  350. connect_scan.assert_called_once_with("dev-01", ip="192.168.1.10")
  351. def test_bacnet_gateway_tools_forward_to_api(self) -> None:
  352. with patch(
  353. "data_collector_mcp.bacnet_server.api_bacnet_point_collect_test",
  354. return_value={"code": 0},
  355. ) as point_test:
  356. self.assertEqual(
  357. bacnet_server.bacnet_point_collect_test(
  358. project_key="dev-01",
  359. ip="192.168.1.20",
  360. bacnet_device_id=12345,
  361. port=47809,
  362. points=[{"object_type": "AnalogInput", "object_id": 1}],
  363. ),
  364. {"code": 0},
  365. )
  366. point_test.assert_called_once_with(
  367. "dev-01",
  368. ip="192.168.1.20",
  369. bacnet_device_id=12345,
  370. port=47809,
  371. points=[{"object_type": "AnalogInput", "object_id": 1}],
  372. )
  373. with patch(
  374. "data_collector_mcp.bacnet_server.api_bacnet_point_search",
  375. return_value={"code": 0},
  376. ) as point_search:
  377. self.assertEqual(
  378. bacnet_server.bacnet_point_search(
  379. project_key="dev-01",
  380. ip="192.168.1.20",
  381. bacnet_device_id=12345,
  382. ),
  383. {"code": 0},
  384. )
  385. point_search.assert_called_once_with(
  386. "dev-01",
  387. ip="192.168.1.20",
  388. bacnet_device_id=12345,
  389. port=47808,
  390. )
  391. with patch(
  392. "data_collector_mcp.bacnet_server.api_bacnet_bbmd_whois",
  393. return_value={"code": 0},
  394. ) as bbmd_whois:
  395. self.assertEqual(bacnet_server.bacnet_bbmd_whois("dev-01"), {"code": 0})
  396. bbmd_whois.assert_called_once_with("dev-01")
  397. def test_bacnet_collector_tools_forward_to_api(self) -> None:
  398. devices = [{"name": "bacnet_1", "ip": "192.168.1.20", "bacnet_device_id": 12345}]
  399. with patch(
  400. "data_collector_mcp.bacnet_server.api_create_bacnet_devices",
  401. return_value={"state": 0},
  402. ) as api_create:
  403. result = bacnet_server.collector_bacnet_device_create(project_key="dev-01", devices=devices)
  404. self.assertEqual(result, {"state": 0})
  405. api_create.assert_called_once_with("dev-01", devices)
  406. with patch(
  407. "data_collector_mcp.bacnet_server.api_edit_bacnet_device",
  408. return_value={"state": 0},
  409. ) as api_edit:
  410. result = bacnet_server.collector_bacnet_device_edit(
  411. project_key="dev-01",
  412. ori_id=9,
  413. name="bacnet_edited",
  414. ip="192.168.1.21",
  415. bacnet_device_id=54321,
  416. device_group_id=10,
  417. )
  418. self.assertEqual(result, {"state": 0})
  419. api_edit.assert_called_once_with(
  420. "dev-01",
  421. {
  422. "ori_id": 9,
  423. "name": "bacnet_edited",
  424. "ip": "192.168.1.21",
  425. "bacnet_device_id": 54321,
  426. "port": 47808,
  427. "bacnet_net": 0,
  428. "asp_ip": "",
  429. "is_persistent": False,
  430. "device_group_id": 10,
  431. "timeout": 3,
  432. "alarm_interval": 90,
  433. "collect_interval": 5,
  434. },
  435. )
  436. points = [{"device_id": 9, "name": "zone_temperature", "object_type": "AnalogInput", "object_id": 1}]
  437. with patch(
  438. "data_collector_mcp.bacnet_server.api_create_bacnet_points",
  439. return_value={"state": 0},
  440. ) as api_point_create:
  441. result = bacnet_server.collector_bacnet_point_create(project_key="dev-01", points=points)
  442. self.assertEqual(result, {"state": 0})
  443. api_point_create.assert_called_once_with("dev-01", points)
  444. with patch(
  445. "data_collector_mcp.bacnet_server.api_edit_bacnet_point",
  446. return_value={"state": 0},
  447. ) as api_point_edit:
  448. result = bacnet_server.collector_bacnet_point_edit(
  449. project_key="dev-01",
  450. ori_id=101,
  451. name="zone_temperature_edited",
  452. object_type="AnalogInput",
  453. object_id=1,
  454. point_id="AI_TEMP_EDITED",
  455. )
  456. self.assertEqual(result, {"state": 0})
  457. payload = api_point_edit.call_args.args[1]
  458. self.assertEqual(api_point_edit.call_args.args[0], "dev-01")
  459. self.assertEqual(payload["ori_id"], 101)
  460. self.assertEqual(payload["name"], "zone_temperature_edited")
  461. self.assertEqual(payload["object_type"], "AnalogInput")
  462. self.assertEqual(payload["object_id"], 1)
  463. self.assertEqual(payload["point_id"], "AI_TEMP_EDITED")
  464. if __name__ == "__main__":
  465. unittest.main()