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

@ -87,6 +87,7 @@ endif()
add_executable(switch-pico
switch-pico.cpp
switch_pro_driver.cpp
switch_haptics.cpp
)
if(SWITCH_PICO_INPUT_BACKEND STREQUAL "BLUEPAD32")
target_sources(switch-pico PRIVATE bluepad32_input_backend.cpp)

View file

@ -102,6 +102,18 @@ RUMBLE (force feedback)
-> [SDL3 haptics] -> [Any controller motors]
```
### HD rumble translation
Nintendo sends two stateful four-byte HD-rumble actuator words. Each word can carry full or relative high/low frequency and amplitude commands with up to three subsamples; amplitude uses a logarithmic curve. The Pico decodes both words once in `SwitchHapticsDecoder`, retains actuator state across packets, and reduces the result to conventional low/strong and high/weak motor magnitudes. SDL3 and Bluepad32 cannot reproduce the original linear-actuator frequencies or left/right spatial effects, but they receive the correct nonlinear band amplitudes.
The UART return frame carries the decoded result rather than raw HD-rumble bytes:
```text
0xBB, 0x02, low-frequency magnitude, high-frequency magnitude, checksum
```
The checksum is the sum of the first four bytes modulo 256. Firmware and Python bridge versions from before this change are not rumble-protocol compatible; controller input framing remains unchanged.
## Hardware wiring (Pico)
- UART1 pins (fixed in firmware):
- **TX**: GPIO4 (Pico pin 6) → RX of your USB-serial adapter.

View file

@ -2,7 +2,6 @@
#include <limits.h>
#include <stddef.h>
#include <string.h>
#include <btstack_run_loop.h>
#include <pico/critical_section.h>
@ -23,10 +22,6 @@ constexpr uint16_t kRumbleDurationMs = 50;
constexpr uint32_t kRumblePollIntervalMs = 5;
constexpr uint kRumbleQueueDepth = 8;
struct RumblePacket {
uint8_t bytes[8];
};
critical_section_t g_state_lock;
queue_t g_rumble_queue;
SwitchInputState g_shared_state;
@ -172,28 +167,9 @@ SwitchInputState map_gamepad(const uni_gamepad_t& gamepad) {
return state;
}
void decode_rumble(const uint8_t bytes[8], uint8_t* left_magnitude, uint8_t* right_magnitude) {
static constexpr uint8_t kNeutralPacket[8] = {0x00, 0x01, 0x40, 0x40, 0x00, 0x01, 0x40, 0x40};
if (memcmp(bytes, kNeutralPacket, sizeof(kNeutralPacket)) == 0) {
*left_magnitude = 0;
*right_magnitude = 0;
return;
}
uint16_t right_raw = static_cast<uint16_t>(((bytes[1] & 0x03) << 8) | bytes[0]);
uint16_t left_raw = static_cast<uint16_t>(((bytes[5] & 0x03) << 8) | bytes[4]);
if (left_raw < 8 && right_raw < 8) {
left_raw = 0;
right_raw = 0;
}
*left_magnitude = static_cast<uint8_t>((left_raw * UINT8_MAX + 511) / 1023);
*right_magnitude = static_cast<uint8_t>((right_raw * UINT8_MAX + 511) / 1023);
}
void process_rumble_timer(btstack_timer_source_t* timer) {
RumblePacket packet{};
RumblePacket latest{};
SwitchRumbleOutput packet{};
SwitchRumbleOutput latest{};
bool have_packet = false;
while (queue_try_remove(&g_rumble_queue, &packet)) {
latest = packet;
@ -202,13 +178,9 @@ void process_rumble_timer(btstack_timer_source_t* timer) {
if (have_packet && g_active_device != nullptr &&
g_active_device->report_parser.play_dual_rumble != nullptr) {
uint8_t left_magnitude = 0;
uint8_t right_magnitude = 0;
decode_rumble(latest.bytes, &left_magnitude, &right_magnitude);
// Bluepad orders the weak (high-frequency) motor before the strong
// (low-frequency) motor; the project decoder names those right/left.
g_active_device->report_parser.play_dual_rumble(
g_active_device, 0, kRumbleDurationMs, right_magnitude, left_magnitude);
g_active_device, 0, kRumbleDurationMs,
latest.high_frequency_magnitude, latest.low_frequency_magnitude);
}
btstack_run_loop_set_timer(timer, kRumblePollIntervalMs);
@ -337,7 +309,7 @@ void bluepad32_input_backend_init() {
}
critical_section_init(&g_state_lock);
queue_init(&g_rumble_queue, sizeof(RumblePacket), kRumbleQueueDepth);
queue_init(&g_rumble_queue, sizeof(SwitchRumbleOutput), kRumbleQueueDepth);
g_shared_state = make_neutral_state();
g_shared_controller_active = false;
g_shared_generation = 0;
@ -387,16 +359,14 @@ void bluepad32_input_backend_report_sent() {
g_consumed_generation = g_last_snapshot_generation;
}
void bluepad32_input_backend_queue_rumble(const uint8_t rumble[8]) {
if (!g_initialized || rumble == nullptr) {
void bluepad32_input_backend_queue_rumble(const SwitchRumbleOutput& rumble) {
if (!g_initialized) {
return;
}
RumblePacket packet{};
memcpy(packet.bytes, rumble, sizeof(packet.bytes));
if (!queue_try_add(&g_rumble_queue, &packet)) {
RumblePacket discarded{};
if (!queue_try_add(&g_rumble_queue, &rumble)) {
SwitchRumbleOutput discarded{};
(void)queue_try_remove(&g_rumble_queue, &discarded);
(void)queue_try_add(&g_rumble_queue, &packet);
(void)queue_try_add(&g_rumble_queue, &rumble);
}
}

View file

@ -3,9 +3,10 @@
#include <stdint.h>
#include "switch_pro_driver.h"
#include "switch_haptics.h"
void bluepad32_input_backend_init();
void bluepad32_input_backend_start();
bool bluepad32_input_backend_snapshot(SwitchInputState* out);
void bluepad32_input_backend_report_sent();
void bluepad32_input_backend_queue_rumble(const uint8_t rumble[8]);
void bluepad32_input_backend_queue_rumble(const SwitchRumbleOutput& rumble);

Binary file not shown.

Binary file not shown.

Binary file not shown.

Binary file not shown.

View file

@ -10,7 +10,6 @@ from .switch_pico_uart import ( # noqa: F401
SwitchDpad,
SwitchUARTClient,
axis_to_stick,
decode_rumble,
discover_serial_ports,
first_serial_port,
str_to_dpad,
@ -24,7 +23,6 @@ __all__ = [
"discover_serial_ports",
"first_serial_port",
"axis_to_stick",
"decode_rumble",
"str_to_dpad",
"trigger_to_button",
]

View file

@ -1,16 +1,17 @@
#!/usr/bin/env python3
"""
Bridge multiple SDL2 controllers to switch-pico over UART and mirror rumble back.
Bridge multiple SDL3 controllers to switch-pico over UART and mirror rumble back.
The framing matches ``switch-pico.cpp``:
- Host -> Pico : 0xAA, buttons (LE16), hat, lx, ly, rx, ry
- Pico -> Host : 0xBB, 0x01, 8 rumble bytes, checksum (sum of first 10 bytes)
- Host -> Pico : UART v2 controller report
- Pico -> Host : 0xBB, 0x02, low-frequency magnitude,
high-frequency magnitude, checksum
Features inspired by ``host/controller_bridge.py``:
- Multiple controllers paired to multiple UART ports
- Rich-powered interactive pairing UI
- Adjustable send frequency, deadzone, and trigger thresholds
- Rumble feedback delivered to SDL2 controllers
- Rumble feedback delivered to SDL3 controllers
"""
from __future__ import annotations
@ -48,15 +49,12 @@ from .switch_pico_uart import (
SwitchReport,
axis_to_stick,
str_to_dpad,
decode_rumble,
discover_serial_ports,
trigger_to_button,
)
RUMBLE_IDLE_TIMEOUT = 0.25 # seconds without packets before forcing rumble off
RUMBLE_STUCK_TIMEOUT = 0.60 # continuous same-energy rumble will be stopped after this
RUMBLE_MIN_ACTIVE = 0.40 # below this, rumble is treated as off/noise
RUMBLE_SCALE = 1.0
RUMBLE_DURATION_MS = 50
CONTROLLER_DB_URL_DEFAULT = "https://raw.githubusercontent.com/mdqinc/SDL_GameControllerDB/refs/heads/master/gamecontrollerdb.txt"
SDL_TRUE = True
SDL_EVENT_GAMEPAD_SENSOR_UPDATE = getattr(sdl3, "SDL_EVENT_GAMEPAD_SENSOR_UPDATE", 0x658)
@ -199,21 +197,16 @@ def interactive_pairing(
return mappings
def apply_rumble(controller: sdl3.SDL_Gamepad, payload: bytes) -> float:
"""Apply rumble payload to SDL controller and return max normalized energy."""
left_norm, right_norm = decode_rumble(payload)
max_norm = max(left_norm, right_norm)
# Treat small rumble as "off" to avoid idle buzz.
if max_norm < RUMBLE_MIN_ACTIVE:
sdl3.SDL_RumbleGamepad(controller, 0, 0, 0)
return 0.0
# Attenuate to feel closer to a real controller; cap at ~25% strength.
scale = RUMBLE_SCALE
low = int(min(1.0, left_norm * scale) * 0xFFFF) # SDL: low_frequency_rumble
high = int(min(1.0, right_norm * scale) * 0xFFFF) # SDL: high_frequency_rumble
duration = 10
sdl3.SDL_RumbleGamepad(controller, low, high, duration)
return max_norm
def apply_rumble(
controller: sdl3.SDL_Gamepad,
low_frequency: float,
high_frequency: float,
) -> bool:
"""Apply normalized low/high rumble magnitudes to an SDL controller."""
low = int(max(0.0, min(1.0, low_frequency)) * 0xFFFF)
high = int(max(0.0, min(1.0, high_frequency)) * 0xFFFF)
sdl3.SDL_RumbleGamepad(controller, low, high, RUMBLE_DURATION_MS)
return low != 0 or high != 0
@dataclass
@ -239,9 +232,7 @@ class ControllerContext:
)
last_send: float = 0.0
last_reopen_attempt: float = 0.0
last_rumble: float = 0.0
last_rumble_change: float = 0.0
last_rumble_energy: float = 0.0
last_rumble_at: float = 0.0
rumble_active: bool = False
axis_offsets: Dict[int, int] = field(default_factory=dict)
swap_abxy: bool = False
@ -1134,7 +1125,6 @@ def handle_removed_port(
ctx.uart = None
ctx.port = None
ctx.rumble_active = False
ctx.last_rumble_energy = 0.0
ctx.last_reopen_attempt = time.monotonic()
console.print(
f"[yellow]UART {path} removed; controller {ctx.controller_index} waiting for reassignment[/yellow]"
@ -1574,32 +1564,25 @@ def service_contexts(
ctx.uart.send_report(ctx.report)
ctx.last_send = now
last_payload = None
latest_rumble = None
while True:
p = ctx.uart.read_rumble_payload()
if not p:
rumble = ctx.uart.read_rumble()
if rumble is None:
break
last_payload = p
latest_rumble = rumble
if last_payload is not None:
# Apply only the freshest rumble payload seen during this tick.
energy = apply_rumble(ctx.controller, last_payload)
ctx.rumble_active = energy >= RUMBLE_MIN_ACTIVE
if ctx.rumble_active and energy != ctx.last_rumble_energy:
ctx.last_rumble_change = now
ctx.last_rumble_energy = energy
ctx.last_rumble = now
elif ctx.rumble_active and (now - ctx.last_rumble) > RUMBLE_IDLE_TIMEOUT:
sdl3.SDL_RumbleGamepad(ctx.controller, 0, 0, 0)
ctx.rumble_active = False
ctx.last_rumble_energy = 0.0
if latest_rumble is not None:
# Apply only the freshest rumble command seen during this tick.
ctx.rumble_active = apply_rumble(
ctx.controller, latest_rumble[0], latest_rumble[1]
)
ctx.last_rumble_at = now
elif (
ctx.rumble_active
and (now - ctx.last_rumble_change) > RUMBLE_STUCK_TIMEOUT
and (now - ctx.last_rumble_at) > RUMBLE_IDLE_TIMEOUT
):
sdl3.SDL_RumbleGamepad(ctx.controller, 0, 0, 0)
ctx.rumble_active = False
ctx.last_rumble_energy = 0.0
except SerialException as exc:
console.print(f"[yellow]UART {ctx.port} disconnected: {exc}[/yellow]")
try:
@ -1609,7 +1592,6 @@ def service_contexts(
sdl3.SDL_RumbleGamepad(ctx.controller, 0, 0, 0)
ctx.uart = None
ctx.rumble_active = False
ctx.last_rumble_energy = 0.0
ctx.last_reopen_attempt = now
except Exception as exc:
console.print(f"[red]UART error on {ctx.port}: {exc}[/red]")

View file

@ -2,12 +2,13 @@
"""
Lightweight helpers for talking to the switch-pico firmware over UART.
This module exposes the raw report structure plus a small convenience wrapper
This module exposes the report structure plus a small convenience wrapper
so other scripts can do things like "press a button" or "move a stick" without
depending on SDL. It mirrors the framing in ``switch-pico.cpp``:
Host -> Pico : 0xAA, buttons (LE16), hat, lx, ly, rx, ry
Pico -> Host : 0xBB, 0x01, 8 rumble bytes, checksum (sum of first 10 bytes)
Host -> Pico : UART v2 controller report
Pico -> Host : 0xBB, 0x02, low-frequency magnitude, high-frequency magnitude,
checksum (sum of the first 4 bytes)
"""
from __future__ import annotations
@ -26,7 +27,7 @@ from serial.tools import list_ports, list_ports_common
UART_HEADER = 0xAA
UART_PROTOCOL_VERSION = 0x02
RUMBLE_HEADER = 0xBB
RUMBLE_TYPE_RUMBLE = 0x01
RUMBLE_TYPE_DECODED = 0x02
UART_BAUD = 921600
IMU_SAMPLES_PER_REPORT = 3
@ -300,15 +301,16 @@ class PicoUART:
"""Send a controller report to the Pico."""
self.serial.write(report.to_bytes())
def read_rumble_payload(self) -> Optional[bytes]:
def read_rumble(self) -> Optional[Tuple[float, float]]:
"""
Drain available UART bytes into an internal buffer, then extract one rumble frame.
Extract one decoded rumble frame as normalized low/high magnitudes.
Frame format:
0: 0xBB (RUMBLE_HEADER)
1: type (0x01 for rumble)
2-9: 8-byte rumble payload
10: checksum (sum of first 10 bytes) & 0xFF
1: type (0x02 for decoded rumble)
2: low-frequency magnitude (0-255)
3: high-frequency magnitude (0-255)
4: checksum (sum of first 4 bytes) & 0xFF
"""
waiting = self.serial.in_waiting
if waiting:
@ -323,18 +325,18 @@ class PicoUART:
self._buffer.clear()
return None
if len(self._buffer) - start < 11:
if len(self._buffer) - start < 5:
if start > 0:
del self._buffer[:start]
return None
frame = self._buffer[start : start + 11]
checksum = compute_checksum(bytes(frame[:10]))
frame = self._buffer[start : start + 5]
checksum = compute_checksum(bytes(frame[:4]))
if frame[1] == RUMBLE_TYPE_RUMBLE and checksum == frame[10]:
payload = bytes(frame[2:10])
del self._buffer[: start + 11]
return payload
if frame[1] == RUMBLE_TYPE_DECODED and checksum == frame[4]:
rumble = (frame[2] / 255.0, frame[3] / 255.0)
del self._buffer[: start + 5]
return rumble
del self._buffer[: start + 1]
@ -343,21 +345,6 @@ class PicoUART:
self.serial.close()
def decode_rumble(payload: bytes) -> Tuple[float, float]:
"""Return normalized rumble amplitudes (0.0-1.0) for left/right."""
if len(payload) < 8:
return 0.0, 0.0
if payload == b"\x00\x01\x40\x40\x00\x01\x40\x40":
return 0.0, 0.0
right_raw = ((payload[1] & 0x03) << 8) | payload[0]
left_raw = ((payload[5] & 0x03) << 8) | payload[4]
if left_raw < 8 and right_raw < 8:
return 0.0, 0.0
left = min(max(left_raw / 1023.0, 0.0), 1.0)
right = min(max(right_raw / 1023.0, 0.0), 1.0)
return left, right
@dataclass
class SwitchControllerState:
"""Mutable controller state with helpers for building reports."""
@ -537,13 +524,10 @@ class SwitchUARTClient:
def poll_rumble(self) -> Optional[Tuple[float, float]]:
"""
Poll for the latest rumble payload and return normalized amplitudes.
Poll for decoded low/high rumble magnitudes normalized to 0.0-1.0.
Returns None if no rumble frame was available.
"""
payload = self.uart.read_rumble_payload()
if payload:
return decode_rumble(payload)
return None
return self.uart.read_rumble()
def close(self) -> None:
if self._auto_thread:

View file

@ -1,5 +1,4 @@
#include <stdio.h>
#include <string.h>
#include "bsp/board.h"
#include "pico/stdlib.h"
#include "tusb.h"
@ -23,7 +22,7 @@
#define UART_TX_PIN 4
#define UART_RX_PIN 5
#define UART_RUMBLE_HEADER 0xBB
#define UART_RUMBLE_RUMBLE_TYPE 0x01
#define UART_RUMBLE_TYPE 0x02
#endif
static bool g_last_mounted = false;
@ -51,22 +50,23 @@ static SwitchInputState neutral_input() {
}
#ifndef SWITCH_PICO_BLUEPAD32
static void send_rumble_uart_frame(const uint8_t rumble[8]) {
uint8_t frame[11];
frame[0] = UART_RUMBLE_HEADER;
frame[1] = UART_RUMBLE_RUMBLE_TYPE;
memcpy(&frame[2], rumble, 8);
static void send_rumble_uart_frame(const SwitchRumbleOutput& rumble) {
uint8_t frame[5] = {
UART_RUMBLE_HEADER,
UART_RUMBLE_TYPE,
rumble.low_frequency_magnitude,
rumble.high_frequency_magnitude,
0,
};
uint8_t checksum = 0;
for (int i = 0; i < 10; ++i) {
checksum = static_cast<uint8_t>(checksum + frame[i]);
for (uint8_t i = 0; i < 4; ++i) {
frame[4] = static_cast<uint8_t>(frame[4] + frame[i]);
}
frame[10] = checksum;
uart_write_blocking(UART_ID, frame, sizeof(frame));
}
#endif
static void on_rumble_from_switch(const uint8_t rumble[8]) {
static void on_rumble_from_switch(const SwitchRumbleOutput& rumble) {
#ifdef SWITCH_PICO_BLUEPAD32
bluepad32_input_backend_queue_rumble(rumble);
#else

306
switch_haptics.cpp Normal file
View file

@ -0,0 +1,306 @@
#include "switch_haptics.h"
#include <cmath>
#include <cstring>
namespace {
enum class CommandAction : uint8_t {
Ignore,
Default,
Substitute,
Sum,
};
struct HapticCommand {
CommandAction amplitude_action;
CommandAction frequency_action;
int16_t amplitude_offset;
int16_t frequency_offset;
};
constexpr HapticCommand kCommands[32] = {
{CommandAction::Default, CommandAction::Default, 0, 0},
{CommandAction::Substitute, CommandAction::Ignore, 0, 0},
{CommandAction::Substitute, CommandAction::Ignore, 240, 0},
{CommandAction::Substitute, CommandAction::Ignore, 224, 0},
{CommandAction::Substitute, CommandAction::Ignore, 208, 0},
{CommandAction::Substitute, CommandAction::Ignore, 192, 0},
{CommandAction::Substitute, CommandAction::Ignore, 176, 0},
{CommandAction::Substitute, CommandAction::Ignore, 160, 0},
{CommandAction::Substitute, CommandAction::Ignore, 144, 0},
{CommandAction::Substitute, CommandAction::Ignore, 128, 0},
{CommandAction::Substitute, CommandAction::Ignore, 112, 0},
{CommandAction::Substitute, CommandAction::Ignore, 96, 0},
{CommandAction::Ignore, CommandAction::Substitute, 0, 5},
{CommandAction::Ignore, CommandAction::Substitute, 0, 5},
{CommandAction::Ignore, CommandAction::Substitute, 0, 0},
{CommandAction::Ignore, CommandAction::Substitute, 0, 7},
{CommandAction::Ignore, CommandAction::Substitute, 0, 7},
{CommandAction::Sum, CommandAction::Sum, 4, 1},
{CommandAction::Sum, CommandAction::Ignore, 4, 0},
{CommandAction::Sum, CommandAction::Sum, 4, -1},
{CommandAction::Sum, CommandAction::Sum, 1, 1},
{CommandAction::Sum, CommandAction::Ignore, 1, 0},
{CommandAction::Sum, CommandAction::Sum, 1, -1},
{CommandAction::Ignore, CommandAction::Sum, 0, 1},
{CommandAction::Ignore, CommandAction::Ignore, 0, 0},
{CommandAction::Ignore, CommandAction::Sum, 0, -1},
{CommandAction::Sum, CommandAction::Sum, -1, 1},
{CommandAction::Sum, CommandAction::Ignore, -1, 0},
{CommandAction::Sum, CommandAction::Sum, -1, -1},
{CommandAction::Sum, CommandAction::Sum, -4, 1},
{CommandAction::Sum, CommandAction::Ignore, -4, 0},
{CommandAction::Sum, CommandAction::Sum, -4, -1},
};
constexpr uint32_t kNeutralWord = 0x40400100u;
constexpr uint8_t kDefaultFrequency = 64;
template <unsigned Shift, uint32_t Mask>
constexpr uint8_t extract(uint32_t word) {
static_assert(Shift < 32u, "32-bit word extraction shift must be bounded");
static_assert(Mask <= 0xffu && Mask <= (0xffffffffu >> Shift),
"word extraction mask must fit the shifted byte");
return static_cast<uint8_t>((word >> Shift) & Mask);
}
uint8_t apply_command(CommandAction action, int16_t offset, uint8_t current,
uint8_t default_value, uint8_t maximum) {
switch (action) {
case CommandAction::Ignore:
return current;
case CommandAction::Default:
return default_value;
case CommandAction::Substitute:
return static_cast<uint8_t>(offset);
case CommandAction::Sum: {
int result = static_cast<int>(current) + static_cast<int>(offset);
if (result < 0) {
result = 0;
} else if (result > maximum) {
result = maximum;
}
return static_cast<uint8_t>(result);
}
}
return default_value;
}
uint8_t host_amplitude_to_lut_index(uint8_t host_index) {
const unsigned index = host_index & 0x7fu;
if (index == 0) {
return 0;
}
if (index < 16) {
return static_cast<uint8_t>(7u + 8u * index);
}
if (index < 32) {
return static_cast<uint8_t>(97u + 2u * index);
}
return static_cast<uint8_t>(128u + index);
}
uint32_t load_little_endian_word(const uint8_t* bytes) {
return static_cast<uint32_t>(bytes[0]) |
(static_cast<uint32_t>(bytes[1]) << 8u) |
(static_cast<uint32_t>(bytes[2]) << 16u) |
(static_cast<uint32_t>(bytes[3]) << 24u);
}
} // namespace
size_t normalize_switch_output_report(uint8_t report_id,
const uint8_t* payload,
size_t payload_size,
uint8_t output[64]) {
if (payload == nullptr || output == nullptr) {
return 0;
}
if (report_id == 0) {
if (payload_size > 64) {
return 0;
}
std::memcpy(output, payload, payload_size);
return payload_size;
}
if (payload_size >= 64) {
return 0;
}
output[0] = report_id;
std::memcpy(output + 1, payload, payload_size);
return payload_size + 1;
}
SwitchHapticsDecoder::SwitchHapticsDecoder() {
reset();
}
void SwitchHapticsDecoder::reset_actuator(ActuatorState& state) {
state.high_amplitude = 0;
state.low_amplitude = 0;
state.high_frequency = kDefaultFrequency;
state.low_frequency = kDefaultFrequency;
state.last_word = 0;
state.have_last_word = false;
}
void SwitchHapticsDecoder::reset() {
reset_actuator(actuators_[0]);
reset_actuator(actuators_[1]);
}
SwitchHapticsDecoder::AmplitudePeak SwitchHapticsDecoder::decode_actuator(
ActuatorState& state, uint32_t word) {
if (word == 0 || word == kNeutralWord) {
reset_actuator(state);
state.last_word = word;
state.have_last_word = true;
return {0, 0};
}
if (state.have_last_word && state.last_word == word) {
return {state.low_amplitude, state.high_amplitude};
}
state.last_word = word;
state.have_last_word = true;
AmplitudePeak peak{0, 0};
bool decoded = false;
const uint8_t frame_count = extract<30u, 0x03u>(word);
const uint32_t data = word & 0x3fffffffu;
if (frame_count == 0) {
state.high_amplitude = 0;
return {state.low_amplitude, 0};
}
const auto record_sample = [&]() {
if (state.low_amplitude > peak.low) {
peak.low = state.low_amplitude;
}
if (state.high_amplitude > peak.high) {
peak.high = state.high_amplitude;
}
};
const auto apply_pair = [&](bool high_band, uint8_t command_index) {
const HapticCommand& command = kCommands[command_index & 0x1fu];
uint8_t& amplitude = high_band ? state.high_amplitude : state.low_amplitude;
uint8_t& frequency = high_band ? state.high_frequency : state.low_frequency;
amplitude = apply_command(command.amplitude_action, command.amplitude_offset,
amplitude, 0, 255);
frequency = apply_command(command.frequency_action, command.frequency_offset,
frequency, kDefaultFrequency, 127);
};
const auto decode_type_1 = [&]() {
const uint8_t high_commands[3] = {
extract<20u, 0x1fu>(word),
extract<10u, 0x1fu>(word),
extract<0u, 0x1fu>(word),
};
const uint8_t low_commands[3] = {
extract<25u, 0x1fu>(word),
extract<15u, 0x1fu>(word),
extract<5u, 0x1fu>(word),
};
for (uint8_t sample = 0; sample < frame_count; ++sample) {
apply_pair(true, high_commands[sample]);
apply_pair(false, low_commands[sample]);
record_sample();
}
decoded = true;
};
if (frame_count == 1) {
if ((data & 0x000fffffu) == 0) {
decode_type_1();
} else if ((data & 0x03u) == 0) {
state.high_frequency = extract<2u, 0x7fu>(word);
state.high_amplitude = host_amplitude_to_lut_index(extract<9u, 0x7fu>(word));
state.low_frequency = extract<16u, 0x7fu>(word);
state.low_amplitude = host_amplitude_to_lut_index(extract<23u, 0x7fu>(word));
record_sample();
decoded = true;
} else if ((data & 0x02u) != 0) {
const bool high_band = extract<0u, 0x01u>(word) != 0;
const bool frequency_selected = extract<2u, 0x01u>(word) != 0;
const uint8_t value = extract<23u, 0x7fu>(word);
if (frequency_selected) {
if (high_band) {
state.high_frequency = value;
} else {
state.low_frequency = value;
}
} else if (high_band) {
state.high_amplitude = host_amplitude_to_lut_index(value);
} else {
state.low_amplitude = host_amplitude_to_lut_index(value);
}
record_sample();
decoded = true;
}
} else if (frame_count == 2) {
if ((data & 0x03ffu) == 0) {
decode_type_1();
} else {
const bool high_band = extract<0u, 0x01u>(word) != 0;
const uint8_t frequency = extract<1u, 0x7fu>(word);
const uint8_t command = extract<18u, 0x1fu>(word);
const uint8_t amplitude = host_amplitude_to_lut_index(extract<23u, 0x7fu>(word));
if (high_band) {
state.high_frequency = frequency;
state.high_amplitude = amplitude;
apply_pair(false, command);
} else {
state.low_frequency = frequency;
state.low_amplitude = amplitude;
apply_pair(true, command);
}
record_sample();
apply_pair(true, extract<8u, 0x1fu>(word));
apply_pair(false, extract<13u, 0x1fu>(word));
record_sample();
decoded = true;
}
} else if (frame_count == 3) {
decode_type_1();
}
if (!decoded) {
return {state.low_amplitude, state.high_amplitude};
}
return peak;
}
uint8_t SwitchHapticsDecoder::amplitude_to_magnitude(uint8_t amplitude_index) {
if (amplitude_index < 2) {
return 0;
}
const double exponent = -8.0 + static_cast<double>(amplitude_index) / 32.0;
const double scaled = std::exp2(exponent) * 255.0;
unsigned magnitude = static_cast<unsigned>(scaled + 0.5);
if (magnitude > 255u) {
magnitude = 255u;
}
return static_cast<uint8_t>(magnitude);
}
SwitchRumbleOutput SwitchHapticsDecoder::decode(const uint8_t payload[8]) {
AmplitudePeak peaks[2] = {
{actuators_[0].low_amplitude, actuators_[0].high_amplitude},
{actuators_[1].low_amplitude, actuators_[1].high_amplitude},
};
if (payload != nullptr) {
peaks[0] = decode_actuator(actuators_[0], load_little_endian_word(payload));
peaks[1] = decode_actuator(actuators_[1], load_little_endian_word(payload + 4));
}
const uint8_t low_peak = peaks[0].low > peaks[1].low ? peaks[0].low : peaks[1].low;
const uint8_t high_peak = peaks[0].high > peaks[1].high ? peaks[0].high : peaks[1].high;
return {amplitude_to_magnitude(low_peak), amplitude_to_magnitude(high_peak)};
}

46
switch_haptics.h Normal file
View file

@ -0,0 +1,46 @@
#ifndef SWITCH_HAPTICS_H
#define SWITCH_HAPTICS_H
#include <stddef.h>
#include <stdint.h>
struct SwitchRumbleOutput {
uint8_t low_frequency_magnitude;
uint8_t high_frequency_magnitude;
};
size_t normalize_switch_output_report(uint8_t report_id,
const uint8_t* payload,
size_t payload_size,
uint8_t output[64]);
class SwitchHapticsDecoder {
public:
SwitchHapticsDecoder();
void reset();
SwitchRumbleOutput decode(const uint8_t payload[8]);
private:
struct ActuatorState {
uint8_t high_amplitude;
uint8_t low_amplitude;
uint8_t high_frequency;
uint8_t low_frequency;
uint32_t last_word;
bool have_last_word;
};
struct AmplitudePeak {
uint8_t low;
uint8_t high;
};
static void reset_actuator(ActuatorState& state);
static AmplitudePeak decode_actuator(ActuatorState& state, uint32_t word);
static uint8_t amplitude_to_magnitude(uint8_t amplitude_index);
ActuatorState actuators_[2];
};
#endif

View file

@ -87,6 +87,7 @@ static uint16_t rightMinX, rightMinY;
static uint16_t rightCenX, rightCenY;
static uint16_t rightMaxX, rightMaxY;
static SwitchRumbleCallback rumble_callback = nullptr;
static SwitchHapticsDecoder rumble_decoder;
static const uint8_t factory_config_data[0xEFF] = {
// serial number
@ -402,14 +403,16 @@ static void read_spi_flash(uint8_t* dest, uint32_t address, uint8_t size) {
}
}
static void forward_rumble_to_host(const uint8_t* report, uint16_t length) {
// Output reports 0x10/0x21 include 8 rumble bytes starting at offset 2.
if (!rumble_callback || length < 10) {
static void forward_decoded_rumble(const uint8_t* report, uint16_t length) {
// Output reports 0x10/0x01 include 8 rumble bytes starting at offset 2.
if (length < 10) {
return;
}
uint8_t rumble[8];
memcpy(rumble, report + 2, sizeof(rumble));
rumble_callback(rumble);
SwitchRumbleOutput rumble = rumble_decoder.decode(report + 2);
if (rumble_callback) {
rumble_callback(rumble);
}
}
static void handle_config_report(uint8_t switchReportID, uint8_t switchReportSubID, const uint8_t *reportData, uint16_t reportLength) {
@ -672,6 +675,7 @@ static void update_switch_report_from_state() {
void switch_pro_init() {
imu_mode = SwitchImuMode::Off;
rumble_decoder.reset();
reset_motion_quaternion();
player_id = 0;
last_report_counter = 0;
@ -931,52 +935,55 @@ uint16_t tud_hid_get_report_cb(uint8_t instance, uint8_t report_id, hid_report_t
return report_size;
}
void tud_hid_set_report_cb(uint8_t instance, uint8_t report_id, hid_report_type_t report_type, uint8_t const *buffer, uint16_t bufsize) {
(void)instance;
if (report_type != HID_REPORT_TYPE_OUTPUT) return;
static void process_output_report(uint8_t callback_report_id,
const uint8_t* payload,
uint16_t payload_size) {
uint8_t normalized[SWITCH_PRO_ENDPOINT_SIZE]{};
size_t normalized_size = normalize_switch_output_report(
callback_report_id, payload, payload_size, normalized);
if (normalized_size < 2) {
return;
}
memset(report_buffer, 0x00, bufsize);
memset(report_buffer, 0x00, sizeof(report_buffer));
uint8_t switchReportID = normalized[0];
uint8_t switchReportSubID = normalized[1];
LOG_PRINTF("[HID] output id=%u switchRID=0x%02x sub=0x%02x len=%u\n",
callback_report_id, switchReportID, switchReportSubID,
static_cast<unsigned>(normalized_size));
uint8_t switchReportID = buffer[0];
uint8_t switchReportSubID = buffer[1];
LOG_PRINTF("[HID] set_report type=%d id=%u switchRID=0x%02x sub=0x%02x len=%u\n",
report_type, report_id, switchReportID, switchReportSubID, bufsize);
if (switchReportID == REPORT_OUTPUT_10 || switchReportID == REPORT_OUTPUT_21) {
forward_rumble_to_host(buffer, bufsize);
if (switchReportID == REPORT_OUTPUT_10 || switchReportID == REPORT_FEATURE) {
forward_decoded_rumble(normalized, static_cast<uint16_t>(normalized_size));
}
if (switchReportID == REPORT_OUTPUT_00) {
// No-op, just acknowledge to clear any stalls.
return;
} else if (switchReportID == REPORT_FEATURE) {
queued_report_id = report_id;
handle_feature_report(switchReportID, switchReportSubID, buffer, bufsize);
}
if (switchReportID == REPORT_FEATURE) {
queued_report_id = 0;
handle_feature_report(switchReportID, switchReportSubID, normalized,
static_cast<uint16_t>(normalized_size));
} else if (switchReportID == REPORT_CONFIGURATION) {
queued_report_id = report_id;
handle_config_report(switchReportID, switchReportSubID, buffer, bufsize);
} else {
queued_report_id = 0;
handle_config_report(switchReportID, switchReportSubID, normalized,
static_cast<uint16_t>(normalized_size));
}
}
void tud_hid_report_received_cb(uint8_t instance, uint8_t report_id, uint8_t const* buffer, uint16_t bufsize) {
void tud_hid_set_report_cb(uint8_t instance, uint8_t report_id,
hid_report_type_t report_type,
const uint8_t* buffer, uint16_t bufsize) {
(void)instance;
// Host sent data on interrupt OUT; mirror the control path handling.
memset(report_buffer, 0x00, bufsize);
uint8_t switchReportID = buffer[0];
uint8_t switchReportSubID = buffer[1];
LOG_PRINTF("[HID] report_received id=%u switchRID=0x%02x sub=0x%02x len=%u\n",
report_id, switchReportID, switchReportSubID, bufsize);
if (switchReportID == REPORT_OUTPUT_10 || switchReportID == REPORT_OUTPUT_21) {
forward_rumble_to_host(buffer, bufsize);
}
if (switchReportID == REPORT_OUTPUT_00) {
if (report_type != HID_REPORT_TYPE_OUTPUT) {
return;
} else if (switchReportID == REPORT_FEATURE) {
queued_report_id = report_id;
handle_feature_report(switchReportID, switchReportSubID, buffer, bufsize);
} else if (switchReportID == REPORT_CONFIGURATION) {
queued_report_id = report_id;
handle_config_report(switchReportID, switchReportSubID, buffer, bufsize);
}
process_output_report(report_id, buffer, bufsize);
}
void tud_hid_report_received_cb(uint8_t instance, uint8_t report_id,
const uint8_t* buffer, uint16_t bufsize) {
(void)instance;
process_output_report(report_id, buffer, bufsize);
}
uint8_t const * tud_hid_descriptor_report_cb(uint8_t itf) {

View file

@ -8,6 +8,7 @@
#include <stdbool.h>
#include <stdint.h>
#include "switch_haptics.h"
#include "switch_pro_descriptors.h"
typedef struct {
@ -66,6 +67,6 @@ bool switch_pro_apply_uart_packet(const uint8_t* packet, uint8_t length, SwitchI
// Driver state helpers
bool switch_pro_is_ready();
// Optional callback fired when the host sends a rumble payload (the raw 8 rumble bytes).
typedef void (*SwitchRumbleCallback)(const uint8_t rumble_data[8]);
// Optional callback fired with decoded rumble intensities from the host.
typedef void (*SwitchRumbleCallback)(const SwitchRumbleOutput& rumble);
void switch_pro_set_rumble_callback(SwitchRumbleCallback cb);

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