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-gateway/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) -> dict: return { "device_type": "S7TCP", "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)