import shutil import subprocess import sys import tempfile import unittest from pathlib import Path ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT / "tools")) from trikke_protocol import ( # noqa: E402 PACKET_TYPE_METADATA, PACKET_TYPE_SAMPLES, StreamParser, sample_to_csv_row, ) class ProtocolContractTest(unittest.TestCase): @classmethod def setUpClass(cls) -> None: compiler = shutil.which("cc") if compiler is None: raise unittest.SkipTest("host C compiler is unavailable") cls.tempdir = tempfile.TemporaryDirectory() executable = Path(cls.tempdir.name) / "protocol_fixture" subprocess.run( [ compiler, "-std=c11", "-Wall", "-Wextra", "-Werror", "-I", str(ROOT / "main"), str(ROOT / "main" / "trikke_protocol.c"), str(ROOT / "tests" / "protocol_fixture.c"), "-o", str(executable), ], check=True, ) cls.encoded = subprocess.run( [str(executable)], check=True, capture_output=True ).stdout @classmethod def tearDownClass(cls) -> None: cls.tempdir.cleanup() def test_c_encoder_to_python_parser_contract(self) -> None: parser = StreamParser() frames = [] stream = b"startup text\r\n" + self.encoded for offset in range(0, len(stream), 7): frames.extend(parser.feed(stream[offset : offset + 7])) self.assertEqual(2, len(frames)) metadata_frame, sample_frame = frames self.assertEqual(PACKET_TYPE_METADATA, metadata_frame.packet_type) self.assertEqual(41, metadata_frame.packet_sequence) self.assertEqual(2, metadata_frame.dropped_sample_count) self.assertEqual(3, metadata_frame.loop_overrun_count) self.assertEqual(100, metadata_frame.metadata.sample_rate_hz) self.assertEqual((-1.5, -4.5, 12.0), metadata_frame.metadata.accel_offset_counts) self.assertEqual(PACKET_TYPE_SAMPLES, sample_frame.packet_type) self.assertEqual(42, sample_frame.packet_sequence) self.assertEqual(2, len(sample_frame.samples)) self.assertEqual(1000, sample_frame.samples[0].sequence) self.assertEqual(2_000_000, sample_frame.samples[0].timestamp_us) self.assertEqual((1, -2, 258), sample_frame.samples[0].accel) self.assertEqual(2_010_000, sample_frame.samples[1].timestamp_us) self.assertEqual((40, -50, 60), sample_frame.samples[1].gyro) self.assertEqual(0xFF, sample_frame.samples[1].gyro_status) self.assertEqual(0, parser.crc_errors) self.assertEqual(len(b"startup text\r\n"), parser.skipped_bytes) row = sample_to_csv_row( sample_frame.samples[0], metadata_frame.metadata, 3 ) self.assertEqual(23, len(row)) self.assertEqual((2, 1, 258), tuple(row[14:17])) self.assertEqual(3, row[-1]) def test_crc_failure_resynchronizes_to_next_frame(self) -> None: first_size = 36 + 48 damaged = bytearray(self.encoded[:first_size]) damaged[-1] ^= 0x80 parser = StreamParser() frames = parser.feed(bytes(damaged) + self.encoded[first_size:]) self.assertEqual(1, parser.startup_crc_errors) self.assertEqual(0, parser.crc_errors) self.assertEqual(1, len(frames)) self.assertEqual(PACKET_TYPE_SAMPLES, frames[0].packet_type) def test_crc_failure_after_sync_is_stream_error(self) -> None: first_size = 36 + 48 damaged = bytearray(self.encoded[first_size:]) damaged[-1] ^= 0x80 parser = StreamParser() frames = parser.feed(self.encoded[:first_size] + bytes(damaged)) self.assertEqual(0, parser.startup_crc_errors) self.assertEqual(1, parser.crc_errors) self.assertEqual(1, len(frames)) self.assertEqual(PACKET_TYPE_METADATA, frames[0].packet_type) if __name__ == "__main__": unittest.main()