test_server_tools.py 23 KB

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