Add stateful HD-rumble decoder with TinyUSB report normalization

Decode Nintendo HD-rumble words once in firmware into conventional low/high magnitudes and normalize stripped control SET_REPORT vs complete interrupt OUT reports through one path.
This commit is contained in:
Joey Yakimowich-Payne 2026-08-29 22:12:24 -06:00
commit bce94d38ce
21 changed files with 902 additions and 180 deletions

View file

@ -0,0 +1,218 @@
#include "switch_haptics.h"
#include <array>
#include <cstdint>
#include <iostream>
namespace {
int failures = 0;
void expect_output(const char* scenario, SwitchRumbleOutput actual,
uint8_t expected_low, uint8_t expected_high) {
if (actual.low_frequency_magnitude == expected_low &&
actual.high_frequency_magnitude == expected_high) {
return;
}
std::cerr << scenario << ": expected low/high "
<< static_cast<unsigned>(expected_low) << "/"
<< static_cast<unsigned>(expected_high) << ", got "
<< static_cast<unsigned>(actual.low_frequency_magnitude) << "/"
<< static_cast<unsigned>(actual.high_frequency_magnitude) << '\n';
++failures;
}
uint32_t type_2(uint8_t high_frequency, uint8_t high_amplitude,
uint8_t low_frequency, uint8_t low_amplitude) {
return (1u << 30u) |
((static_cast<uint32_t>(low_amplitude) & 0x7fu) << 23u) |
((static_cast<uint32_t>(low_frequency) & 0x7fu) << 16u) |
((static_cast<uint32_t>(high_amplitude) & 0x7fu) << 9u) |
((static_cast<uint32_t>(high_frequency) & 0x7fu) << 2u);
}
uint32_t type_1_one_sample(uint8_t high_command, uint8_t low_command) {
return (1u << 30u) |
((static_cast<uint32_t>(low_command) & 0x1fu) << 25u) |
((static_cast<uint32_t>(high_command) & 0x1fu) << 20u);
}
uint32_t type_1_three_samples(uint8_t high_0, uint8_t low_0,
uint8_t high_1, uint8_t low_1,
uint8_t high_2, uint8_t low_2) {
return (3u << 30u) |
((static_cast<uint32_t>(low_0) & 0x1fu) << 25u) |
((static_cast<uint32_t>(high_0) & 0x1fu) << 20u) |
((static_cast<uint32_t>(low_1) & 0x1fu) << 15u) |
((static_cast<uint32_t>(high_1) & 0x1fu) << 10u) |
((static_cast<uint32_t>(low_2) & 0x1fu) << 5u) |
(static_cast<uint32_t>(high_2) & 0x1fu);
}
std::array<uint8_t, 8> payload(uint32_t left, uint32_t right) {
std::array<uint8_t, 8> bytes{};
const uint32_t words[2] = {left, right};
for (unsigned actuator = 0; actuator < 2; ++actuator) {
const unsigned offset = actuator * 4u;
bytes[offset] = static_cast<uint8_t>(words[actuator]);
bytes[offset + 1u] = static_cast<uint8_t>(words[actuator] >> 8u);
bytes[offset + 2u] = static_cast<uint8_t>(words[actuator] >> 16u);
bytes[offset + 3u] = static_cast<uint8_t>(words[actuator] >> 24u);
}
return bytes;
}
void test_neutral_and_per_actuator_reset() {
constexpr uint32_t neutral = 0x40400100u;
SwitchHapticsDecoder decoder;
auto frame = payload(neutral, neutral);
expect_output("explicit neutral", decoder.decode(frame.data()), 0, 0);
frame = payload(type_2(90, 16, 50, 127), type_2(100, 32, 40, 16));
expect_output("active actuators", decoder.decode(frame.data()), 250, 32);
decoder.reset();
frame = payload(1u << 5u, 1u << 5u);
expect_output("explicit decoder reset", decoder.decode(frame.data()), 0, 0);
frame = payload(type_2(90, 16, 50, 127), type_2(100, 32, 40, 16));
decoder.decode(frame.data());
frame = payload(0, type_2(100, 32, 40, 16));
expect_output("zero resets only left actuator", decoder.decode(frame.data()), 16, 32);
frame = payload(0, neutral);
expect_output("neutral resets right actuator", decoder.decode(frame.data()), 0, 0);
}
void test_type_2_full_state_and_band_mapping() {
constexpr uint32_t neutral = 0x40400100u;
SwitchHapticsDecoder decoder;
const auto frame = payload(type_2(100, 32, 20, 16), neutral);
expect_output("type-2 low/high mapping", decoder.decode(frame.data()), 16, 32);
}
void test_type_1_relative_update_and_idempotence() {
constexpr uint32_t neutral = 0x40400100u;
SwitchHapticsDecoder decoder;
auto frame = payload(type_2(64, 16, 64, 16), neutral);
expect_output("relative update initial state", decoder.decode(frame.data()), 16, 16);
frame = payload(type_1_one_sample(17, 20), neutral);
expect_output("type-1 relative update", decoder.decode(frame.data()), 17, 18);
expect_output("identical delta is idempotent", decoder.decode(frame.data()), 17, 18);
}
void test_subsample_peak_and_repeated_current_state() {
constexpr uint32_t neutral = 0x40400100u;
SwitchHapticsDecoder decoder;
auto frame = payload(type_2(64, 16, 64, 16), neutral);
decoder.decode(frame.data());
frame = payload(type_1_three_samples(17, 17, 29, 29, 24, 24), neutral);
expect_output("peak across three subsamples", decoder.decode(frame.data()), 18, 18);
expect_output("repeat returns final cumulative state", decoder.decode(frame.data()), 16, 16);
}
void test_left_right_peak_combination() {
SwitchHapticsDecoder decoder;
const auto frame = payload(type_2(90, 1, 50, 127), type_2(100, 32, 40, 1));
expect_output("independent actuator band peaks", decoder.decode(frame.data()), 250, 32);
}
void test_type_3_and_type_4_frames() {
constexpr uint32_t neutral = 0x40400100u;
SwitchHapticsDecoder decoder;
auto frame = payload(type_2(64, 16, 64, 16), neutral);
decoder.decode(frame.data());
const uint32_t type3 = (2u << 30u) | 1u | (70u << 1u) |
(24u << 8u) | (17u << 13u) |
(20u << 18u) | (32u << 23u);
frame = payload(type3, neutral);
expect_output("type-3 full plus relative samples", decoder.decode(frame.data()), 18, 32);
const uint32_t type4_low_amplitude = (1u << 30u) | 2u | (32u << 23u);
frame = payload(type4_low_amplitude, neutral);
expect_output("type-4 low amplitude selection", decoder.decode(frame.data()), 32, 32);
const uint32_t type4_high_amplitude = (1u << 30u) | 3u | (127u << 23u);
frame = payload(type4_high_amplitude, neutral);
expect_output("type-4 high amplitude selection", decoder.decode(frame.data()), 32, 250);
}
void test_malformed_and_reserved_words_preserve_state() {
constexpr uint32_t neutral = 0x40400100u;
SwitchHapticsDecoder decoder;
auto frame = payload(type_2(100, 32, 20, 16), neutral);
decoder.decode(frame.data());
frame = payload((1u << 30u) | 1u, neutral);
expect_output("reserved type discriminator", decoder.decode(frame.data()), 16, 32);
frame = payload(1u << 5u, neutral);
expect_output("zero-frame word clears high band", decoder.decode(frame.data()), 16, 0);
}
void test_output_report_normalization() {
const uint8_t stripped[] = {
0x0a,
0x00, 0x01, 0x40, 0x40, 0x00, 0x01, 0x40, 0x40,
};
uint8_t output[64]{};
size_t size = normalize_switch_output_report(0x01, stripped, sizeof(stripped), output);
if (size != sizeof(stripped) + 1 || output[0] != 0x01 ||
output[1] != 0x0a || output[2] != 0x00 || output[9] != 0x40) {
std::cerr << "stripped 0x01 report normalization failed\n";
++failures;
}
size = normalize_switch_output_report(0x10, stripped, sizeof(stripped), output);
if (size != sizeof(stripped) + 1 || output[0] != 0x10 ||
output[1] != 0x0a || output[2] != 0x00 || output[9] != 0x40) {
std::cerr << "stripped 0x10 report normalization failed\n";
++failures;
}
const uint8_t complete[] = {
0x10, 0x0a,
0x00, 0x01, 0x40, 0x40, 0x00, 0x01, 0x40, 0x40,
};
size = normalize_switch_output_report(0, complete, sizeof(complete), output);
if (size != sizeof(complete) || output[0] != 0x10 ||
output[1] != 0x0a || output[9] != 0x40) {
std::cerr << "complete interrupt report normalization failed\n";
++failures;
}
std::array<uint8_t, 64> oversized{};
if (normalize_switch_output_report(0x01, oversized.data(), oversized.size(), output) != 0) {
std::cerr << "oversized stripped report was accepted\n";
++failures;
}
}
} // namespace
int main() {
test_neutral_and_per_actuator_reset();
test_type_2_full_state_and_band_mapping();
test_type_1_relative_update_and_idempotence();
test_subsample_peak_and_repeated_current_state();
test_left_right_peak_combination();
test_type_3_and_type_4_frames();
test_malformed_and_reserved_words_preserve_state();
test_output_report_normalization();
if (failures != 0) {
std::cerr << failures << " haptics test(s) failed\n";
return 1;
}
return 0;
}

