You cannot select more than 25 topics Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
MAI/tests/test_error_output.py

162 lines
5.4 KiB
Python

import math
import socket
import struct
import sys
import unittest
from pathlib import Path
from unittest.mock import Mock, patch
PROJECT_ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(PROJECT_ROOT))
from error_output import (
ErrorOutputSender,
build_error_payload,
decode_guidance_v1_response,
encode_error_payload,
)
class ErrorOutputTests(unittest.TestCase):
def test_pixel_units_from_normalized_guidance_error(self):
payload = build_error_payload(
{"frame_id": 7, "frame_w": 1280, "frame_h": 720, "error_x": 0.25, "error_y": -0.5},
units="px",
timestamp=1.0,
)
self.assertEqual(payload["unit"], "px")
self.assertEqual(payload["x"], 160.0)
self.assertEqual(payload["y"], -180.0)
self.assertAlmostEqual(payload["mag"], math.hypot(160.0, -180.0))
def test_degree_units_use_fov(self):
payload = build_error_payload(
{"frame_w": 100, "frame_h": 100, "error_x": 1.0, "error_y": 0.0},
units="deg",
hfov_deg=90,
vfov_deg=60,
timestamp=1.0,
)
self.assertEqual(payload["unit"], "deg")
self.assertAlmostEqual(payload["x"], 45.0, places=5)
self.assertAlmostEqual(payload["y"], 0.0, places=5)
def test_meter_units_need_range(self):
payload = build_error_payload(
{"frame_w": 100, "frame_h": 100, "error_x": 1.0, "error_y": 0.0},
units="m",
hfov_deg=90,
range_m=10,
timestamp=1.0,
)
self.assertTrue(payload["valid"])
self.assertAlmostEqual(payload["x"], 10.0, places=5)
def test_binary_packet_magic(self):
payload = build_error_payload({"frame_id": 3, "frame_w": 100, "frame_h": 100}, timestamp=1.0)
data = encode_error_payload(payload, "bin")
self.assertEqual(data[:4], b"FPVE")
self.assertEqual(struct.unpack("<I", data[4:8])[0], 3)
def test_csv_packet_has_selected_unit(self):
payload = build_error_payload({"frame_id": 3, "frame_w": 100, "frame_h": 100}, units="norm", timestamp=1.0)
text = encode_error_payload(payload, "csv").decode("ascii")
self.assertIn(",norm,", text)
def test_guidance_v1_packet_matches_document_layout(self):
payload = build_error_payload(
{
"active": True,
"det_count": 2,
"frame_w": 100,
"frame_h": 100,
"error_x": 0.5,
"error_y": -0.4,
"box_w": 20,
"box_h": 10,
},
object_id=7,
timestamp=1.0,
)
data = encode_error_payload(payload, "guidance_v1")
self.assertEqual(len(data), 10)
self.assertEqual(
struct.unpack("<BBBhhbbb", data),
(1, 7, 2, 20, 25, 40, 50, 2),
)
def test_guidance_v1_no_target_zeros_measurements(self):
payload = build_error_payload(
{
"active": False,
"frame_w": 100,
"frame_h": 100,
"error_x": 1.0,
"error_y": 1.0,
},
object_id=3,
timestamp=1.0,
)
self.assertEqual(
struct.unpack("<BBBhhbbb", encode_error_payload(payload, "guidance_v1")),
(1, 3, 0, 0, 0, 0, 0, 0),
)
def test_guidance_v1_response_validation(self):
self.assertEqual(
decode_guidance_v1_response(bytes([2, 1])),
{"descriptor": 2, "response": 1},
)
self.assertIsNone(decode_guidance_v1_response(bytes([1, 1])))
self.assertIsNone(decode_guidance_v1_response(bytes([2, 3])))
def test_sender_uses_selected_udp_port(self):
receiver = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
receiver.bind(("127.0.0.1", 0))
receiver.settimeout(1.0)
port = receiver.getsockname()[1]
try:
with patch.multiple(
"error_output",
ERROR_OUTPUT_ENABLE=True,
ERROR_OUTPUT_PROTOCOL="guidance_v1",
ERROR_OUTPUT_HOST="127.0.0.1",
ERROR_OUTPUT_PORT=port,
ERROR_OUTPUT_OBJECT_ID=9,
ERROR_OUTPUT_EVERY=1,
):
sender = ErrorOutputSender()
sender.start()
try:
sender.send(
{
"frame_id": 1,
"active": True,
"det_count": 1,
"frame_w": 100,
"frame_h": 100,
"error_x": 0.0,
"error_y": 0.0,
"box_w": 10,
"box_h": 10,
}
)
data, _ = receiver.recvfrom(64)
finally:
sender.close()
self.assertEqual(len(data), 10)
self.assertEqual(data[:3], bytes([1, 9, 1]))
finally:
receiver.close()
def test_sender_skips_unverified_target(self):
sender = ErrorOutputSender()
sender.enabled = True
sender._sock = Mock()
self.assertIsNone(sender.send({"frame_id": 1, "active": False}))
sender._sock.sendto.assert_not_called()
if __name__ == "__main__":
unittest.main()