import hashlib 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 IntegrityTracker, PACKET_FLAG_TIMESTAMP_DELTA_SATURATED, 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, ) fixture = subprocess.run([str(executable)], capture_output=True) if fixture.returncode != 0: stderr = fixture.stderr.decode(errors="replace").strip() raise AssertionError( f"protocol fixture exited {fixture.returncode}: {stderr}" ) cls.encoded = fixture.stdout transport_executable = Path(cls.tempdir.name) / "transport_fixture" subprocess.run( [ compiler, "-std=c11", "-Wall", "-Wextra", "-Werror", "-I", str(ROOT / "main"), str(ROOT / "main" / "trikke_transport.c"), str(ROOT / "tests" / "transport_fixture.c"), "-o", str(transport_executable), ], check=True, ) transport_fixture = subprocess.run( [str(transport_executable)], capture_output=True ) if transport_fixture.returncode != 0: stderr = transport_fixture.stderr.decode(errors="replace").strip() raise AssertionError( "transport fixture exited " f"{transport_fixture.returncode}: {stderr}" ) cls.transport_fixture_passed = True @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(4, len(frames)) metadata_frame, sample_frame, full_frame, saturated_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) self.assertEqual(0, parser.buffered_bytes) self.assertEqual(8, len(full_frame.samples)) self.assertEqual(43, full_frame.packet_sequence) self.assertEqual(2_007, full_frame.samples[-1].sequence) self.assertEqual(3_070_000, full_frame.samples[-1].timestamp_us) self.assertEqual(0, full_frame.flags) self.assertEqual(44, saturated_frame.packet_sequence) self.assertEqual( PACKET_FLAG_TIMESTAMP_DELTA_SATURATED, saturated_frame.flags, ) self.assertEqual(4_655_350, saturated_frame.samples[-1].timestamp_us) 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_transport_state_machine_contract(self) -> None: self.assertTrue(self.transport_fixture_passed) def test_integrity_sequence_wrap_classification(self) -> None: self.assertEqual( (0, 0), IntegrityTracker._classify_sequence(0xFFFFFFFF, 0) ) self.assertEqual( (2, 0), IntegrityTracker._classify_sequence(0xFFFFFFFE, 1) ) self.assertEqual((0, 1), IntegrityTracker._classify_sequence(1000, 0)) 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(3, 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 second_size = 36 + 2 * 20 damaged = bytearray(self.encoded[first_size : first_size + second_size]) damaged[-1] ^= 0x80 parser = StreamParser() frames = parser.feed( self.encoded[:first_size] + bytes(damaged) + self.encoded[first_size + second_size :] ) self.assertEqual(0, parser.startup_crc_errors) self.assertEqual(1, parser.crc_errors) self.assertEqual(3, len(frames)) self.assertEqual(PACKET_TYPE_METADATA, frames[0].packet_type) def test_trailing_partial_frame_is_observable(self) -> None: parser = StreamParser() frames = parser.feed(self.encoded[:-5]) self.assertEqual(3, len(frames)) self.assertEqual(36 + 2 * 20 - 5, parser.buffered_bytes) def test_hardware_outage_validation_artifacts(self) -> None: expected = { "forced_outage_3s.trk": { "sha256": "01482816cdaa668e4681c33c8baa1df331d733b9bbcbc4f448ece25e88185ad6", "sample_count": 2144, "first_sequence": 0, "last_sequence": 2143, "max_dropped": 0, "timing_anomalies": 2, "gaps": [], }, "forced_outage_7s.trk": { "sha256": "2ea8a5742944bdebc13bec2ccdbceba75f0bb71e48c856b0f86285878e190cd3", "sample_count": 1840, "first_sequence": 0, "last_sequence": 1977, "max_dropped": 138, "timing_anomalies": 3, "gaps": [(511, 650, 1_390_000)], }, "direct_usb_stall.trk": { "sha256": "40f874b7eaa7f705524ecdd75f832e8a724252366633116ac015fc75dfd16558", "sample_count": 864, "first_sequence": 8, "last_sequence": 2065, "max_dropped": 1194, "timing_anomalies": 1, "gaps": [(511, 1706, 11_950_000)], }, "direct_usb_3c95f3d.trk": { "sha256": "f495486f094a758bb785e145026e3934d60b52dd5083aef7f5d896195d967869", "sample_count": 1680, "first_sequence": 0, "last_sequence": 1679, "max_dropped": 0, "timing_anomalies": 8, "gaps": [], }, } for name, contract in expected.items(): with self.subTest(fixture=name): data = (ROOT / "tests" / "fixtures" / name).read_bytes() self.assertEqual( contract["sha256"], hashlib.sha256(data).hexdigest() ) parser = StreamParser() frames = [] for offset in range(0, len(data), 257): frames.extend(parser.feed(data[offset : offset + 257])) self.assertEqual(0, parser.startup_crc_errors) self.assertEqual(0, parser.crc_errors) self.assertEqual(0, parser.header_errors) self.assertEqual(0, parser.skipped_bytes) self.assertEqual(0, parser.buffered_bytes) self.assertTrue(frames) samples = [sample for frame in frames for sample in frame.samples] self.assertEqual(contract["sample_count"], len(samples)) self.assertEqual(contract["first_sequence"], samples[0].sequence) self.assertEqual(contract["last_sequence"], samples[-1].sequence) self.assertEqual( contract["max_dropped"], max(frame.dropped_sample_count for frame in frames), ) self.assertEqual( 0, max(frame.loop_overrun_count for frame in frames) ) self.assertFalse( any( frame.flags & PACKET_FLAG_TIMESTAMP_DELTA_SATURATED for frame in frames ) ) integrity = IntegrityTracker() for frame in frames: integrity.observe(frame) self.assertEqual(0, integrity.packet_gap_count) self.assertEqual(0, integrity.packet_reset_count) self.assertEqual( contract["max_dropped"], integrity.sample_gap_count ) self.assertEqual(0, integrity.sample_reset_count) self.assertEqual( contract["timing_anomalies"], integrity.timing_anomaly_count, ) self.assertEqual( contract["max_dropped"], integrity.final_dropped_sample_count, ) self.assertEqual(0, integrity.final_loop_overrun_count) gaps = [ ( left.sequence, right.sequence, right.timestamp_us - left.timestamp_us, ) for left, right in zip(samples, samples[1:]) if right.sequence != left.sequence + 1 ] self.assertEqual(contract["gaps"], gaps) def test_exact_usb_raw_wire_evidence(self) -> None: wire = ( ROOT / "tests" / "fixtures" / "direct_usb_3c95f3d.wire" ).read_bytes() validated = ( ROOT / "tests" / "fixtures" / "direct_usb_3c95f3d.trk" ).read_bytes() self.assertEqual( "3bdaeadff7962c6eac48c4ebeda285c8eb359e439d5e1728104add2009122c03", hashlib.sha256(wire).hexdigest(), ) parser = StreamParser() frames = [] for offset in range(0, len(wire), 113): frames.extend(parser.feed(wire[offset : offset + 113])) self.assertEqual(214, len(frames)) self.assertEqual(validated, b"".join(frame.raw for frame in frames)) self.assertEqual(0, parser.startup_crc_errors) self.assertEqual(0, parser.crc_errors) self.assertEqual(1, parser.header_errors) self.assertEqual(563, parser.skipped_bytes) self.assertEqual(0, parser.buffered_bytes) if __name__ == "__main__": unittest.main()