Files
trikkeSensors/tests/test_trikke_protocol.py
T

503 lines
20 KiB
Python

import hashlib
import json
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,
PACKET_TYPE_STATUS,
StreamParser,
sample_to_csv_row,
)
from trikke_ble import BleFrameReassembler, encode_ack, encode_begin_session # noqa: E402
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
ble_protocol_executable = Path(cls.tempdir.name) / "ble_protocol_fixture"
subprocess.run(
[
compiler,
"-std=c11",
"-Wall",
"-Wextra",
"-Werror",
"-I",
str(ROOT / "main"),
str(ROOT / "main" / "trikke_ble_protocol.c"),
str(ROOT / "tests" / "ble_protocol_fixture.c"),
"-o",
str(ble_protocol_executable),
],
check=True,
)
ble_protocol_fixture = subprocess.run(
[str(ble_protocol_executable)], capture_output=True
)
if ble_protocol_fixture.returncode != 0:
stderr = ble_protocol_fixture.stderr.decode(errors="replace").strip()
raise AssertionError(
"BLE protocol fixture exited "
f"{ble_protocol_fixture.returncode}: {stderr}"
)
cls.ble_protocol_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(5, len(frames))
metadata_frame, sample_frame, full_frame, saturated_frame, status_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)
self.assertEqual(PACKET_TYPE_STATUS, status_frame.packet_type)
self.assertEqual(45, status_frame.packet_sequence)
self.assertEqual(5, status_frame.status.sensor_read_failure_count)
self.assertEqual(6, status_frame.status.queue_overflow_count)
self.assertEqual(7, status_frame.status.transport_begin_retry_count)
self.assertEqual(8, status_frame.status.transport_disconnect_count)
self.assertEqual(9, status_frame.status.transport_send_failure_count)
self.assertEqual(10, status_frame.status.transport_replay_count)
self.assertEqual(11, status_frame.status.transport_invalid_ack_count)
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_ble_fragment_and_ack_contract(self) -> None:
self.assertTrue(self.ble_protocol_fixture_passed)
def test_ble_reassembly_and_replay_contract(self) -> None:
frame = self.encoded[: 36 + 48]
sequence = int.from_bytes(frame[12:16], "little")
def fragment(offset: int, size: int) -> bytes:
data = frame[offset : offset + size]
return (
sequence.to_bytes(4, "little")
+ offset.to_bytes(2, "little")
+ len(frame).to_bytes(2, "little")
+ data
)
reassembler = BleFrameReassembler()
self.assertIsNone(reassembler.feed(fragment(0, 20)))
# A replay from offset zero discards the partial attempt cleanly.
self.assertIsNone(reassembler.feed(fragment(0, 40)))
self.assertEqual(frame, reassembler.feed(fragment(40, len(frame) - 40)))
self.assertEqual(b"ACK1" + sequence.to_bytes(4, "little"), encode_ack(sequence))
self.assertEqual(
b"BGN1\x08\x07\x06\x05\x04\x03\x02\x01",
encode_begin_session(0x0102030405060708),
)
self.assertEqual(0, reassembler.rejected_fragment_count)
self.assertIsNone(reassembler.feed(fragment(20, 20)))
self.assertEqual(1, reassembler.rejected_fragment_count)
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(4, 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(4, 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(4, len(frames))
self.assertEqual(36 + 32 - 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": [],
},
"direct_usb_73e5680.trk": {
"sha256": "fd34bb3bf8f92a64960024f1287e553c03076ff3714fe629ec629bec81ddf821",
"sample_count": 1184,
"first_sequence": 0,
"last_sequence": 1183,
"max_dropped": 0,
"timing_anomalies": 2,
"gaps": [],
},
"ble_reconnect_483ace3.trk": {
"sha256": "c4a0d795cdf9d490acaca0144c3ad33f85bbfb2214d3f7abdfe515c4e5f25398",
"sample_count": 2576,
"first_sequence": 1992,
"last_sequence": 4567,
"max_dropped": 0,
"timing_anomalies": 1293,
"gaps": [],
"final_status": {
"sensor_read_failure_count": 0,
"queue_overflow_count": 0,
"transport_begin_retry_count": 187,
"transport_disconnect_count": 1,
"transport_send_failure_count": 0,
"transport_replay_count": 1,
"transport_invalid_ack_count": 0,
},
},
"ble_mtu_race_4bf00eb.trk": {
"sha256": "4a231b54a3c7320fd69cac869b830e94aca8f704f946881b9cdfd21991bfa41f",
"sample_count": 1848,
"first_sequence": 0,
"last_sequence": 4845,
"max_dropped": 2998,
"timing_anomalies": 412,
"gaps": [(1023, 4022, 29_990_016)],
"final_status": {
"sensor_read_failure_count": 0,
"queue_overflow_count": 2998,
"transport_begin_retry_count": 2471,
"transport_disconnect_count": 6,
"transport_send_failure_count": 0,
"transport_replay_count": 6,
"transport_invalid_ack_count": 0,
},
},
"ride_20260820_090607.trk": {
"sha256": "c938934c0905748d2f8d8be61cff1ff446d8415b34850628654967398fd7d91f",
"sample_count": 19984,
"first_sequence": 0,
"last_sequence": 19983,
"max_dropped": 0,
"timing_anomalies": 8745,
"accel_overruns": 6,
"gaps": [],
"final_status": {
"sensor_read_failure_count": 0,
"queue_overflow_count": 0,
"transport_begin_retry_count": 2,
"transport_disconnect_count": 1,
"transport_send_failure_count": 0,
"transport_replay_count": 1,
"transport_invalid_ack_count": 0,
},
},
}
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,
)
if "accel_overruns" in contract:
self.assertEqual(
contract["accel_overruns"],
integrity.accel_overrun_count,
)
self.assertEqual(
contract["max_dropped"],
integrity.final_dropped_sample_count,
)
self.assertEqual(0, integrity.final_loop_overrun_count)
if "final_status" in contract:
self.assertIsNotNone(integrity.final_status)
for field, value in contract["final_status"].items():
self.assertEqual(
value,
getattr(integrity.final_status, field),
)
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_android_reconnect_session_sidecar(self) -> None:
data = (
ROOT
/ "tests"
/ "fixtures"
/ "ride_20260820_090607.session.json"
).read_bytes()
self.assertEqual(
"b16a8f0c213a14fd2760b6efba1c6288f3631e57a331c553893e281dce047888",
hashlib.sha256(data).hexdigest(),
)
summary = json.loads(data)
self.assertTrue(summary["complete"])
self.assertIsNone(summary["error"])
self.assertEqual("ride_20260820_090607.trk", summary["captureFile"])
self.assertEqual(2578, summary["frames"])
self.assertEqual(19984, summary["samples"])
self.assertEqual(1, summary["duplicateReplays"])
self.assertEqual(0, summary["packetGaps"])
self.assertEqual(0, summary["sampleGaps"])
self.assertEqual(0, summary["droppedSamples"])
self.assertEqual(0, summary["queueOverflows"])
self.assertEqual(1, summary["transportDisconnects"])
self.assertEqual(1, summary["transportReplays"])
def test_exact_usb_raw_wire_evidence(self) -> None:
expected = {
"direct_usb_3c95f3d": {
"sha256": "3bdaeadff7962c6eac48c4ebeda285c8eb359e439d5e1728104add2009122c03",
"frames": 214,
"skipped": 563,
},
"direct_usb_73e5680": {
"sha256": "82d6d17bbf0729e9bfc53f337ec9adf70bcb5e5898b039685eda5dfa19cf4eea",
"frames": 151,
"skipped": 3589,
},
}
for stem, contract in expected.items():
with self.subTest(fixture=stem):
wire = (
ROOT / "tests" / "fixtures" / f"{stem}.wire"
).read_bytes()
validated = (
ROOT / "tests" / "fixtures" / f"{stem}.trk"
).read_bytes()
self.assertEqual(
contract["sha256"], 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(contract["frames"], 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(contract["skipped"], parser.skipped_bytes)
self.assertEqual(0, parser.buffered_bytes)
if __name__ == "__main__":
unittest.main()