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_udp_dump_capture.py

131 lines
4.8 KiB
Python

import socket
import tempfile
import unittest
from pathlib import Path
import cv2
from udp_dump_capture import LiveMikUdpCapture, UdpDumpCapture
def payload(flags, sequence, number, value, data=b""):
header = bytes((0, flags, sequence, number)) + value.to_bytes(4, "little")
return header + data
def packet(flags, sequence, number, value, data=b"", port=59004):
data = payload(flags, sequence, number, value, data)
return int(port).to_bytes(2, "little") + len(data).to_bytes(2, "little") + data
def frame_array(width=4, height=2, padding=2, labels=()):
rows = []
value = 100
for _ in range(height):
row = b"".join((value + index * 100).to_bytes(2, "little") for index in range(width))
rows.append(row + b"\xff" * padding)
value += width * 100
label_data = b"".join(labels)
video_header = (
width.to_bytes(2, "little")
+ height.to_bytes(2, "little")
+ bytes((UdpDumpCapture.PIXEL_INT16, 0, padding, 0))
)
return len(labels).to_bytes(4, "little") + label_data + video_header + b"".join(rows)
def dump_for(data, port=59004):
split = min(13, len(data))
return b"".join((
packet(2, 7, 0, len(data), port=port),
packet(0, 7, 1, 0, data[:split], port=port),
packet(0, 7, 2, split, data[split:], port=port),
packet(1, 7, 3, len(data), port=port),
))
class UdpDumpCaptureTests(unittest.TestCase):
def test_live_receiver_uses_same_mik_packet_assembly(self):
data = frame_array()
split = min(13, len(data))
packets = (
payload(2, 7, 0, len(data)),
payload(0, 7, 1, 0, data[:split]),
payload(0, 7, 2, split, data[split:]),
payload(1, 7, 3, len(data)),
)
cap = LiveMikUdpCapture("127.0.0.1", 0, fps=25, width=4, height=2)
sender = socket.socket(socket.AF_INET, socket.SOCK_DGRAM)
try:
for item in packets:
sender.sendto(item, ("127.0.0.1", cap.port))
ok, frame = cap.read()
self.assertTrue(ok)
self.assertEqual(frame.shape, (2, 4, 3))
self.assertEqual(cap.frames_read, 1)
finally:
sender.close()
cap.release()
def test_reads_spec_packet_log_and_strips_row_padding(self):
label = bytes(range(40))
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "camera-dump"
path.write_bytes(dump_for(frame_array(labels=(label,))))
cap = UdpDumpCapture(path, fps=25)
self.assertTrue(cap.isOpened())
self.assertEqual(cap.get(cv2.CAP_PROP_FRAME_WIDTH), 4)
self.assertEqual(cap.get(cv2.CAP_PROP_FRAME_HEIGHT), 2)
self.assertEqual(cap.pixel_id, UdpDumpCapture.PIXEL_INT16)
self.assertEqual(cap.row_padding, 2)
self.assertEqual(cap.last_labels, [label])
ok, frame = cap.read()
self.assertTrue(ok)
self.assertEqual(frame.shape, (2, 4, 3))
self.assertLess(int(frame[0, 0, 0]), int(frame[-1, -1, 0]))
self.assertEqual(cap.get(cv2.CAP_PROP_POS_MSEC), 40)
self.assertEqual(cap.read(), (False, None))
cap.release()
def test_accepts_consistent_mik_dump_from_another_port(self):
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "camera-40404.udp"
path.write_bytes(dump_for(frame_array(), port=40404))
cap = UdpDumpCapture(path)
self.assertTrue(cap.isOpened())
self.assertEqual(cap.port, 40404)
self.assertTrue(cap.read()[0])
cap.release()
def test_packet_gap_drops_array_and_recovers_at_next_start(self):
data = frame_array()
broken = packet(2, 1, 0, len(data)) + packet(0, 1, 2, 0, data)
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "recover.dump"
path.write_bytes(broken + dump_for(data))
cap = UdpDumpCapture(path)
self.assertTrue(cap.isOpened())
self.assertEqual(cap.dropped_arrays, 1)
cap.release()
def test_truncated_packet_fails_without_exception(self):
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "broken.dump"
path.write_bytes((59004).to_bytes(2, "little") + b"\x10\x00\x00")
cap = UdpDumpCapture(path)
self.assertFalse(cap.isOpened())
self.assertIn("truncated", cap.last_error)
def test_unknown_port_is_not_opened(self):
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "not-a-dump"
path.write_bytes(b"nope")
cap = UdpDumpCapture(path)
self.assertFalse(cap.isOpened())
if __name__ == "__main__":
unittest.main()