View file

@ -31,7 +31,7 @@ class RecordingUART:
def send_report(self, report: SwitchReport) -> None:
self.sent_imu.append(tuple(report.imu_samples))
def read_rumble_payload(self) -> bytes | None:
def read_rumble(self) -> tuple[float, float] | None:
return None

View file

@ -0,0 +1,31 @@
from __future__ import annotations
import shutil
import subprocess
from pathlib import Path
def test_switch_haptics_native(tmp_path: Path) -> None:
root = Path(__file__).resolve().parents[1]
compiler = shutil.which("c++") or shutil.which("g++")
assert compiler is not None, "a host C++ compiler is required"
executable = tmp_path / "switch_haptics_test"
subprocess.run(
[
compiler,
"-std=c++17",
"-Wall",
"-Wextra",
"-Werror",
"-pedantic",
f"-I{root}",
str(root / "switch_haptics.cpp"),
str(root / "tests" / "switch_haptics_test.cpp"),
"-o",
str(executable),
],
check=True,
cwd=root,
)
subprocess.run([str(executable)], check=True, cwd=root)

View file

@ -6,8 +6,11 @@ from switch_pico_bridge.switch_pico_uart import (
SwitchReport,
IMUSample,
SwitchDpad,
PicoUART,
UART_HEADER,
UART_PROTOCOL_VERSION,
RUMBLE_HEADER,
RUMBLE_TYPE_DECODED,
ACCEL_LSB_PER_G,
GYRO_LSB_PER_RAD_S,
MS2_PER_G,
@ -15,6 +18,36 @@ from switch_pico_bridge.switch_pico_uart import (
)
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) -> 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()
return uart, serial_port
def test_v2_frame_with_imu_samples():
"""V2 frame with 3 IMU samples should be 48 bytes with correct layout."""
r = SwitchReport(
@ -118,3 +151,34 @@ def test_max_imu_samples_capped():
assert len(data) == 48 # 3 samples, not 5
assert data[10] == 3
assert data[2] == 44 # payload_len for 3 samples
def test_decoded_rumble_frame_survives_fragmented_input():
frame = make_rumble_frame(64, 192)
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))
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))
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))
assert uart.read_rumble() == pytest.approx((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))
assert uart.read_rumble() == (0.0, 0.0)
assert uart.read_rumble() == (1.0, 1.0)

