Emulate up to four controllers on the UART firmware

UART protocol v3 adds a slot byte to input frames (v2 still accepted as
slot 0) and slot-tagged BB 03 rumble frames. The regular firmware exposes
SWITCH_PICO_UART_CONTROLLERS (default 4) Switch Pro interfaces like the
AIO build. The bridge shares one serial port across controllers, maps
index:port[:slot], and demuxes rumble by slot.
This commit is contained in:
Joey Yakimowich-Payne 2026-09-22 10:25:30 -06:00
commit b6a017eb06
14 changed files with 737 additions and 411 deletions

View file

@ -681,8 +681,10 @@ void test_uart_parser_is_pure() {
packet.back() = static_cast<uint8_t>(packet.back() + packet[i]);
}
ControllerState parsed{};
expect(switch_pro_apply_uart_packet(packet.data(), packet.size(), parsed),
uint8_t slot = 0xff;
expect(switch_pro_apply_uart_packet(packet.data(), packet.size(), parsed, slot),
"valid UART packet was rejected");
expect(slot == 0, "v2 UART packet must map to slot 0");
expect(parsed.button_east && parsed.button_left_shoulder && parsed.dpad_down &&
parsed.dpad_left,
"UART buttons or hat were parsed incorrectly");
@ -707,12 +709,44 @@ void test_uart_parser_is_pure() {
ControllerState unchanged{};
unchanged.button_system = true;
unchanged.left_stick_x = 123;
uint8_t unchanged_slot = 0xff;
packet.back() ^= 0xffu;
expect(!switch_pro_apply_uart_packet(packet.data(), packet.size(),
unchanged),
unchanged, unchanged_slot),
"invalid UART checksum was accepted");
expect(unchanged.button_system && unchanged.left_stick_x == 123,
"failed UART parse modified its output reference");
expect(unchanged.button_system && unchanged.left_stick_x == 123 &&
unchanged_slot == 0xff,
"failed UART parse modified its output references");
// v3 inserts a slot byte between the length and the payload.
std::array<uint8_t, 13> slotted{};
slotted[0] = 0xaa;
slotted[1] = 0x03;
slotted[2] = 8;
slotted[3] = 2;
std::copy(packet.begin() + 3, packet.begin() + 10, slotted.begin() + 4);
for (unsigned i = 0; i < slotted.size() - 1; ++i) {
slotted.back() = static_cast<uint8_t>(slotted.back() + slotted[i]);
}
ControllerState slotted_state{};
expect(switch_pro_apply_uart_packet(slotted.data(), slotted.size(),
slotted_state, slot),
"valid v3 UART packet was rejected");
expect(slot == 2, "v3 slot byte was not reported");
expect(slotted_state.button_east && slotted_state.button_left_shoulder &&
slotted_state.dpad_down && slotted_state.dpad_left &&
slotted_state.right_stick_y ==
controller_axis_from_unsigned(0x7878),
"v3 payload offsets were parsed incorrectly");
slotted[3] = SWITCH_PICO_HID_INSTANCE_COUNT;
slotted.back() = 0;
for (unsigned i = 0; i < slotted.size() - 1; ++i) {
slotted.back() = static_cast<uint8_t>(slotted.back() + slotted[i]);
}
expect(!switch_pro_apply_uart_packet(slotted.data(), slotted.size(),
slotted_state, slot),
"out-of-range v3 slot was accepted");
}
void test_motion_backpressure_retries_without_advancing_state() {

View file

@ -28,10 +28,10 @@ class RecordingUART:
def __init__(self) -> None:
self.sent_imu: list[tuple[IMUSample, ...]] = []
def send_report(self, report: SwitchReport) -> None:
def send_report(self, report: SwitchReport, slot: int = 0) -> None:
self.sent_imu.append(tuple(report.imu_samples))
def read_rumble(self) -> tuple[float, float] | None:
def read_rumble(self) -> tuple[int, float, float] | None:
return None
@ -101,9 +101,8 @@ def test_sensor_buffer_retains_latest_three_samples() -> None:
def test_service_republishes_latest_imu_window(monkeypatch: MonkeyPatch) -> None:
uart = RecordingUART()
controller = cast(sdl3.SDL_Gamepad, object())
ctx = bridge.ControllerContext(
controller, 7, 0, "dualsense", "/dev/null", cast(PicoUART, cast(object, uart))
)
ctx = bridge.ControllerContext(controller, 7, 0, "dualsense", "/dev/null")
links = {"/dev/null": bridge.UartLink("/dev/null", cast(PicoUART, cast(object, uart)))}
ctx.sensors_enabled = True
samples = [
IMUSample(1, 2, 3, 4, 5, 6),
@ -123,8 +122,8 @@ def test_service_republishes_latest_imu_window(monkeypatch: MonkeyPatch) -> None
args = Namespace(baud=UART_BAUD)
console = Console(file=StringIO())
bridge.service_contexts(1.0, args, config, contexts, [], console)
bridge.service_contexts(2.0, args, config, contexts, [], console)
bridge.service_contexts(1.0, args, config, contexts, links, console)
bridge.service_contexts(2.0, args, config, contexts, links, console)
assert uart.sent_imu == [tuple(samples), tuple(samples)]
assert ctx.imu_samples == samples

View file

@ -1,4 +1,4 @@
"""Tests for UART v2 protocol serialization in switch_pico_uart."""
"""Tests for UART v3 protocol serialization in switch_pico_uart."""
import struct
import pytest
@ -9,8 +9,10 @@ from switch_pico_bridge.switch_pico_uart import (
PicoUART,
UART_HEADER,
UART_PROTOCOL_VERSION,
UART_SLOT_COUNT,
RUMBLE_HEADER,
RUMBLE_TYPE_DECODED,
RUMBLE_TYPE_SLOT,
ACCEL_LSB_PER_G,
GYRO_LSB_PER_RAD_S,
MS2_PER_G,
@ -35,7 +37,12 @@ class BufferedSerial:
self._data.extend(data)
def make_rumble_frame(low: int, high: int) -> bytes:
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)])
@ -48,8 +55,8 @@ def make_uart(data: bytes = b"") -> tuple[PicoUART, BufferedSerial]:
return uart, serial_port
def test_v2_frame_with_imu_samples():
"""V2 frame with 3 IMU samples should be 48 bytes with correct layout."""
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=[
@ -59,35 +66,49 @@ def test_v2_frame_with_imu_samples():
],
)
data = r.to_bytes()
assert len(data) == 48, f"Expected 48 bytes, got {len(data)}"
assert len(data) == 49, f"Expected 49 bytes, got {len(data)}"
assert data[0] == UART_HEADER # 0xAA
assert data[1] == UART_PROTOCOL_VERSION # 0x02
assert data[1] == UART_PROTOCOL_VERSION # 0x03
assert data[2] == 44 # payload_len
assert data[10] == 3 # imu_count
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 11)
ax0 = struct.unpack_from("<h", data, 11)[0]
# Verify first sample accel_x (int16 LE at byte 12)
ax0 = struct.unpack_from("<h", data, 12)[0]
assert ax0 == 100, f"Expected accel_x=100, got {ax0}"
# Verify first sample gyro_z (int16 LE at bytes 21-22)
gz0 = struct.unpack_from("<h", data, 21)[0]
# Verify first sample gyro_z (int16 LE at bytes 22-23)
gz0 = struct.unpack_from("<h", data, 22)[0]
assert gz0 == 0, f"Expected gyro_z=0, got {gz0}"
def test_v2_frame_no_imu():
"""V2 frame with no IMU samples should be 12 bytes."""
def test_v3_frame_no_imu():
"""V3 frame with no IMU samples should be 13 bytes."""
r = SwitchReport(
buttons=0x0004, hat=SwitchDpad.CENTER, lx=128, ly=128, rx=128, ry=128
)
data = r.to_bytes()
assert len(data) == 12, f"Expected 12 bytes, got {len(data)}"
assert len(data) == 13, f"Expected 13 bytes, got {len(data)}"
assert data[0] == UART_HEADER
assert data[1] == UART_PROTOCOL_VERSION
assert data[2] == 8 # payload_len
assert data[10] == 0 # imu_count
assert data[3] == 0 # slot
assert data[11] == 0 # imu_count
assert data[-1] == compute_checksum(data[:-1])
def test_v3_frame_addresses_slot():
"""The slot byte selects which emulated controller receives the report."""
data = SwitchReport(buttons=0x0001).to_bytes(slot=3)
assert data[3] == 3
assert struct.unpack_from("<H", data, 4)[0] == 0x0001
assert data[-1] == compute_checksum(data[:-1])
with pytest.raises(ValueError):
SwitchReport().to_bytes(slot=UART_SLOT_COUNT)
with pytest.raises(ValueError):
SwitchReport().to_bytes(slot=-1)
def test_checksum_validation():
"""Checksum should match sum of all preceding bytes & 0xFF."""
r = SwitchReport(buttons=0x0001)
@ -96,7 +117,7 @@ def test_checksum_validation():
assert data[-1] == expected_checksum
# Corrupt a byte and verify mismatch
corrupted = bytearray(data)
corrupted[3] ^= 0xFF # flip bits in first payload byte
corrupted[4] ^= 0xFF # flip bits in first payload byte
recalculated = sum(corrupted[:-1]) & 0xFF
assert corrupted[-1] != recalculated, "Checksum should not match corrupted data"
@ -126,21 +147,21 @@ def test_imu_sample_dataclass():
s2 = IMUSample(accel_x=99999)
r = SwitchReport(imu_samples=[s2])
data = r.to_bytes()
ax = struct.unpack_from("<h", data, 11)[0]
ax = struct.unpack_from("<h", data, 12)[0]
assert ax == 32767, f"Expected clamped value 32767, got {ax}"
def test_backward_compat_switch_report():
"""SwitchReport with no imu_samples produces valid v2 frame (backward compat)."""
def test_switch_report_payload_layout():
"""Buttons and axes land at the documented v3 payload offsets."""
r = SwitchReport(buttons=0x000A, lx=200, ly=50, rx=128, ry=128)
data = r.to_bytes()
assert len(data) == 12
assert data[1] == 0x02 # still v2
# Buttons at bytes 3-4
buttons = struct.unpack_from("<H", data, 3)[0]
assert len(data) == 13
assert data[1] == 0x03
# Buttons at bytes 4-5
buttons = struct.unpack_from("<H", data, 4)[0]
assert buttons == 0x000A
# lx at byte 6
assert data[6] == 200
# lx at byte 7
assert data[7] == 200
def test_max_imu_samples_capped():
@ -148,37 +169,51 @@ def test_max_imu_samples_capped():
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) == 48 # 3 samples, not 5
assert data[10] == 3
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)
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((64 / 255.0, 192 / 255.0))
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((12 / 255.0, 34 / 255.0))
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))
uart, _ = make_uart(bytes(corrupted) + make_rumble_frame(75, 100, slot=2))
assert uart.read_rumble() == pytest.approx((75 / 255.0, 100 / 255.0))
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))
uart, _ = make_uart(make_rumble_frame(0, 0) + make_rumble_frame(255, 255, slot=3))
assert uart.read_rumble() == (0.0, 0.0)
assert uart.read_rumble() == (1.0, 1.0)
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))

