Files
ReinLoopTest/ReinLoop/tests/test_identification_config.py
2026-08-03 11:16:49 +08:00

134 lines
4.6 KiB
Python

import importlib.util
import csv
import io
from pathlib import Path
import sys
import types
import unittest
from unittest.mock import patch
MODULE_PATH = Path(__file__).resolve().parents[1] / "core" / "identification_config.py"
SPEC = importlib.util.spec_from_file_location(
"identification_config_under_test", MODULE_PATH
)
IDENTIFICATION_CONFIG = importlib.util.module_from_spec(SPEC)
SPEC.loader.exec_module(IDENTIFICATION_CONFIG)
validate_identification_config = IDENTIFICATION_CONFIG.validate_identification_config
parse_identification_config_csv = IDENTIFICATION_CONFIG.parse_identification_config_csv
download_identification_config = IDENTIFICATION_CONFIG.download_identification_config
VALID_CONFIG = {
"q_in_val": 50.0,
"dt": 0.1,
"n_order": 6,
"t_c": 2.5,
"levels": [10, 20, 30, 40, 50, 60, 70, 80],
"dead_area": 240.0,
"xa_full": 1000.0,
"V_val": 5.0,
"repeat": 2,
}
def config_csv(config):
output = io.StringIO(newline="")
writer = csv.writer(output)
writer.writerow(("parameter", "value"))
for key in (
"q_in_val", "dt", "n_order", "t_c", "levels", "dead_area",
"xa_full", "V_val", "repeat"):
value = config[key]
if key == "levels":
value = ",".join(str(item) for item in value)
writer.writerow((key, value))
return output.getvalue()
class IdentificationConfigTests(unittest.TestCase):
def test_accepts_and_normalizes_valid_config(self):
result = validate_identification_config(VALID_CONFIG)
self.assertEqual(result["repeat"], 2)
self.assertEqual(result["levels"], [
10.0, 20.0, 30.0, 40.0, 50.0, 60.0, 70.0, 80.0
])
def test_rejects_missing_field(self):
config = dict(VALID_CONFIG)
config.pop("repeat")
with self.assertRaisesRegex(ValueError, "缺少"):
validate_identification_config(config)
def test_parses_parameter_value_csv(self):
result = parse_identification_config_csv(config_csv(VALID_CONFIG))
self.assertEqual(result, VALID_CONFIG)
def test_rejects_non_power_of_two_levels(self):
config = dict(VALID_CONFIG, levels=[10, 20, 30])
with self.assertRaisesRegex(ValueError, "2 的整数次幂"):
validate_identification_config(config)
def test_rejects_symbol_period_shorter_than_sample_period(self):
config = dict(VALID_CONFIG, dt=0.1, t_c=0.05)
with self.assertRaisesRegex(ValueError, "t_c 必须大于等于 dt"):
validate_identification_config(config)
def test_rejects_travel_scan_above_xa_full(self):
config = dict(VALID_CONFIG, xa_full=999.0)
with self.assertRaisesRegex(ValueError, "1000"):
validate_identification_config(config)
def test_rejects_dead_area_at_or_above_xa_full(self):
config = dict(VALID_CONFIG, dead_area=1000.0)
with self.assertRaisesRegex(ValueError, "dead_area"):
validate_identification_config(config)
def test_download_requests_customer_config_and_validates_it(self):
calls = []
requests_module = types.ModuleType("requests")
requests_module.RequestException = Exception
def device_post(payload, timeout):
calls.append(("device_post", payload, timeout))
return {"success": True, "url": "https://temp/config"}
def get(url, timeout):
calls.append(("get", url, timeout))
return types.SimpleNamespace(
text=config_csv(VALID_CONFIG),
raise_for_status=lambda: None
)
requests_module.get = get
api_module = types.ModuleType("api")
api_module.device_post = device_post
with patch.dict(sys.modules, {
"requests": requests_module,
"api": api_module,
}):
result = download_identification_config(timeout=7)
self.assertEqual(result["repeat"], 2)
self.assertEqual(calls[0], ("device_post", {"type": "getIdentificationConfig"}, 7))
self.assertEqual(calls[1], ("get", "https://temp/config", 7))
def test_download_reports_cloud_rejection(self):
requests_module = types.ModuleType("requests")
requests_module.RequestException = Exception
api_module = types.ModuleType("api")
api_module.device_post = lambda *args, **kwargs: {"success": False, "errMsg": "配置不存在"}
with patch.dict(sys.modules, {
"requests": requests_module,
"api": api_module,
}):
with self.assertRaisesRegex(ValueError, "配置不存在"):
download_identification_config()
if __name__ == "__main__":
unittest.main()