"""Tests for UART v3 protocol serialization in switch_pico_uart.""" import struct import pytest from switch_pico_bridge.switch_pico_uart import ( SwitchReport, IMUSample, SwitchDpad, PicoUART, UART_HEADER, UART_PROTOCOL_VERSION, UART_SLOT_COUNT, UartLinkStats, RUMBLE_HEADER, RUMBLE_TYPE_DECODED, RUMBLE_TYPE_SLOT, ACCEL_LSB_PER_G, GYRO_LSB_PER_RAD_S, MS2_PER_G, compute_checksum, ) class BufferedSerial: def __init__(self, data: bytes = b""): self._data = bytearray(data) @property def in_waiting(self) -> int: return len(self._data) def read(self, size: int) -> bytes: data = bytes(self._data[:size]) del self._data[:size] return data def feed(self, data: bytes) -> None: self._data.extend(data) def make_rumble_frame(low: int, high: int, slot: int = 0) -> bytes: frame = bytes([RUMBLE_HEADER, RUMBLE_TYPE_SLOT, slot, low, high]) return frame + bytes([compute_checksum(frame)]) def make_legacy_rumble_frame(low: int, high: int) -> bytes: frame = bytes([RUMBLE_HEADER, RUMBLE_TYPE_DECODED, low, high]) return frame + bytes([compute_checksum(frame)]) def make_uart(data: bytes = b"") -> tuple[PicoUART, BufferedSerial]: uart = object.__new__(PicoUART) serial_port = BufferedSerial(data) uart.serial = serial_port uart._buffer = bytearray() uart.last_stats = None return uart, serial_port def test_v3_frame_with_imu_samples(): """V3 frame with 3 IMU samples should be 49 bytes with correct layout.""" r = SwitchReport( buttons=0, imu_samples=[ IMUSample(100, -200, 4096, 50, -50, 0), IMUSample(101, -201, 4097, 51, -51, 1), IMUSample(102, -202, 4098, 52, -52, 2), ], ) data = r.to_bytes() assert len(data) == 49, f"Expected 49 bytes, got {len(data)}" assert data[0] == UART_HEADER # 0xAA assert data[1] == UART_PROTOCOL_VERSION # 0x03 assert data[2] == 44 # payload_len assert data[3] == 0 # slot assert data[11] == 3 # imu_count # Verify checksum assert data[-1] == compute_checksum(data[:-1]) # Verify first sample accel_x (int16 LE at byte 12) ax0 = struct.unpack_from("3 IMU samples should cap at 3.""" samples = [IMUSample(i, 0, 0, 0, 0, 0) for i in range(5)] r = SwitchReport(imu_samples=samples) data = r.to_bytes() assert len(data) == 49 # 3 samples, not 5 assert data[11] == 3 assert data[2] == 44 # payload_len for 3 samples def test_decoded_rumble_frame_survives_fragmented_input(): frame = make_rumble_frame(64, 192, slot=1) uart, serial_port = make_uart(frame[:3]) assert uart.read_rumble() is None serial_port.feed(frame[3:]) assert uart.read_rumble() == pytest.approx((1, 64 / 255.0, 192 / 255.0)) def test_decoded_rumble_frame_resynchronizes_after_garbage(): uart, _ = make_uart(b"\x00\xffnot-a-frame" + make_rumble_frame(12, 34)) assert uart.read_rumble() == pytest.approx((0, 12 / 255.0, 34 / 255.0)) def test_decoded_rumble_frame_rejects_bad_checksum(): corrupted = bytearray(make_rumble_frame(25, 50)) corrupted[-1] ^= 0x01 uart, _ = make_uart(bytes(corrupted) + make_rumble_frame(75, 100, slot=2)) assert uart.read_rumble() == pytest.approx((2, 75 / 255.0, 100 / 255.0)) def test_decoded_rumble_zero_and_full_magnitudes(): uart, _ = make_uart(make_rumble_frame(0, 0) + make_rumble_frame(255, 255, slot=3)) assert uart.read_rumble() == (0, 0.0, 0.0) assert uart.read_rumble() == (3, 1.0, 1.0) def test_legacy_rumble_frame_maps_to_slot_zero(): """Pre-multi-controller firmware sends 5-byte frames without a slot byte.""" uart, _ = make_uart(make_legacy_rumble_frame(10, 20) + make_rumble_frame(30, 40, slot=1)) assert uart.read_rumble() == pytest.approx((0, 10 / 255.0, 20 / 255.0)) assert uart.read_rumble() == pytest.approx((1, 30 / 255.0, 40 / 255.0)) def test_rumble_frame_with_out_of_range_slot_is_skipped(): uart, _ = make_uart(make_rumble_frame(1, 2, slot=UART_SLOT_COUNT) + make_rumble_frame(3, 4)) assert uart.read_rumble() == pytest.approx((0, 3 / 255.0, 4 / 255.0)) def test_reboot_bootsel_frame_matches_firmware_contract(): """0xAA 0xFE len(8) cmd(1) 'BOOTSEL' checksum: 12 bytes, the parser's minimum frame.""" frame = PicoUART.reboot_bootsel_frame() assert frame[:3] == bytes([UART_HEADER, 0xFE, 8]) assert frame[3] == 0x01 assert frame[4:11] == b"BOOTSEL" assert len(frame) == 12 assert frame[-1] == compute_checksum(frame[:-1]) def test_stats_request_frame_matches_firmware_contract(): frame = PicoUART.stats_request_frame() assert frame[:3] == bytes([UART_HEADER, 0xFE, 8]) assert frame[3] == 0x02 assert frame[4:11] == b"STATS\0\0" assert frame[-1] == compute_checksum(frame[:-1]) def test_stats_reply_is_captured_without_disturbing_rumble_parsing(): counters = (1000, 2, 30, 0, 300, 900) stats = bytes([RUMBLE_HEADER, 0x05]) + struct.pack("<6I", *counters) stats += bytes([compute_checksum(stats)]) uart, _ = make_uart(make_rumble_frame(5, 6, slot=1) + stats + make_rumble_frame(7, 8)) assert uart.read_rumble() == pytest.approx((1, 5 / 255.0, 6 / 255.0)) assert uart.last_stats is None assert uart.read_rumble() == pytest.approx((0, 7 / 255.0, 8 / 255.0)) assert uart.last_stats == UartLinkStats(*counters) later = UartLinkStats(1010, 2, 30, 1, 303, 909) assert later.delta(uart.last_stats) == UartLinkStats(10, 0, 0, 1, 3, 9)