View file

@ -14,17 +14,22 @@ from switch_pico_bridge.switch_pico_uart import PicoUART, SwitchReport, UART_BAU
class RecordingUART:
def __init__(self) -> None:
self.rumble: list[tuple[float, float]] = []
self.rumble: list[tuple[int, float, float]] = []
self.sent: list[tuple[int, int]] = []
def send_report(self, _report: SwitchReport) -> None:
pass
def send_report(self, report: SwitchReport, slot: int = 0) -> None:
self.sent.append((slot, report.buttons))
def read_rumble(self) -> tuple[float, float] | None:
def read_rumble(self) -> tuple[int, float, float] | None:
if not self.rumble:
return None
return self.rumble.pop(0)
def make_links(port: str, uart: RecordingUART) -> dict[str, bridge.UartLink]:
return {port: bridge.UartLink(port, cast(PicoUART, cast(object, uart)))}
def make_config() -> bridge.BridgeConfig:
return bridge.BridgeConfig(
interval=10.0,
@ -88,29 +93,87 @@ def test_repeated_constant_rumble_stays_active_until_idle_timeout(
uart = RecordingUART()
controller = cast(sdl3.SDL_Gamepad, object())
ctx = bridge.ControllerContext(
controller,
7,
0,
"controller",
"/dev/null",
cast(PicoUART, cast(object, uart)),
)
ctx = bridge.ControllerContext(controller, 7, 0, "controller", "/dev/null")
links = make_links("/dev/null", uart)
contexts = {ctx.instance_id: ctx}
args = Namespace(baud=UART_BAUD)
console = Console(file=StringIO())
magnitude = (64 / 255.0, 192 / 255.0)
magnitude = (0, 64 / 255.0, 192 / 255.0)
uart.rumble.append(magnitude)
bridge.service_contexts(1.0, args, make_config(), contexts, [], console)
bridge.service_contexts(1.0, args, make_config(), contexts, links, console)
uart.rumble.append(magnitude)
bridge.service_contexts(1.7, args, make_config(), contexts, [], console)
bridge.service_contexts(1.71, args, make_config(), contexts, [], console)
bridge.service_contexts(1.7, args, make_config(), contexts, links, console)
bridge.service_contexts(1.71, args, make_config(), contexts, links, console)
assert calls == [(16448, 49344, 50), (16448, 49344, 50)]
assert ctx.rumble_active
bridge.service_contexts(1.96, args, make_config(), contexts, [], console)
bridge.service_contexts(1.96, args, make_config(), contexts, links, console)
assert calls[-1] == (0, 0, 0)
assert not ctx.rumble_active
def test_shared_port_routes_reports_and_rumble_by_slot(
monkeypatch: pytest.MonkeyPatch,
) -> None:
calls: list[tuple[object, int, int]] = []
monkeypatch.setattr(
bridge.sdl3,
"SDL_RumbleGamepad",
lambda controller, low, high, _duration: calls.append((controller, low, high)) or True,
)
monkeypatch.setattr(bridge, "poll_controller_buttons", lambda _ctx, _map: None)
uart = RecordingUART()
pad_a = cast(sdl3.SDL_Gamepad, object())
pad_b = cast(sdl3.SDL_Gamepad, object())
ctx_a = bridge.ControllerContext(pad_a, 7, 0, "a", "COM11", slot=0)
ctx_b = bridge.ControllerContext(pad_b, 8, 1, "b", "COM11", slot=2)
ctx_a.report.buttons = 0x0001
ctx_b.report.buttons = 0x0002
contexts = {7: ctx_a, 8: ctx_b}
links = make_links("COM11", uart)
args = Namespace(baud=UART_BAUD)
console = Console(file=StringIO())
# Slot 2 rumbles, slot 0 is idle, slot 3 has no controller attached.
uart.rumble.extend([(0, 0.0, 0.0), (2, 1.0, 0.5), (3, 1.0, 1.0)])
bridge.service_contexts(20.0, args, make_config(), contexts, links, console)
assert sorted(uart.sent) == [(0, 0x0001), (2, 0x0002)]
assert calls == [(pad_a, 0, 0), (pad_b, 0xFFFF, 0x7FFF)]
assert not ctx_a.rumble_active
assert ctx_b.rumble_active
def test_auto_pairing_spreads_controllers_across_ports_then_fills_slots() -> None:
pairing = bridge.PairingState(
mapping_by_index={},
available_ports=["COM11", "COM12"],
slots_per_port=2,
auto_pairing_enabled=True,
)
console = Console(file=StringIO())
assignments = [bridge.assign_port_for_index(pairing, idx, console) for idx in range(5)]
assert assignments == [("COM11", 0), ("COM12", 0), ("COM11", 1), ("COM12", 1), None]
# Releasing a slot makes exactly that slot reusable.
del pairing.mapping_by_index[2]
assert bridge.assign_port_for_index(pairing, 9, console) == ("COM11", 1)
def test_explicit_mappings_fill_omitted_slots_and_reject_conflicts() -> None:
parser = bridge.build_arg_parser()
resolved = bridge.resolve_mapping_slots(
[(0, "COM11", None), (1, "COM11", 3), (2, "COM11", None)], 4, parser
)
assert resolved == {0: ("COM11", 0), 1: ("COM11", 3), 2: ("COM11", 1)}
with pytest.raises(SystemExit):
bridge.resolve_mapping_slots([(0, "COM11", 1), (1, "COM11", 1)], 4, parser)
with pytest.raises(SystemExit):
bridge.resolve_mapping_slots([(0, "COM11", None), (1, "COM11", None)], 1, parser)
assert bridge.parse_mapping("2:COM11:3") == (2, "COM11", 3)
assert bridge.parse_mapping("0:/dev/ttyUSB0") == (0, "/dev/ttyUSB0", None)