| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143 |
- import json
- import logging
- import socket
- import threading
- import time
- import urllib.error
- import urllib.request
- from contextlib import closing
- from unittest.mock import patch
- import uvicorn
- from snap7.server import Server
- from snap7.type import SrvArea
- from app.main import app
- REQUEST_TIMEOUT = 10
- logging.getLogger("snap7.server").setLevel(logging.CRITICAL)
- def find_free_port() -> int:
- with closing(socket.socket(socket.AF_INET, socket.SOCK_STREAM)) as sock:
- sock.bind(("127.0.0.1", 0))
- return int(sock.getsockname()[1])
- class ApiTestServer:
- def __init__(self) -> None:
- self.port = find_free_port()
- self.base_url = f"http://127.0.0.1:{self.port}"
- config = uvicorn.Config(app, host="127.0.0.1", port=self.port, log_level="critical", access_log=False)
- self.server = uvicorn.Server(config)
- self.thread = threading.Thread(target=self.server.run, daemon=True)
- def __enter__(self) -> "ApiTestServer":
- self.thread.start()
- deadline = time.time() + 10
- while time.time() < deadline:
- try:
- with urllib.request.urlopen(f"{self.base_url}/api/dc-debugtool/health", timeout=1) as response:
- if response.status == 200:
- return self
- except OSError:
- time.sleep(0.05)
- raise RuntimeError("API test server did not start")
- def __exit__(self, exc_type, exc_value, traceback) -> None:
- self.server.should_exit = True
- self.thread.join(timeout=5)
- def post_json(self, path: str, payload: dict) -> tuple[int, dict]:
- data = json.dumps(payload).encode("utf-8")
- request = urllib.request.Request(
- f"{self.base_url}{path}",
- data=data,
- headers={"Content-Type": "application/json"},
- method="POST",
- )
- try:
- with urllib.request.urlopen(request, timeout=REQUEST_TIMEOUT) as response:
- body = response.read().decode("utf-8")
- return response.status, json.loads(body)
- except urllib.error.HTTPError as exc:
- body = exc.read().decode("utf-8")
- return exc.code, json.loads(body)
- class Snap7TestServer:
- def __init__(
- self,
- db_data: bytearray | None = None,
- input_data: bytearray | None = None,
- output_data: bytearray | None = None,
- memory_data: bytearray | None = None,
- ) -> None:
- self.port = find_free_port()
- self.db_data = db_data
- self.input_data = input_data
- self.output_data = output_data
- self.memory_data = memory_data
- self.server = Server()
- def __enter__(self) -> "Snap7TestServer":
- if self.db_data is not None:
- self.server.register_area(SrvArea.DB, 1, self.db_data)
- if self.input_data is not None:
- self.server.register_area(SrvArea.PE, 0, self.input_data)
- if self.output_data is not None:
- self.server.register_area(SrvArea.PA, 0, self.output_data)
- if self.memory_data is not None:
- self.server.register_area(SrvArea.MK, 0, self.memory_data)
- self.server.start(tcp_port=self.port)
- time.sleep(0.2)
- return self
- def __exit__(self, exc_type, exc_value, traceback) -> None:
- self.server.stop()
- destroy = getattr(self.server, "destroy", None)
- if callable(destroy):
- destroy()
- class FakeSnap7Client:
- def get_connected(self) -> bool:
- return False
- def disconnect(self) -> None:
- pass
- def destroy(self) -> None:
- pass
- def patch_connect_scan(successful: set[tuple[int, int, str]]):
- def fake_connect_client(ip: str, port: int, rock: int, slot: int, tsap_conn_type: str) -> FakeSnap7Client:
- if port != 102:
- raise AssertionError("connect_scan must use fixed S7 port 102")
- if (rock, slot, tsap_conn_type) not in successful:
- raise RuntimeError("connection failed")
- return FakeSnap7Client()
- return patch("app.services.s7_service.connect_client", side_effect=fake_connect_client)
- def s7_device_payload(port: int, device_type: str = "S7-1200") -> dict:
- return {
- "device_type": device_type,
- "ip": "127.0.0.1",
- "port": port,
- "rock": 0,
- "slot": 1,
- "tsap_conn_type": "PG",
- }
- def assert_response_contract(testcase, status: int, data: dict) -> None:
- testcase.assertEqual(200, status, data)
- testcase.assertIn("code", data)
- testcase.assertIn("msg", data)
- testcase.assertIn("data", data)
- testcase.assertIsInstance(data["data"], dict)
|