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:
parent
20c9c98f89
commit
bce94d38ce
21 changed files with 902 additions and 180 deletions
218
tests/switch_haptics_test.cpp
Normal file
218
tests/switch_haptics_test.cpp
Normal 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;
|
||||
}
|
||||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
31
tests/test_switch_haptics_native.py
Normal file
31
tests/test_switch_haptics_native.py
Normal 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)
|
||||
|
|
@ -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
101
tests/test_uart_rumble.py
Normal 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
|
||||
Loading…
Add table
Add a link
Reference in a new issue