common.py 4.6 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143
  1. import json
  2. import logging
  3. import socket
  4. import threading
  5. import time
  6. import urllib.error
  7. import urllib.request
  8. from contextlib import closing
  9. from unittest.mock import patch
  10. import uvicorn
  11. from snap7.server import Server
  12. from snap7.type import SrvArea
  13. from app.main import app
  14. REQUEST_TIMEOUT = 10
  15. logging.getLogger("snap7.server").setLevel(logging.CRITICAL)
  16. def find_free_port() -> int:
  17. with closing(socket.socket(socket.AF_INET, socket.SOCK_STREAM)) as sock:
  18. sock.bind(("127.0.0.1", 0))
  19. return int(sock.getsockname()[1])
  20. class ApiTestServer:
  21. def __init__(self) -> None:
  22. self.port = find_free_port()
  23. self.base_url = f"http://127.0.0.1:{self.port}"
  24. config = uvicorn.Config(app, host="127.0.0.1", port=self.port, log_level="critical", access_log=False)
  25. self.server = uvicorn.Server(config)
  26. self.thread = threading.Thread(target=self.server.run, daemon=True)
  27. def __enter__(self) -> "ApiTestServer":
  28. self.thread.start()
  29. deadline = time.time() + 10
  30. while time.time() < deadline:
  31. try:
  32. with urllib.request.urlopen(f"{self.base_url}/api/dc-debugtool/health", timeout=1) as response:
  33. if response.status == 200:
  34. return self
  35. except OSError:
  36. time.sleep(0.05)
  37. raise RuntimeError("API test server did not start")
  38. def __exit__(self, exc_type, exc_value, traceback) -> None:
  39. self.server.should_exit = True
  40. self.thread.join(timeout=5)
  41. def post_json(self, path: str, payload: dict) -> tuple[int, dict]:
  42. data = json.dumps(payload).encode("utf-8")
  43. request = urllib.request.Request(
  44. f"{self.base_url}{path}",
  45. data=data,
  46. headers={"Content-Type": "application/json"},
  47. method="POST",
  48. )
  49. try:
  50. with urllib.request.urlopen(request, timeout=REQUEST_TIMEOUT) as response:
  51. body = response.read().decode("utf-8")
  52. return response.status, json.loads(body)
  53. except urllib.error.HTTPError as exc:
  54. body = exc.read().decode("utf-8")
  55. return exc.code, json.loads(body)
  56. class Snap7TestServer:
  57. def __init__(
  58. self,
  59. db_data: bytearray | None = None,
  60. input_data: bytearray | None = None,
  61. output_data: bytearray | None = None,
  62. memory_data: bytearray | None = None,
  63. ) -> None:
  64. self.port = find_free_port()
  65. self.db_data = db_data
  66. self.input_data = input_data
  67. self.output_data = output_data
  68. self.memory_data = memory_data
  69. self.server = Server()
  70. def __enter__(self) -> "Snap7TestServer":
  71. if self.db_data is not None:
  72. self.server.register_area(SrvArea.DB, 1, self.db_data)
  73. if self.input_data is not None:
  74. self.server.register_area(SrvArea.PE, 0, self.input_data)
  75. if self.output_data is not None:
  76. self.server.register_area(SrvArea.PA, 0, self.output_data)
  77. if self.memory_data is not None:
  78. self.server.register_area(SrvArea.MK, 0, self.memory_data)
  79. self.server.start(tcp_port=self.port)
  80. time.sleep(0.2)
  81. return self
  82. def __exit__(self, exc_type, exc_value, traceback) -> None:
  83. self.server.stop()
  84. destroy = getattr(self.server, "destroy", None)
  85. if callable(destroy):
  86. destroy()
  87. class FakeSnap7Client:
  88. def get_connected(self) -> bool:
  89. return False
  90. def disconnect(self) -> None:
  91. pass
  92. def destroy(self) -> None:
  93. pass
  94. def patch_connect_scan(successful: set[tuple[int, int, str]]):
  95. def fake_connect_client(ip: str, port: int, rock: int, slot: int, tsap_conn_type: str) -> FakeSnap7Client:
  96. if port != 102:
  97. raise AssertionError("connect_scan must use fixed S7 port 102")
  98. if (rock, slot, tsap_conn_type) not in successful:
  99. raise RuntimeError("connection failed")
  100. return FakeSnap7Client()
  101. return patch("app.services.s7_service.connect_client", side_effect=fake_connect_client)
  102. def s7_device_payload(port: int, device_type: str = "S7-1200") -> dict:
  103. return {
  104. "device_type": device_type,
  105. "ip": "127.0.0.1",
  106. "port": port,
  107. "rock": 0,
  108. "slot": 1,
  109. "tsap_conn_type": "PG",
  110. }
  111. def assert_response_contract(testcase, status: int, data: dict) -> None:
  112. testcase.assertEqual(200, status, data)
  113. testcase.assertIn("code", data)
  114. testcase.assertIn("msg", data)
  115. testcase.assertIn("data", data)
  116. testcase.assertIsInstance(data["data"], dict)