Compare commits
10 commits
ae65251cb4
...
4f7e4a8de6
| Author | SHA1 | Date | |
|---|---|---|---|
| 4f7e4a8de6 | |||
| 5c18c75d33 | |||
| 3f6bf3dee2 | |||
|
e13ba506cf |
|||
|
91c619691d |
|||
|
bd973253a6 |
|||
|
22da7bce8f |
|||
|
d81e8f90c0 |
|||
|
0db04be858 |
|||
|
2604ff274b |
11 changed files with 2341 additions and 250 deletions
27
README.md
27
README.md
|
|
@ -15,10 +15,8 @@ Raspberry Pi Pico firmware that emulates a Switch Pro controller over USB and a
|
|||
5. Connect the Pico to the Switch (dock USB-A or USB-C OTG); the Switch should see it as a wired Pro Controller.
|
||||
|
||||
## Planned features
|
||||
- IMU support for motion controls (gyro/accelerometer passthrough for controllers that support it to the Switch) (WIP branch on `gyrov3`).
|
||||
|
||||
## Limitations
|
||||
- No motion controls/IMU passthrough yet (planned).
|
||||
- No NFC/amiibo/IR support.
|
||||
- Rumble is best-effort: it depends on the Switch sending rumble and SDL2 being able to drive haptics on your specific controller.
|
||||
- Requires a host computer running the bridge; the Pico is not a Bluetooth/USB host for controllers.
|
||||
|
|
@ -188,6 +186,9 @@ Options:
|
|||
- `--swap-abxy-guid GUID` (repeatable) to flip AB/XY for a specific physical controller (GUID is stable across runs).
|
||||
- `--swap-hotkey x` to pick the runtime hotkey that prompts you to toggle ABXY layout for a specific connected controller (default `x`; empty string disables).
|
||||
- `--sdl-mapping path/to/gamecontrollerdb.txt` to load extra SDL mappings (defaults to `switch_pico_bridge/controller_db/gamecontrollerdb.txt`).
|
||||
- `--debug-imu` to print raw gyroscope and accelerometer readings every ~200ms (useful for verifying sensor data and troubleshooting).
|
||||
- `--no-imu` to disable sensor reading entirely (useful for controllers without gyro, or if motion causes issues).
|
||||
- `--gyro-scale FLOAT` to adjust gyroscope sensitivity (default 1.0; reduce below 1.0 if camera rotates too fast; increase above 1.0 for more sensitivity).
|
||||
|
||||
### Runtime hotkeys
|
||||
- By default, pressing `z` in the terminal re-samples every connected controller's sticks and re-applies neutral offsets. Change/disable with `--zero-hotkey`.
|
||||
|
|
@ -228,6 +229,28 @@ with SwitchUARTClient("/dev/cu.usbserial-0001") as client:
|
|||
### Linux tips
|
||||
- You may need udev permissions for `/dev/ttyUSB*`/`/dev/ttyACM*` (add user to `dialout`/`uucp` or use `udev` rules).
|
||||
|
||||
## IMU / Motion Controls
|
||||
|
||||
The bridge supports gyroscope and accelerometer passthrough from controllers that have motion sensors (e.g. the Nintendo Switch Pro Controller, DualSense). Motion data is forwarded to the Pico which injects it into the emulated Switch Pro Controller's HID reports.
|
||||
|
||||
### Requirements
|
||||
- A controller with gyro/accelerometer support (SDL2 must be able to enable sensors on it).
|
||||
- The Switch will automatically use motion data once the controller is recognised as a Pro Controller.
|
||||
|
||||
### Gyro bias calibration
|
||||
On startup, the bridge collects the first 200 gyro readings while the controller is stationary and averages them to compute a per-axis bias (zero-rate offset). Gyro output is zeroed during this ~1 second calibration window, then bias is subtracted from all subsequent readings. Keep the controller still when starting the bridge for best results. Use `--no-gyro-bias` to skip calibration and use raw values directly.
|
||||
|
||||
### CLI flags
|
||||
- `--debug-imu`: Print raw sensor values (m/s² and rad/s) and converted Switch integer counts every ~200ms. Useful for verifying the sensor is detected and producing sensible data.
|
||||
- `--no-imu`: Disable IMU entirely. The bridge sends zero motion data to the Pico, which sends zero-filled IMU bytes to the Switch. Buttons and sticks are unaffected.
|
||||
- `--gyro-scale FLOAT` (default 1.0): Multiply all gyro values by this factor before sending. Reduce below 1.0 if the camera moves too fast; increase above 1.0 for more sensitivity.
|
||||
|
||||
### Troubleshooting
|
||||
- **Gyro not detected**: Run with `--debug-imu`. If no IMU readings appear, SDL2 cannot see sensors on your controller (may not be supported or driver issue). On Linux, the `hid-nintendo` kernel driver routes Pro Controller IMU to a separate evdev device that SDL2 cannot read; use Windows or macOS for gyro passthrough.
|
||||
- **Wild camera swinging**: Start with `--gyro-scale 0.3` and increase gradually. Ensure the controller is still during the first second of startup (bias calibration).
|
||||
- **Verifying Pico output**: Use `python tools/read_pro_imu.py --vid 0x057E --pid 0x2009` to read raw IMU bytes directly from the Pico's USB HID output and confirm non-zero values appear.
|
||||
- **SDL2 accuracy**: SDL2 (version < 2.32.7) has a known inaccuracy bug with Switch Pro Controller gyro data. Updating the SDL2 shared library to 2.32.7 or later improves accuracy.
|
||||
|
||||
## References
|
||||
- GP2040-CE (controller firmware ecosystem): https://github.com/OpenStickCommunity/GP2040-CE
|
||||
- nxbt (Switch controller research/tools): https://github.com/Brikwerk/nxbt
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ requires-python = ">=3.9"
|
|||
authors = [{ name = "Switch Pico Maintainers" }]
|
||||
dependencies = [
|
||||
"pyserial",
|
||||
"pysdl2",
|
||||
"PySDL3",
|
||||
"rich",
|
||||
]
|
||||
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -12,6 +12,7 @@ depending on SDL. It mirrors the framing in ``switch-pico.cpp``:
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import struct
|
||||
import time
|
||||
import threading
|
||||
|
|
@ -23,9 +24,28 @@ import serial
|
|||
from serial.tools import list_ports, list_ports_common
|
||||
|
||||
UART_HEADER = 0xAA
|
||||
UART_PROTOCOL_VERSION = 0x02
|
||||
RUMBLE_HEADER = 0xBB
|
||||
RUMBLE_TYPE_RUMBLE = 0x01
|
||||
UART_BAUD = 921600
|
||||
IMU_SAMPLES_PER_REPORT = 3
|
||||
|
||||
MS2_PER_G = 9.80665
|
||||
RAD_TO_DEG = 180.0 / math.pi
|
||||
ACCEL_LSB_PER_G = 4096.0
|
||||
GYRO_LSB_PER_RAD_S = 818.5
|
||||
|
||||
try:
|
||||
import sdl3 as _sdl3 # type: ignore[import-not-found]
|
||||
|
||||
_sensor_accel = getattr(_sdl3, "SDL_SENSOR_ACCEL", 1)
|
||||
_sensor_gyro = getattr(_sdl3, "SDL_SENSOR_GYRO", 2)
|
||||
except ImportError:
|
||||
_sensor_accel = 1
|
||||
_sensor_gyro = 2
|
||||
|
||||
SENSOR_ACCEL: int = _sensor_accel
|
||||
SENSOR_GYRO: int = _sensor_gyro
|
||||
|
||||
|
||||
class SwitchButton(IntFlag):
|
||||
|
|
@ -62,9 +82,9 @@ def _is_usb_serial_path(path: str) -> bool:
|
|||
"""Heuristic for USB serial path prefixes."""
|
||||
lower = path.lower()
|
||||
usb_prefixes = (
|
||||
"/dev/ttyusb", # Linux USB serial
|
||||
"/dev/ttyacm", # Linux CDC ACM
|
||||
"/dev/cu.usb", # macOS cu/tty USB adapters
|
||||
"/dev/ttyusb", # Linux USB serial
|
||||
"/dev/ttyacm", # Linux CDC ACM
|
||||
"/dev/cu.usb", # macOS cu/tty USB adapters
|
||||
"/dev/tty.usb",
|
||||
)
|
||||
if lower.startswith(usb_prefixes):
|
||||
|
|
@ -144,6 +164,7 @@ def first_serial_port(
|
|||
return None
|
||||
return ports[0]["device"]
|
||||
|
||||
|
||||
def clamp_byte(value: Union[int, float]) -> int:
|
||||
"""Clamp a numeric value to the 0-255 byte range."""
|
||||
return max(0, min(255, int(value)))
|
||||
|
|
@ -202,6 +223,21 @@ def str_to_dpad(flags: Mapping[str, bool]) -> SwitchDpad:
|
|||
return SwitchDpad.CENTER
|
||||
|
||||
|
||||
def compute_checksum(data: bytes) -> int:
|
||||
"""Compute UART checksum as sum of bytes modulo 256."""
|
||||
return sum(data) & 0xFF
|
||||
|
||||
|
||||
@dataclass
|
||||
class IMUSample:
|
||||
accel_x: int = 0
|
||||
accel_y: int = 0
|
||||
accel_z: int = 0
|
||||
gyro_x: int = 0
|
||||
gyro_y: int = 0
|
||||
gyro_z: int = 0
|
||||
|
||||
|
||||
@dataclass
|
||||
class SwitchReport:
|
||||
buttons: int = 0
|
||||
|
|
@ -210,13 +246,38 @@ class SwitchReport:
|
|||
ly: int = 128
|
||||
rx: int = 128
|
||||
ry: int = 128
|
||||
imu_samples: List[IMUSample] = field(default_factory=list)
|
||||
|
||||
def to_bytes(self) -> bytes:
|
||||
"""Serialize the report into the UART packet format."""
|
||||
return struct.pack(
|
||||
"<BHBBBBB", UART_HEADER, self.buttons & 0xFFFF, self.hat & 0xFF, self.lx, self.ly, self.rx, self.ry
|
||||
"""Serialize the report into UART v2 framed packet format."""
|
||||
count = min(len(self.imu_samples), IMU_SAMPLES_PER_REPORT)
|
||||
payload = struct.pack(
|
||||
"<HBBBBBB",
|
||||
self.buttons & 0xFFFF,
|
||||
int(self.hat) & 0xFF,
|
||||
clamp_byte(self.lx),
|
||||
clamp_byte(self.ly),
|
||||
clamp_byte(self.rx),
|
||||
clamp_byte(self.ry),
|
||||
count,
|
||||
)
|
||||
|
||||
for i in range(count):
|
||||
sample = self.imu_samples[i]
|
||||
payload += struct.pack(
|
||||
"<hhhhhh",
|
||||
max(-32768, min(32767, int(sample.accel_x))),
|
||||
max(-32768, min(32767, int(sample.accel_y))),
|
||||
max(-32768, min(32767, int(sample.accel_z))),
|
||||
max(-32768, min(32767, int(sample.gyro_x))),
|
||||
max(-32768, min(32767, int(sample.gyro_y))),
|
||||
max(-32768, min(32767, int(sample.gyro_z))),
|
||||
)
|
||||
|
||||
payload_len = len(payload)
|
||||
frame = bytes([UART_HEADER, UART_PROTOCOL_VERSION, payload_len]) + payload
|
||||
return frame + bytes([compute_checksum(frame)])
|
||||
|
||||
|
||||
class PicoUART:
|
||||
def __init__(self, port: str, baudrate: int = UART_BAUD) -> None:
|
||||
|
|
@ -267,15 +328,15 @@ class PicoUART:
|
|||
del self._buffer[:start]
|
||||
return None
|
||||
|
||||
frame = self._buffer[start:start + 11]
|
||||
checksum = sum(frame[:10]) & 0xFF
|
||||
frame = self._buffer[start : start + 11]
|
||||
checksum = compute_checksum(bytes(frame[:10]))
|
||||
|
||||
if frame[1] == RUMBLE_TYPE_RUMBLE and checksum == frame[10]:
|
||||
payload = bytes(frame[2:10])
|
||||
del self._buffer[:start + 11]
|
||||
del self._buffer[: start + 11]
|
||||
return payload
|
||||
|
||||
del self._buffer[:start + 1]
|
||||
del self._buffer[: start + 1]
|
||||
|
||||
def close(self) -> None:
|
||||
"""Close the UART connection."""
|
||||
|
|
@ -434,14 +495,20 @@ class SwitchUARTClient:
|
|||
self.state.move_right_stick(x, y)
|
||||
self.send()
|
||||
|
||||
def press_for(self, duration: float, *buttons: SwitchButton | SwitchDpad | int) -> None:
|
||||
def press_for(
|
||||
self, duration: float, *buttons: SwitchButton | SwitchDpad | int
|
||||
) -> None:
|
||||
"""Press buttons/hat for a duration, then release."""
|
||||
self.press(*buttons)
|
||||
time.sleep(max(0.0, duration))
|
||||
self.release(*buttons)
|
||||
|
||||
def move_left_stick_for(
|
||||
self, x: Union[int, float], y: Union[int, float], duration: float, neutral_after: bool = True
|
||||
self,
|
||||
x: Union[int, float],
|
||||
y: Union[int, float],
|
||||
duration: float,
|
||||
neutral_after: bool = True,
|
||||
) -> None:
|
||||
"""Move left stick for a duration, optionally returning it to neutral afterward."""
|
||||
self.move_left_stick(x, y)
|
||||
|
|
@ -451,7 +518,11 @@ class SwitchUARTClient:
|
|||
self.send()
|
||||
|
||||
def move_right_stick_for(
|
||||
self, x: Union[int, float], y: Union[int, float], duration: float, neutral_after: bool = True
|
||||
self,
|
||||
x: Union[int, float],
|
||||
y: Union[int, float],
|
||||
duration: float,
|
||||
neutral_after: bool = True,
|
||||
) -> None:
|
||||
"""Move right stick for a duration, optionally returning it to neutral afterward."""
|
||||
self.move_right_stick(x, y)
|
||||
|
|
|
|||
|
|
@ -61,11 +61,13 @@ static void on_rumble_from_switch(const uint8_t rumble[8]) {
|
|||
}
|
||||
|
||||
// Consume UART bytes and forward complete frames to the Switch Pro driver.
|
||||
static void poll_uart_frames() {
|
||||
static uint8_t buffer[8];
|
||||
static bool poll_uart_frames() {
|
||||
static uint8_t buffer[64];
|
||||
static uint8_t index = 0;
|
||||
static uint8_t expected_len = 0;
|
||||
static absolute_time_t last_byte_time = {0};
|
||||
static bool has_last_byte = false;
|
||||
bool new_data = false;
|
||||
|
||||
while (uart_is_readable(UART_ID)) {
|
||||
uint8_t byte = uart_getc(UART_ID);
|
||||
|
|
@ -73,6 +75,7 @@ static void poll_uart_frames() {
|
|||
uint64_t now = to_ms_since_boot(get_absolute_time());
|
||||
if (has_last_byte && (now - to_ms_since_boot(last_byte_time)) > 20) {
|
||||
index = 0; // stale data, restart frame
|
||||
expected_len = 0;
|
||||
}
|
||||
last_byte_time = get_absolute_time();
|
||||
has_last_byte = true;
|
||||
|
|
@ -83,11 +86,26 @@ static void poll_uart_frames() {
|
|||
}
|
||||
}
|
||||
|
||||
buffer[index++] = byte;
|
||||
if (index >= sizeof(buffer)) {
|
||||
index = 0;
|
||||
expected_len = 0;
|
||||
}
|
||||
|
||||
buffer[index++] = byte;
|
||||
if (index == 3) {
|
||||
expected_len = static_cast<uint8_t>(buffer[2] + 4u);
|
||||
if (expected_len < 12 || expected_len > sizeof(buffer)) {
|
||||
index = 0;
|
||||
expected_len = 0;
|
||||
continue;
|
||||
}
|
||||
}
|
||||
|
||||
if (expected_len > 0 && index >= expected_len) {
|
||||
SwitchInputState parsed{};
|
||||
if (switch_pro_apply_uart_packet(buffer, sizeof(buffer), &parsed)) {
|
||||
if (switch_pro_apply_uart_packet(buffer, expected_len, &parsed)) {
|
||||
g_user_state = parsed;
|
||||
new_data = true;
|
||||
LOG_PRINTF("[UART] packet buttons=0x%04x hat=%u lx=%u ly=%u rx=%u ry=%u\n",
|
||||
(parsed.button_a ? SWITCH_PRO_MASK_A : 0) |
|
||||
(parsed.button_b ? SWITCH_PRO_MASK_B : 0) |
|
||||
|
|
@ -104,14 +122,17 @@ static void poll_uart_frames() {
|
|||
(parsed.button_l3 ? SWITCH_PRO_MASK_L3 : 0) |
|
||||
(parsed.button_r3 ? SWITCH_PRO_MASK_R3 : 0),
|
||||
parsed.dpad_up ? SWITCH_PRO_HAT_UP :
|
||||
parsed.dpad_down ? SWITCH_PRO_HAT_DOWN :
|
||||
parsed.dpad_left ? SWITCH_PRO_HAT_LEFT :
|
||||
parsed.dpad_right ? SWITCH_PRO_HAT_RIGHT : SWITCH_PRO_HAT_NOTHING,
|
||||
parsed.lx >> 8, parsed.ly >> 8, parsed.rx >> 8, parsed.ry >> 8);
|
||||
parsed.dpad_down ? SWITCH_PRO_HAT_DOWN :
|
||||
parsed.dpad_left ? SWITCH_PRO_HAT_LEFT :
|
||||
parsed.dpad_right ? SWITCH_PRO_HAT_RIGHT : SWITCH_PRO_HAT_NOTHING,
|
||||
parsed.lx >> 8, parsed.ly >> 8, parsed.rx >> 8, parsed.ry >> 8);
|
||||
}
|
||||
index = 0;
|
||||
expected_len = 0;
|
||||
}
|
||||
}
|
||||
|
||||
return new_data;
|
||||
}
|
||||
|
||||
static void log_usb_state() {
|
||||
|
|
@ -146,7 +167,8 @@ int main() {
|
|||
|
||||
while (true) {
|
||||
tud_task(); // USB device tasks
|
||||
poll_uart_frames(); // Pull controller state from UART1
|
||||
bool new_data = poll_uart_frames(); // Pull controller state from UART1
|
||||
(void)new_data;
|
||||
SwitchInputState state = g_user_state;
|
||||
switch_pro_set_input(state);
|
||||
switch_pro_task(); // Push state to the Switch host
|
||||
|
|
|
|||
|
|
@ -184,12 +184,39 @@ static std::map<uint32_t, const uint8_t*> spi_flash_data = {
|
|||
|
||||
static inline uint16_t scale16To12(uint16_t pos) { return pos >> 4; }
|
||||
|
||||
static void fill_imu_report_data(const SwitchInputState& state) {
|
||||
if (state.imu_sample_count == 0) {
|
||||
memset(switch_report.imuData, 0x00, sizeof(switch_report.imuData));
|
||||
return;
|
||||
}
|
||||
uint8_t sample_count = state.imu_sample_count > 3 ? 3 : state.imu_sample_count;
|
||||
// If fewer than 3 samples, duplicate the last one to fill all 3 slots
|
||||
uint8_t* dst = switch_report.imuData;
|
||||
for (uint8_t i = 0; i < 3; ++i) {
|
||||
const SwitchImuSample& s = (i < sample_count) ? state.imu_samples[i] : state.imu_samples[sample_count - 1];
|
||||
dst[0] = static_cast<uint8_t>(s.accel_x & 0xFF);
|
||||
dst[1] = static_cast<uint8_t>((s.accel_x >> 8) & 0xFF);
|
||||
dst[2] = static_cast<uint8_t>(s.accel_y & 0xFF);
|
||||
dst[3] = static_cast<uint8_t>((s.accel_y >> 8) & 0xFF);
|
||||
dst[4] = static_cast<uint8_t>(s.accel_z & 0xFF);
|
||||
dst[5] = static_cast<uint8_t>((s.accel_z >> 8) & 0xFF);
|
||||
dst[6] = static_cast<uint8_t>(s.gyro_x & 0xFF);
|
||||
dst[7] = static_cast<uint8_t>((s.gyro_x >> 8) & 0xFF);
|
||||
dst[8] = static_cast<uint8_t>(s.gyro_y & 0xFF);
|
||||
dst[9] = static_cast<uint8_t>((s.gyro_y >> 8) & 0xFF);
|
||||
dst[10] = static_cast<uint8_t>(s.gyro_z & 0xFF);
|
||||
dst[11] = static_cast<uint8_t>((s.gyro_z >> 8) & 0xFF);
|
||||
dst += 12;
|
||||
}
|
||||
}
|
||||
|
||||
static SwitchInputState make_neutral_state() {
|
||||
SwitchInputState s{};
|
||||
s.lx = SWITCH_PRO_JOYSTICK_MID;
|
||||
s.ly = SWITCH_PRO_JOYSTICK_MID;
|
||||
s.rx = SWITCH_PRO_JOYSTICK_MID;
|
||||
s.ry = SWITCH_PRO_JOYSTICK_MID;
|
||||
s.imu_sample_count = 0;
|
||||
return s;
|
||||
}
|
||||
|
||||
|
|
@ -485,6 +512,7 @@ static void update_switch_report_from_state() {
|
|||
switch_report.inputs.rightStick.setX(std::min(std::max(scaleRightStickX,rightMinX), rightMaxX));
|
||||
switch_report.inputs.rightStick.setY(-std::min(std::max(scaleRightStickY,rightMinY), rightMaxY));
|
||||
|
||||
fill_imu_report_data(g_input_state);
|
||||
switch_report.rumbleReport = 0x09;
|
||||
}
|
||||
|
||||
|
|
@ -594,14 +622,13 @@ void switch_pro_task() {
|
|||
switch_report.timestamp = last_report_counter;
|
||||
void * inputReport = &switch_report;
|
||||
uint16_t report_size = sizeof(switch_report);
|
||||
if (memcmp(last_report, inputReport, report_size) != 0) {
|
||||
if (tud_hid_ready() && send_report(0, inputReport, report_size) == true ) {
|
||||
memcpy(last_report, inputReport, report_size);
|
||||
report_sent = true;
|
||||
}
|
||||
|
||||
last_report_timer = now;
|
||||
if (tud_hid_ready() && send_report(0, inputReport, report_size) == true ) {
|
||||
memcpy(last_report, inputReport, report_size);
|
||||
g_input_state.imu_sample_count = 0;
|
||||
report_sent = true;
|
||||
}
|
||||
|
||||
last_report_timer = now;
|
||||
}
|
||||
} else {
|
||||
if (!is_initialized) {
|
||||
|
|
@ -617,24 +644,71 @@ void switch_pro_task() {
|
|||
}
|
||||
|
||||
bool switch_pro_apply_uart_packet(const uint8_t* packet, uint8_t length, SwitchInputState* out_state) {
|
||||
// Packet format: 0xAA, buttons(2 LE), hat, lx, ly, rx, ry
|
||||
if (length < 8 || packet[0] != 0xAA) {
|
||||
// v2 format: 0xAA + 0x02 + payload_len + payload... + checksum
|
||||
if (length < 12) {
|
||||
return false;
|
||||
}
|
||||
if (packet[0] != 0xAA) {
|
||||
return false;
|
||||
}
|
||||
if (packet[1] != 0x02) {
|
||||
return false;
|
||||
}
|
||||
|
||||
uint8_t payload_len = packet[2];
|
||||
if ((uint16_t)payload_len + 4u != length) {
|
||||
return false;
|
||||
}
|
||||
|
||||
uint16_t sum = 0;
|
||||
for (uint16_t i = 0; i < (uint16_t)(3u + payload_len); ++i) {
|
||||
sum += packet[i];
|
||||
}
|
||||
if ((sum & 0xFF) != packet[length - 1]) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// payload: buttons(2 LE), hat, lx, ly, rx, ry, imu_count, [imu_samples...]
|
||||
if (payload_len < 8) {
|
||||
return false;
|
||||
}
|
||||
|
||||
SwitchProOutReport out{};
|
||||
out.buttons = static_cast<uint16_t>(packet[1]) | (static_cast<uint16_t>(packet[2]) << 8);
|
||||
out.hat = packet[3];
|
||||
out.lx = packet[4];
|
||||
out.ly = packet[5];
|
||||
out.rx = packet[6];
|
||||
out.ry = packet[7];
|
||||
out.buttons = static_cast<uint16_t>(packet[3]) | (static_cast<uint16_t>(packet[4]) << 8);
|
||||
out.hat = packet[5];
|
||||
out.lx = packet[6];
|
||||
out.ly = packet[7];
|
||||
out.rx = packet[8];
|
||||
out.ry = packet[9];
|
||||
uint8_t imu_count = packet[10];
|
||||
if (imu_count > 3) {
|
||||
imu_count = 3;
|
||||
}
|
||||
|
||||
uint16_t required_payload_len = static_cast<uint16_t>(8u + static_cast<uint16_t>(imu_count) * 12u);
|
||||
if (payload_len < required_payload_len) {
|
||||
return false;
|
||||
}
|
||||
|
||||
auto expand_axis = [](uint8_t v) -> uint16_t {
|
||||
return static_cast<uint16_t>(v) << 8 | v;
|
||||
};
|
||||
|
||||
SwitchInputState state = make_neutral_state();
|
||||
state.imu_sample_count = imu_count;
|
||||
|
||||
auto read_int16 = [](const uint8_t* src) -> int16_t {
|
||||
return static_cast<int16_t>(static_cast<uint16_t>(src[0]) | (static_cast<uint16_t>(src[1]) << 8));
|
||||
};
|
||||
for (uint8_t i = 0; i < imu_count; ++i) {
|
||||
const uint8_t* base = &packet[11 + i * 12];
|
||||
state.imu_samples[i].accel_x = read_int16(base + 0);
|
||||
state.imu_samples[i].accel_y = read_int16(base + 2);
|
||||
state.imu_samples[i].accel_z = read_int16(base + 4);
|
||||
state.imu_samples[i].gyro_x = read_int16(base + 6);
|
||||
state.imu_samples[i].gyro_y = read_int16(base + 8);
|
||||
state.imu_samples[i].gyro_z = read_int16(base + 10);
|
||||
}
|
||||
|
||||
switch (out.hat) {
|
||||
case SWITCH_PRO_HAT_UP: state.dpad_up = true; break;
|
||||
|
|
@ -668,11 +742,10 @@ bool switch_pro_apply_uart_packet(const uint8_t* packet, uint8_t length, SwitchI
|
|||
state.rx = expand_axis(out.rx);
|
||||
state.ry = expand_axis(out.ry);
|
||||
|
||||
if (out_state) {
|
||||
*out_state = state;
|
||||
} else {
|
||||
switch_pro_set_input(state);
|
||||
if (!out_state) {
|
||||
return false;
|
||||
}
|
||||
*out_state = state;
|
||||
return true;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -10,6 +10,15 @@
|
|||
#include <stdint.h>
|
||||
#include "switch_pro_descriptors.h"
|
||||
|
||||
typedef struct {
|
||||
int16_t accel_x;
|
||||
int16_t accel_y;
|
||||
int16_t accel_z;
|
||||
int16_t gyro_x;
|
||||
int16_t gyro_y;
|
||||
int16_t gyro_z;
|
||||
} SwitchImuSample;
|
||||
|
||||
typedef struct {
|
||||
bool dpad_up;
|
||||
bool dpad_down;
|
||||
|
|
@ -35,6 +44,9 @@ typedef struct {
|
|||
uint16_t ly;
|
||||
uint16_t rx;
|
||||
uint16_t ry;
|
||||
|
||||
uint8_t imu_sample_count; // 0-3
|
||||
SwitchImuSample imu_samples[3];
|
||||
} SwitchInputState;
|
||||
|
||||
// Initialize USB state and calibration before entering the main loop.
|
||||
|
|
|
|||
0
tests/__init__.py
Normal file
0
tests/__init__.py
Normal file
120
tests/test_uart_protocol.py
Normal file
120
tests/test_uart_protocol.py
Normal file
|
|
@ -0,0 +1,120 @@
|
|||
"""Tests for UART v2 protocol serialization in switch_pico_uart."""
|
||||
|
||||
import struct
|
||||
import pytest
|
||||
from switch_pico_bridge.switch_pico_uart import (
|
||||
SwitchReport,
|
||||
IMUSample,
|
||||
SwitchDpad,
|
||||
UART_HEADER,
|
||||
UART_PROTOCOL_VERSION,
|
||||
ACCEL_LSB_PER_G,
|
||||
GYRO_LSB_PER_RAD_S,
|
||||
MS2_PER_G,
|
||||
compute_checksum,
|
||||
)
|
||||
|
||||
|
||||
def test_v2_frame_with_imu_samples():
|
||||
"""V2 frame with 3 IMU samples should be 48 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) == 48, f"Expected 48 bytes, got {len(data)}"
|
||||
assert data[0] == UART_HEADER # 0xAA
|
||||
assert data[1] == UART_PROTOCOL_VERSION # 0x02
|
||||
assert data[2] == 44 # payload_len
|
||||
assert data[10] == 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]
|
||||
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]
|
||||
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."""
|
||||
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 data[0] == UART_HEADER
|
||||
assert data[1] == UART_PROTOCOL_VERSION
|
||||
assert data[2] == 8 # payload_len
|
||||
assert data[10] == 0 # imu_count
|
||||
assert data[-1] == compute_checksum(data[:-1])
|
||||
|
||||
|
||||
def test_checksum_validation():
|
||||
"""Checksum should match sum of all preceding bytes & 0xFF."""
|
||||
r = SwitchReport(buttons=0x0001)
|
||||
data = r.to_bytes()
|
||||
expected_checksum = sum(data[:-1]) & 0xFF
|
||||
assert data[-1] == expected_checksum
|
||||
# Corrupt a byte and verify mismatch
|
||||
corrupted = bytearray(data)
|
||||
corrupted[3] ^= 0xFF # flip bits in first payload byte
|
||||
recalculated = sum(corrupted[:-1]) & 0xFF
|
||||
assert corrupted[-1] != recalculated, "Checksum should not match corrupted data"
|
||||
|
||||
|
||||
def test_accel_scale_gravity():
|
||||
"""1G (9.80665 m/s²) should convert to ~4096 raw counts."""
|
||||
# convert_accel_to_raw(9.80665) ≈ 4096
|
||||
raw = int(round((MS2_PER_G / MS2_PER_G) * ACCEL_LSB_PER_G))
|
||||
assert abs(raw - 4096) <= 5, f"Expected ~4096 for 1G, got {raw}"
|
||||
|
||||
|
||||
def test_gyro_scale_one_rad():
|
||||
"""1.0 rad/s should convert to ~818 raw counts."""
|
||||
raw = int(round(1.0 * GYRO_LSB_PER_RAD_S))
|
||||
assert abs(raw - 818) <= 5, f"Expected ~818 for 1 rad/s, got {raw}"
|
||||
|
||||
|
||||
def test_imu_sample_dataclass():
|
||||
"""IMUSample fields accept int16 range values."""
|
||||
s = IMUSample(
|
||||
accel_x=32767, accel_y=-32768, accel_z=0, gyro_x=100, gyro_y=-100, gyro_z=1000
|
||||
)
|
||||
assert s.accel_x == 32767
|
||||
assert s.accel_y == -32768
|
||||
assert s.gyro_z == 1000
|
||||
# Values outside int16 range are clamped in to_bytes()
|
||||
s2 = IMUSample(accel_x=99999)
|
||||
r = SwitchReport(imu_samples=[s2])
|
||||
data = r.to_bytes()
|
||||
ax = struct.unpack_from("<h", data, 11)[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)."""
|
||||
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 buttons == 0x000A
|
||||
# lx at byte 6
|
||||
assert data[6] == 200
|
||||
|
||||
|
||||
def test_max_imu_samples_capped():
|
||||
"""Providing >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) == 48 # 3 samples, not 5
|
||||
assert data[10] == 3
|
||||
assert data[2] == 44 # payload_len for 3 samples
|
||||
199
tools/read_pro_imu.py
Executable file
199
tools/read_pro_imu.py
Executable file
|
|
@ -0,0 +1,199 @@
|
|||
#!/usr/bin/env python3
|
||||
"""
|
||||
Read raw IMU samples from a Nintendo Switch Pro Controller (or Pico spoof) over USB.
|
||||
|
||||
Uses the `hidapi` (pyhidapi) package. Press Ctrl+C to exit.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import struct
|
||||
import sys
|
||||
from typing import List, Tuple
|
||||
|
||||
DEFAULT_VENDOR_ID = 0x057E
|
||||
DEFAULT_PRODUCT_ID = 0x2009 # Switch Pro Controller (USB)
|
||||
|
||||
try:
|
||||
import hid # from pyhidapi
|
||||
except ImportError:
|
||||
hid = None
|
||||
|
||||
|
||||
def list_devices(filter_vid=None, filter_pid=None):
|
||||
devices = hid.enumerate()
|
||||
for d in devices:
|
||||
if filter_vid and d["vendor_id"] != filter_vid:
|
||||
continue
|
||||
if filter_pid and d["product_id"] != filter_pid:
|
||||
continue
|
||||
print(
|
||||
f"VID=0x{d['vendor_id']:04X} PID=0x{d['product_id']:04X} "
|
||||
f"path={d.get('path')} "
|
||||
f"serial={d.get('serial_number')} "
|
||||
f"manufacturer={d.get('manufacturer_string')} "
|
||||
f"product={d.get('product_string')} "
|
||||
f"interface={d.get('interface_number')}"
|
||||
)
|
||||
return devices
|
||||
|
||||
|
||||
def find_device(vendor_id: int, product_id: int):
|
||||
for dev in hid.enumerate():
|
||||
if dev["vendor_id"] == vendor_id and dev["product_id"] == product_id:
|
||||
return dev
|
||||
return None
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Read raw 0x30 reports (IMU) from a Switch Pro Controller / Pico."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vid",
|
||||
type=lambda x: int(x, 0),
|
||||
default=DEFAULT_VENDOR_ID,
|
||||
help="Vendor ID (default 0x057E)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--pid",
|
||||
type=lambda x: int(x, 0),
|
||||
default=DEFAULT_PRODUCT_ID,
|
||||
help="Product ID (default 0x2009)",
|
||||
)
|
||||
parser.add_argument("--path", help="Explicit HID path to open (overrides VID/PID).")
|
||||
parser.add_argument(
|
||||
"--count",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Stop after this many 0x30 reports (0 = infinite).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--timeout", type=int, default=3000, help="Read timeout ms (default 3000)."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--list", action="store_true", help="List detected HID devices and exit."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--plot",
|
||||
action="store_true",
|
||||
help="Plot accel/gyro traces after capture (requires matplotlib).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--save-prefix",
|
||||
help="If set, save accel/gyro plots as '<prefix>_accel.png' and '<prefix>_gyro.png'.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
if hid is None:
|
||||
print(
|
||||
"pyhidapi is required for this tool. Install it with: pip install pyhidapi",
|
||||
file=sys.stderr,
|
||||
)
|
||||
sys.exit(1)
|
||||
|
||||
if args.list:
|
||||
list_devices()
|
||||
return
|
||||
|
||||
if args.path:
|
||||
dev_info = {
|
||||
"path": bytes(args.path, encoding="utf-8"),
|
||||
"vendor_id": args.vid,
|
||||
"product_id": args.pid,
|
||||
}
|
||||
else:
|
||||
dev_info = find_device(args.vid, args.pid)
|
||||
if not dev_info:
|
||||
print(
|
||||
f"No HID device found for VID=0x{args.vid:04X} PID=0x{args.pid:04X}. "
|
||||
"Use --list to inspect devices or --path to target a specific one.",
|
||||
file=sys.stderr,
|
||||
)
|
||||
sys.exit(1)
|
||||
|
||||
device = hid.device()
|
||||
device.open_path(dev_info["path"])
|
||||
device.set_nonblocking(False)
|
||||
print(
|
||||
f"Reading raw 0x30 reports from device (VID=0x{args.vid:04X} PID=0x{args.pid:04X})... "
|
||||
"Ctrl+C to stop."
|
||||
)
|
||||
accel_series: List[Tuple[int, int, int]] = []
|
||||
gyro_series: List[Tuple[int, int, int]] = []
|
||||
try:
|
||||
read_count = 0
|
||||
while args.count == 0 or read_count < args.count:
|
||||
data = device.read(64, timeout_ms=args.timeout)
|
||||
if not data:
|
||||
print(f"(timeout after {args.timeout} ms, no data)")
|
||||
continue
|
||||
if data[0] != 0x30:
|
||||
print(f"(non-0x30 report id=0x{data[0]:02X}, len={len(data)})")
|
||||
continue
|
||||
samples = []
|
||||
offset = 13 # accel_x starts at byte 13
|
||||
for _ in range(3):
|
||||
ax, ay, az, gx, gy, gz = struct.unpack_from(
|
||||
"<hhhhhh", bytes(data), offset
|
||||
)
|
||||
samples.append((ax, ay, az, gx, gy, gz))
|
||||
offset += 12
|
||||
print(samples)
|
||||
accel_series.extend((s[0], s[1], s[2]) for s in samples)
|
||||
gyro_series.extend((s[3], s[4], s[5]) for s in samples)
|
||||
read_count += 1
|
||||
except KeyboardInterrupt:
|
||||
pass
|
||||
finally:
|
||||
device.close()
|
||||
|
||||
if args.plot:
|
||||
try:
|
||||
import matplotlib.pyplot as plt
|
||||
except Exception as exc: # pragma: no cover - optional dependency
|
||||
print(f"Unable to plot (matplotlib not available): {exc}", file=sys.stderr)
|
||||
return
|
||||
|
||||
if accel_series and gyro_series:
|
||||
# Each sample is a tuple of three axes; plot per axis vs sample index.
|
||||
accel_x = [s[0] for s in accel_series]
|
||||
accel_y = [s[1] for s in accel_series]
|
||||
accel_z = [s[2] for s in accel_series]
|
||||
gyro_x = [s[0] for s in gyro_series]
|
||||
gyro_y = [s[1] for s in gyro_series]
|
||||
gyro_z = [s[2] for s in gyro_series]
|
||||
|
||||
fig1, ax1 = plt.subplots()
|
||||
ax1.plot(accel_x, label="ax")
|
||||
ax1.plot(accel_y, label="ay")
|
||||
ax1.plot(accel_z, label="az")
|
||||
ax1.set_title("Accel (counts)")
|
||||
ax1.set_xlabel("Sample")
|
||||
ax1.set_ylabel("Counts")
|
||||
ax1.legend()
|
||||
|
||||
fig2, ax2 = plt.subplots()
|
||||
ax2.plot(gyro_x, label="gx")
|
||||
ax2.plot(gyro_y, label="gy")
|
||||
ax2.plot(gyro_z, label="gz")
|
||||
ax2.set_title("Gyro (counts)")
|
||||
ax2.set_xlabel("Sample")
|
||||
ax2.set_ylabel("Counts")
|
||||
ax2.legend()
|
||||
|
||||
if args.save_prefix:
|
||||
fig1.savefig(
|
||||
f"{args.save_prefix}_accel.png", dpi=150, bbox_inches="tight"
|
||||
)
|
||||
fig2.savefig(
|
||||
f"{args.save_prefix}_gyro.png", dpi=150, bbox_inches="tight"
|
||||
)
|
||||
print(
|
||||
f"Saved plots to {args.save_prefix}_accel.png and {args.save_prefix}_gyro.png"
|
||||
)
|
||||
|
||||
plt.show()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Loading…
Add table
Add a link
Reference in a new issue