101
tests/test_uart_rumble.py Normal file
View file

@ -0,0 +1,101 @@
"""Focused tests for decoded UART rumble delivery to SDL3."""
from argparse import Namespace
from io import StringIO
from typing import cast
import pytest
import sdl3
from rich.console import Console
import switch_pico_bridge.controller_uart_bridge as bridge
from switch_pico_bridge.switch_pico_uart import PicoUART, SwitchReport, UART_BAUD
class RecordingUART:
def __init__(self) -> None:
self.rumble: list[tuple[float, float]] = []
def send_report(self, _report: SwitchReport) -> None:
pass
def read_rumble(self) -> tuple[float, float] | None:
if not self.rumble:
return None
return self.rumble.pop(0)
def make_config() -> bridge.BridgeConfig:
return bridge.BridgeConfig(
interval=10.0,
deadzone_raw=0,
trigger_threshold=0,
zero_sticks=False,
zero_hotkey="",
swap_hotkey="",
button_map_default={},
button_map_swapped={},
swap_abxy_indices=set(),
swap_abxy_ids=set(),
swap_abxy_global=False,
no_imu=True,
)
def test_apply_rumble_maps_low_and_high_with_50ms_duration(
monkeypatch: pytest.MonkeyPatch,
) -> None:
calls: list[tuple[int, int, int]] = []
monkeypatch.setattr(
bridge.sdl3,
"SDL_RumbleGamepad",
lambda _controller, low, high, duration: calls.append((low, high, duration)),
)
controller = cast(sdl3.SDL_Gamepad, object())
assert bridge.apply_rumble(controller, 1.0, 0.5)
assert calls[-1] == (0xFFFF, 0x7FFF, 50)
assert not bridge.apply_rumble(controller, 0.0, 0.0)
assert calls[-1] == (0, 0, 50)
def test_repeated_constant_rumble_stays_active_until_idle_timeout(
monkeypatch: pytest.MonkeyPatch,
) -> None:
calls: list[tuple[int, int, int]] = []
monkeypatch.setattr(
bridge.sdl3,
"SDL_RumbleGamepad",
lambda _controller, low, high, duration: calls.append((low, high, duration)),
)
monkeypatch.setattr(bridge, "poll_controller_buttons", lambda _ctx, _map: None)
uart = RecordingUART()
controller = cast(sdl3.SDL_Gamepad, object())
ctx = bridge.ControllerContext(
controller,
7,
0,
"controller",
"/dev/null",
cast(PicoUART, cast(object, uart)),
)
contexts = {ctx.instance_id: ctx}
args = Namespace(baud=UART_BAUD)
console = Console(file=StringIO())
magnitude = (64 / 255.0, 192 / 255.0)
uart.rumble.append(magnitude)
bridge.service_contexts(1.0, args, make_config(), contexts, [], 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)
assert calls == [(16448, 49344, 50), (16448, 49344, 50)]
assert ctx.rumble_active
bridge.service_contexts(1.96, args, make_config(), contexts, [], console)
assert calls[-1] == (0, 0, 0)
assert not ctx.rumble_active