diff --git a/src/switch_pico_bridge/controller_uart_bridge.py b/src/switch_pico_bridge/controller_uart_bridge.py index 4f13857..71550f6 100644 --- a/src/switch_pico_bridge/controller_uart_bridge.py +++ b/src/switch_pico_bridge/controller_uart_bridge.py @@ -34,6 +34,14 @@ from rich.table import Table from .switch_pico_uart import ( UART_BAUD, + MS2_PER_G, + RAD_TO_DEG, + ACCEL_LSB_PER_G, + GYRO_LSB_PER_RAD_S, + SENSOR_ACCEL, + SENSOR_GYRO, + IMU_SAMPLES_PER_REPORT, + IMUSample, PicoUART, SwitchButton, SwitchDpad, @@ -49,9 +57,11 @@ 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 -CONTROLLER_DB_URL_DEFAULT = ( - "https://raw.githubusercontent.com/mdqinc/SDL_GameControllerDB/refs/heads/master/gamecontrollerdb.txt" -) +CONTROLLER_DB_URL_DEFAULT = "https://raw.githubusercontent.com/mdqinc/SDL_GameControllerDB/refs/heads/master/gamecontrollerdb.txt" +SDL_TRUE = getattr(sdl2, "SDL_TRUE", 1) +SDL_CONTROLLERSENSORUPDATE = getattr(sdl2, "SDL_CONTROLLERSENSORUPDATE", 0x658) +GYRO_BIAS_SAMPLES = 200 +IMU_BUFFER_SIZE = 32 def parse_mapping(value: str) -> Tuple[int, str]: @@ -62,7 +72,9 @@ def parse_mapping(value: str) -> Tuple[int, str]: try: idx = int(idx_str, 10) except ValueError as exc: - raise argparse.ArgumentTypeError(f"Invalid controller index '{idx_str}'") from exc + raise argparse.ArgumentTypeError( + f"Invalid controller index '{idx_str}'" + ) from exc if not port: raise argparse.ArgumentTypeError("Serial port cannot be empty") return idx, port.strip() @@ -84,9 +96,13 @@ def download_controller_db(console: Console, destination: Path, url: str) -> boo try: destination.write_bytes(data) except Exception as exc: - console.print(f"[red]Failed to write controller database to {destination}: {exc}[/red]") + console.print( + f"[red]Failed to write controller database to {destination}: {exc}[/red]" + ) return False - console.print(f"[green]Updated controller database ({len(data)} bytes) at {destination}[/green]") + console.print( + f"[green]Updated controller database ({len(data)} bytes) at {destination}[/green]" + ) return True @@ -98,7 +114,9 @@ def parse_hotkey(value: str) -> str: if not value: return "" if len(value) != 1: - raise argparse.ArgumentTypeError("Hotkeys must be a single character (or empty to disable).") + raise argparse.ArgumentTypeError( + "Hotkeys must be a single character (or empty to disable)." + ) return value @@ -151,7 +169,9 @@ def interactive_pairing( mappings: List[Tuple[int, str]] = [] for controller_idx in controller_info: if not available: - console.print("[bold red]No more UART devices available for pairing.[/bold red]") + console.print( + "[bold red]No more UART devices available for pairing.[/bold red]" + ) break table = Table( @@ -174,7 +194,9 @@ def interactive_pairing( idx = int(selection) port = available.pop(idx) mappings.append((controller_idx, port["device"])) - console.print(f"[bold green]Paired controller {controller_idx} with {port['device']}[/bold green]") + console.print( + f"[bold green]Paired controller {controller_idx} with {port['device']}[/bold green]" + ) return mappings @@ -188,7 +210,7 @@ def apply_rumble(controller: sdl2.SDL_GameController, payload: bytes) -> float: 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 + 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 sdl2.SDL_GameControllerRumble(controller, low, high, duration) @@ -204,9 +226,18 @@ class ControllerContext: port: Optional[str] uart: Optional[PicoUART] report: SwitchReport = field(default_factory=SwitchReport) - dpad: Dict[str, bool] = field(default_factory=lambda: {"up": False, "down": False, "left": False, "right": False}) + dpad: Dict[str, bool] = field( + default_factory=lambda: { + "up": False, + "down": False, + "left": False, + "right": False, + } + ) button_state: Dict[int, bool] = field(default_factory=dict) - last_trigger_state: Dict[str, bool] = field(default_factory=lambda: {"left": False, "right": False}) + last_trigger_state: Dict[str, bool] = field( + default_factory=lambda: {"left": False, "right": False} + ) last_send: float = 0.0 last_reopen_attempt: float = 0.0 last_rumble: float = 0.0 @@ -215,6 +246,16 @@ class ControllerContext: rumble_active: bool = False axis_offsets: Dict[int, int] = field(default_factory=dict) swap_abxy: bool = False + sensors_supported: bool = False + sensors_enabled: bool = False + imu_samples: List[IMUSample] = field(default_factory=list) + last_accel: Tuple[float, float, float] = (0.0, 0.0, 0.0) + gyro_bias_x: float = 0.0 + gyro_bias_y: float = 0.0 + gyro_bias_z: float = 0.0 + gyro_bias_samples: int = 0 + gyro_bias_locked: bool = False + last_debug_imu_print: float = 0.0 def capture_stick_offsets(controller: sdl2.SDL_GameController) -> Dict[int, int]: @@ -227,7 +268,9 @@ def capture_stick_offsets(controller: sdl2.SDL_GameController) -> Dict[int, int] def format_axis_offsets(offsets: Dict[int, int]) -> str: """Return a human-friendly summary of per-axis offsets (for logging).""" - return ", ".join(f"{label}={offsets.get(axis, 0):+d}" for axis, label in STICK_AXIS_LABELS) + return ", ".join( + f"{label}={offsets.get(axis, 0):+d}" for axis, label in STICK_AXIS_LABELS + ) def calibrate_axis_value(value: int, axis: int, ctx: ControllerContext) -> int: @@ -240,10 +283,63 @@ def calibrate_axis_value(value: int, axis: int, ctx: ControllerContext) -> int: return max(-32768, min(32767, value - offset)) +def clamp_int16(value: float) -> int: + return int(max(-32768.0, min(32767.0, round(value)))) + + +def convert_accel_to_raw(m_s2: float) -> int: + return clamp_int16((m_s2 / MS2_PER_G) * ACCEL_LSB_PER_G) + + +def convert_gyro_to_raw(rad_s: float, scale: float = 1.0) -> int: + return clamp_int16(rad_s * GYRO_LSB_PER_RAD_S * scale) + + +def initialize_controller_sensors(ctx: ControllerContext, console: Console) -> None: + """Enable SDL accelerometer and gyroscope on the controller if available.""" + if not all( + hasattr(sdl2, name) + for name in ( + "SDL_GameControllerHasSensor", + "SDL_GameControllerSetSensorEnabled", + "SDL_GameControllerGetSensorData", + ) + ): + ctx.sensors_supported = False + ctx.sensors_enabled = False + return + if SENSOR_ACCEL is None or SENSOR_GYRO is None: + ctx.sensors_supported = False + return + accel_ok = bool(sdl2.SDL_GameControllerHasSensor(ctx.controller, SENSOR_ACCEL)) + gyro_ok = bool(sdl2.SDL_GameControllerHasSensor(ctx.controller, SENSOR_GYRO)) + ctx.sensors_supported = accel_ok and gyro_ok + if not ctx.sensors_supported: + console.print( + f"[yellow]Controller {ctx.controller_index} has no accelerometer/gyro sensors[/yellow]" + ) + return + accel_enabled = ( + sdl2.SDL_GameControllerSetSensorEnabled(ctx.controller, SENSOR_ACCEL, SDL_TRUE) + == 0 + ) + gyro_enabled = ( + sdl2.SDL_GameControllerSetSensorEnabled(ctx.controller, SENSOR_GYRO, SDL_TRUE) + == 0 + ) + ctx.sensors_enabled = accel_enabled and gyro_enabled + if not ctx.sensors_enabled: + console.print( + f"[yellow]Controller {ctx.controller_index} failed to enable sensors[/yellow]" + ) + + class HotkeyMonitor: """Platform-aware helper that watches for configured hotkeys without blocking the main loop.""" - def __init__(self, console: Console, key_messages: Optional[Dict[str, str]] = None) -> None: + def __init__( + self, console: Console, key_messages: Optional[Dict[str, str]] = None + ) -> None: self.console = console self._platform = os.name self._fd: Optional[int] = None @@ -272,7 +368,9 @@ class HotkeyMonitor: try: import msvcrt # type: ignore except ImportError: - self.console.print("[yellow]Hotkeys disabled: msvcrt unavailable.[/yellow]") + self.console.print( + "[yellow]Hotkeys disabled: msvcrt unavailable.[/yellow]" + ) return False self._msvcrt = msvcrt self._active = True @@ -286,9 +384,10 @@ class HotkeyMonitor: import termios import tty - self._fd = sys.stdin.fileno() - self._orig_termios = termios.tcgetattr(self._fd) - tty.setcbreak(self._fd) + fd = sys.stdin.fileno() + self._fd = fd + self._orig_termios = termios.tcgetattr(fd) + tty.setcbreak(fd) self._active = True self._started = True self._print_instructions() @@ -297,7 +396,11 @@ class HotkeyMonitor: def suspend(self) -> None: if not self._active: return - if self._platform != "nt" and self._fd is not None and self._orig_termios is not None: + if ( + self._platform != "nt" + and self._fd is not None + and self._orig_termios is not None + ): import termios termios.tcsetattr(self._fd, termios.TCSADRAIN, self._orig_termios) @@ -317,7 +420,11 @@ class HotkeyMonitor: self._active = True def stop(self) -> None: - if self._platform != "nt" and self._fd is not None and self._orig_termios is not None: + if ( + self._platform != "nt" + and self._fd is not None + and self._orig_termios is not None + ): import termios termios.tcsetattr(self._fd, termios.TCSADRAIN, self._orig_termios) @@ -340,13 +447,19 @@ class HotkeyMonitor: def _print_instructions(self) -> None: if not self._keys: return - instructions = " | ".join(f"'{key.upper()}' to {message}" for key, message in self._keys.items()) + instructions = " | ".join( + f"'{key.upper()}' to {message}" for key, message in self._keys.items() + ) self.console.print(f"[magenta]Hotkeys active: {instructions}[/magenta]") def _read_key(self) -> Optional[str]: if self._platform == "nt": - if self._msvcrt and self._msvcrt.kbhit(): - ch = self._msvcrt.getwch() + if self._msvcrt is None: + return None + kbhit = getattr(self._msvcrt, "kbhit", None) + getwch = getattr(self._msvcrt, "getwch", None) + if callable(kbhit) and callable(getwch) and kbhit(): + ch = str(getwch()) if ch == "\x03": raise KeyboardInterrupt return ch @@ -362,7 +475,11 @@ class HotkeyMonitor: return ch -def zero_context_sticks(ctx: ControllerContext, console: Optional[Console] = None, reason: str = "Zeroed stick centers") -> None: +def zero_context_sticks( + ctx: ControllerContext, + console: Optional[Console] = None, + reason: str = "Zeroed stick centers", +) -> None: """Capture and store the current stick positions for a controller.""" offsets = capture_stick_offsets(ctx.controller) ctx.axis_offsets = offsets @@ -372,7 +489,9 @@ def zero_context_sticks(ctx: ControllerContext, console: Optional[Console] = Non ) -def zero_all_context_sticks(contexts: Dict[int, ControllerContext], console: Console) -> None: +def zero_all_context_sticks( + contexts: Dict[int, ControllerContext], console: Console +) -> None: """Zero every connected controller's sticks.""" if not contexts: console.print("[yellow]No controllers available to zero right now.[/yellow]") @@ -391,10 +510,14 @@ def controller_display_name(ctx: ControllerContext) -> str: return str(name) -def toggle_abxy_for_context(ctx: ControllerContext, config: BridgeConfig, console: Console) -> None: +def toggle_abxy_for_context( + ctx: ControllerContext, config: BridgeConfig, console: Console +) -> None: """Toggle the ABXY layout for a single controller.""" if config.swap_abxy_global: - console.print("[yellow]Global --swap-abxy is enabled; disable it to use per-controller toggles.[/yellow]") + console.print( + "[yellow]Global --swap-abxy is enabled; disable it to use per-controller toggles.[/yellow]" + ) return swapped = ctx.swap_abxy action = "standard" if swapped else "swapped" @@ -412,9 +535,13 @@ def prompt_swap_abxy_controller( ) -> None: """Prompt the user to choose a controller whose ABXY layout should be toggled.""" if not contexts: - console.print("[yellow]No controllers connected to toggle ABXY layout.[/yellow]") + console.print( + "[yellow]No controllers connected to toggle ABXY layout.[/yellow]" + ) return - controllers = sorted(contexts.values(), key=lambda ctx: (ctx.controller_index, ctx.instance_id)) + controllers = sorted( + contexts.values(), key=lambda ctx: (ctx.controller_index, ctx.instance_id) + ) table = Table(title="Toggle ABXY layout for a controller") table.add_column("Choice", justify="center") table.add_column("SDL Index", justify="center") @@ -459,7 +586,9 @@ def open_controller(index: int) -> Tuple[sdl2.SDL_GameController, int, str]: """Open an SDL GameController by index and return it with instance ID and GUID string.""" controller = sdl2.SDL_GameControllerOpen(index) if not controller: - raise RuntimeError(f"Failed to open controller {index}: {sdl2.SDL_GetError().decode()}") + raise RuntimeError( + f"Failed to open controller {index}: {sdl2.SDL_GetError().decode()}" + ) joystick = sdl2.SDL_GameControllerGetJoystick(controller) instance_id = sdl2.SDL_JoystickInstanceID(joystick) guid_str = guid_string_from_joystick(joystick) @@ -501,7 +630,9 @@ def open_uart_or_warn(port: str, baud: int, console: Console) -> Optional[PicoUA def build_arg_parser() -> argparse.ArgumentParser: """Construct the CLI argument parser for the bridge.""" - parser = argparse.ArgumentParser(description="Bridge SDL2 controllers to switch-pico UART (with rumble)") + parser = argparse.ArgumentParser( + description="Bridge SDL2 controllers to switch-pico UART (with rumble)" + ) parser.add_argument( "--map", action="append", @@ -514,8 +645,16 @@ def build_arg_parser() -> argparse.ArgumentParser: nargs="+", help="Serial ports to auto-pair with controllers in ascending index order.", ) - parser.add_argument("--interactive", action="store_true", help="Launch an interactive pairing UI using Rich.") - parser.add_argument("--all-ports", action="store_true", help="Include non-USB serial ports when listing devices.") + parser.add_argument( + "--interactive", + action="store_true", + help="Launch an interactive pairing UI using Rich.", + ) + parser.add_argument( + "--all-ports", + action="store_true", + help="Include non-USB serial ports when listing devices.", + ) parser.add_argument( "--frequency", type=float, @@ -622,10 +761,28 @@ def build_arg_parser() -> argparse.ArgumentParser: default=[], help="Path to an SDL2 controller mapping database (e.g. controllerdb.txt). Repeatable.", ) + parser.add_argument( + "--debug-imu", + action="store_true", + help="Print raw IMU sensor readings and converted values for debugging.", + ) + parser.add_argument( + "--no-imu", + action="store_true", + help="Disable IMU/sensor reading entirely. Useful for controllers without gyro.", + ) + parser.add_argument( + "--gyro-scale", + type=float, + default=1.0, + help="Scale factor for gyro sensitivity (default 1.0). Reduce below 1.0 if camera moves too fast.", + ) return parser -def poll_controller_buttons(ctx: ControllerContext, button_map: Dict[int, SwitchButton]) -> None: +def poll_controller_buttons( + ctx: ControllerContext, button_map: Dict[int, SwitchButton] +) -> None: """Update button/hat state based on current SDL controller readings.""" changed = False for sdl_button, switch_bit in button_map.items(): @@ -665,6 +822,9 @@ class BridgeConfig: swap_abxy_indices: set[int] swap_abxy_ids: set[str] swap_abxy_global: bool + debug_imu: bool = False + no_imu: bool = False + gyro_scale: float = 1.0 @dataclass @@ -680,7 +840,9 @@ class PairingState: include_port_mfr: List[str] = field(default_factory=list) -def load_button_maps(console: Console, args: argparse.Namespace) -> Tuple[Dict[int, SwitchButton], Dict[int, SwitchButton], set[int]]: +def load_button_maps( + console: Console, args: argparse.Namespace +) -> Tuple[Dict[int, SwitchButton], Dict[int, SwitchButton], set[int]]: """Load SDL controller mappings and return button map variants.""" default_mapping = Path(__file__).parent / "controller_db" / "gamecontrollerdb.txt" if args.update_controller_db or not default_mapping.exists(): @@ -695,13 +857,19 @@ def load_button_maps(console: Console, args: argparse.Namespace) -> Tuple[Dict[i button_map_swapped[sdl2.SDL_CONTROLLER_BUTTON_B] = SwitchButton.A button_map_swapped[sdl2.SDL_CONTROLLER_BUTTON_X] = SwitchButton.Y button_map_swapped[sdl2.SDL_CONTROLLER_BUTTON_Y] = SwitchButton.X - swap_abxy_indices = {idx for idx in args.swap_abxy_index if idx is not None and idx >= 0} + swap_abxy_indices = { + idx for idx in args.swap_abxy_index if idx is not None and idx >= 0 + } for mapping_path in mappings_to_load: try: loaded = sdl2.SDL_GameControllerAddMappingsFromFile(mapping_path.encode()) - console.print(f"[green]Loaded {loaded} SDL mapping(s) from {mapping_path}[/green]") + console.print( + f"[green]Loaded {loaded} SDL mapping(s) from {mapping_path}[/green]" + ) except Exception as exc: - console.print(f"[red]Failed to load SDL mapping {mapping_path}: {exc}[/red]") + console.print( + f"[red]Failed to load SDL mapping {mapping_path}: {exc}[/red]" + ) return button_map_default, button_map_swapped, swap_abxy_indices @@ -710,7 +878,9 @@ def build_bridge_config(console: Console, args: argparse.Namespace) -> BridgeCon interval = 1.0 / max(args.frequency, 1.0) deadzone_raw = int(max(0.0, min(args.deadzone, 1.0)) * 32767) trigger_threshold = int(max(0.0, min(args.trigger_threshold, 1.0)) * 32767) - button_map_default, button_map_swapped, swap_abxy_indices = load_button_maps(console, args) + button_map_default, button_map_swapped, swap_abxy_indices = load_button_maps( + console, args + ) swap_abxy_guids = {g.lower() for g in args.swap_abxy_guid} return BridgeConfig( interval=interval, @@ -724,6 +894,9 @@ def build_bridge_config(console: Console, args: argparse.Namespace) -> BridgeCon swap_abxy_indices=swap_abxy_indices, swap_abxy_ids=set(swap_abxy_guids), # filled later once stable IDs are known swap_abxy_global=bool(args.swap_abxy), + debug_imu=bool(args.debug_imu), + no_imu=bool(args.no_imu), + gyro_scale=float(args.gyro_scale), ) @@ -734,7 +907,14 @@ def initialize_sdl(parser: argparse.ArgumentParser) -> None: set_hint("SDL_JOYSTICK_HIDAPI_SWITCH", "1") # Use controller button labels so Nintendo layouts (ABXY) map correctly on Linux. set_hint("SDL_GAMECONTROLLER_USE_BUTTON_LABELS", "1") - if sdl2.SDL_Init(sdl2.SDL_INIT_GAMECONTROLLER | sdl2.SDL_INIT_JOYSTICK | sdl2.SDL_INIT_EVERYTHING) != 0: + if ( + sdl2.SDL_Init( + sdl2.SDL_INIT_GAMECONTROLLER + | sdl2.SDL_INIT_JOYSTICK + | sdl2.SDL_INIT_EVERYTHING + ) + != 0 + ): parser.error(f"SDL init failed: {sdl2.SDL_GetError().decode(errors='ignore')}") @@ -752,8 +932,12 @@ def detect_controllers( if sdl2.SDL_IsGameController(index): name = sdl2.SDL_GameControllerNameForIndex(index) name_str = name.decode() if isinstance(name, bytes) else str(name) - if include_controller_name and all(substr not in name_str.lower() for substr in include_controller_name): - console.print(f"[yellow]Skipping controller {index} ({name_str}) due to name filter[/yellow]") + if include_controller_name and all( + substr not in name_str.lower() for substr in include_controller_name + ): + console.print( + f"[yellow]Skipping controller {index} ({name_str}) due to name filter[/yellow]" + ) continue console.print(f"[cyan]Detected controller {index}: {name_str}[/cyan]") controller_indices.append(index) @@ -761,14 +945,22 @@ def detect_controllers( else: name = sdl2.SDL_JoystickNameForIndex(index) name_str = name.decode() if isinstance(name, bytes) else str(name) - if include_controller_name and all(substr not in name_str.lower() for substr in include_controller_name): - console.print(f"[yellow]Skipping joystick {index} ({name_str}) due to name filter[/yellow]") + if include_controller_name and all( + substr not in name_str.lower() for substr in include_controller_name + ): + console.print( + f"[yellow]Skipping joystick {index} ({name_str}) due to name filter[/yellow]" + ) continue - console.print(f"[yellow]Found joystick {index} (not a GameController): {name_str}[/yellow]") + console.print( + f"[yellow]Found joystick {index} (not a GameController): {name_str}[/yellow]" + ) return controller_indices, controller_names -def list_controllers_with_guids(console: Console, parser: argparse.ArgumentParser) -> None: +def list_controllers_with_guids( + console: Console, parser: argparse.ArgumentParser +) -> None: """Print detected controllers with their GUID strings and exit.""" count = sdl2.SDL_NumJoysticks() if count < 0: @@ -783,10 +975,16 @@ def list_controllers_with_guids(console: Console, parser: argparse.ArgumentParse table.add_column("GUID") for idx in range(count): is_gc = sdl2.SDL_IsGameController(idx) - name = sdl2.SDL_GameControllerNameForIndex(idx) if is_gc else sdl2.SDL_JoystickNameForIndex(idx) + name = ( + sdl2.SDL_GameControllerNameForIndex(idx) + if is_gc + else sdl2.SDL_JoystickNameForIndex(idx) + ) name_str = name.decode() if isinstance(name, bytes) else str(name) guid_str = guid_string_for_device_index(idx) - table.add_row(str(idx), "GameController" if is_gc else "Joystick", name_str, guid_str) + table.add_row( + str(idx), "GameController" if is_gc else "Joystick", name_str, guid_str + ) console.print(table) @@ -825,7 +1023,9 @@ def prepare_pairing_state( elif auto_pairing_enabled: if args.ports: available_ports.extend(list(args.ports)) - console.print(f"[green]Prepared {len(available_ports)} specified UART port(s) for auto-pairing.[/green]") + console.print( + f"[green]Prepared {len(available_ports)} specified UART port(s) for auto-pairing.[/green]" + ) else: # Passive mode: grab whatever UARTs exist now, and keep looking later. discovered = discover_serial_ports( @@ -840,7 +1040,9 @@ def prepare_pairing_state( for info in discovered: console.print(f" {info['device']} ({info['description']})") else: - console.print("[yellow]No UART devices detected yet; waiting for hotplug...[/yellow]") + console.print( + "[yellow]No UART devices detected yet; waiting for hotplug...[/yellow]" + ) mapping_by_index = {index: port for index, port in mappings} return PairingState( @@ -855,7 +1057,9 @@ def prepare_pairing_state( ) -def assign_port_for_index(pairing: PairingState, idx: int, console: Console) -> Optional[str]: +def assign_port_for_index( + pairing: PairingState, idx: int, console: Console +) -> Optional[str]: """Return the UART assigned to a controller index, auto-pairing if allowed.""" if idx in pairing.mapping_by_index: return pairing.mapping_by_index[idx] @@ -877,12 +1081,21 @@ def ports_in_use(pairing: PairingState, contexts: Dict[int, ControllerContext]) return used -def handle_removed_port(path: str, pairing: PairingState, contexts: Dict[int, ControllerContext], console: Console) -> None: +def handle_removed_port( + path: str, + pairing: PairingState, + contexts: Dict[int, ControllerContext], + console: Console, +) -> None: """Clear mappings/contexts for a UART path that disappeared.""" if path in pairing.available_ports: pairing.available_ports.remove(path) - console.print(f"[yellow]UART {path} removed; dropping from available pool[/yellow]") - indices_to_clear = [idx for idx, mapped in pairing.mapping_by_index.items() if mapped == path] + console.print( + f"[yellow]UART {path} removed; dropping from available pool[/yellow]" + ) + indices_to_clear = [ + idx for idx, mapped in pairing.mapping_by_index.items() if mapped == path + ] for idx in indices_to_clear: pairing.mapping_by_index.pop(idx, None) pairing.auto_assigned_indices.discard(idx) @@ -900,10 +1113,14 @@ def handle_removed_port(path: str, pairing: PairingState, contexts: Dict[int, Co 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]") + console.print( + f"[yellow]UART {path} removed; controller {ctx.controller_index} waiting for reassignment[/yellow]" + ) -def discover_new_ports(pairing: PairingState, contexts: Dict[int, ControllerContext], console: Console) -> None: +def discover_new_ports( + pairing: PairingState, contexts: Dict[int, ControllerContext], console: Console +) -> None: """Scan for new serial ports and add unused ones to the available pool.""" if not pairing.auto_discover_ports: return @@ -927,7 +1144,9 @@ def discover_new_ports(pairing: PairingState, contexts: Dict[int, ControllerCont if path in in_use or path in pairing.available_ports: continue pairing.available_ports.append(path) - console.print(f"[green]Discovered UART {path} ({info['description']}); available for pairing.[/green]") + console.print( + f"[green]Discovered UART {path} ({info['description']}); available for pairing.[/green]" + ) def pair_waiting_contexts( @@ -975,7 +1194,9 @@ def open_initial_contexts( if index >= sdl2.SDL_NumJoysticks() or not sdl2.SDL_IsGameController(index): name = sdl2.SDL_JoystickNameForIndex(index) name_str = name.decode() if isinstance(name, bytes) else str(name) - console.print(f"[yellow]Index {index} is not a GameController ({name_str}). Trying raw open failed.[/yellow]") + console.print( + f"[yellow]Index {index} is not a GameController ({name_str}). Trying raw open failed.[/yellow]" + ) continue port = assign_port_for_index(pairing, index, console) if port is None and not pairing.auto_pairing_enabled: @@ -986,13 +1207,19 @@ def open_initial_contexts( console.print(f"[red]Failed to open controller {index}: {exc}[/red]") continue stable_id = guid - should_swap = index in config.swap_abxy_indices or stable_id in config.swap_abxy_ids + should_swap = ( + index in config.swap_abxy_indices or stable_id in config.swap_abxy_ids + ) uart = open_uart_or_warn(port, args.baud, console) if port else None if uart: uarts.append(uart) - console.print(f"[green]Controller {index} (id {stable_id}, inst {instance_id}) paired to {port}[/green]") + console.print( + f"[green]Controller {index} (id {stable_id}, inst {instance_id}) paired to {port}[/green]" + ) elif port: - console.print(f"[yellow]Controller {index} (id {stable_id}, inst {instance_id}) waiting for UART {port}[/yellow]") + console.print( + f"[yellow]Controller {index} (id {stable_id}, inst {instance_id}) waiting for UART {port}[/yellow]" + ) else: console.print( f"[yellow]Controller {index} (id {stable_id}, inst {instance_id}) connected; waiting for an available UART[/yellow]" @@ -1006,13 +1233,17 @@ def open_initial_contexts( uart=uart, swap_abxy=should_swap, ) + if not config.no_imu: + initialize_controller_sensors(ctx, console) if config.zero_sticks: zero_context_sticks(ctx, console) contexts[instance_id] = ctx return contexts, uarts -def handle_axis_motion(event: sdl2.SDL_Event, contexts: Dict[int, ControllerContext], config: BridgeConfig) -> None: +def handle_axis_motion( + event: sdl2.SDL_Event, contexts: Dict[int, ControllerContext], config: BridgeConfig +) -> None: """Process axis motion event into stick/trigger state.""" ctx = contexts.get(event.caxis.which) if not ctx: @@ -1045,6 +1276,79 @@ def handle_axis_motion(event: sdl2.SDL_Event, contexts: Dict[int, ControllerCont ctx.last_trigger_state["right"] = pressed +def handle_sensor_update( + event: sdl2.SDL_Event, + contexts: Dict[int, ControllerContext], + config: BridgeConfig, +) -> None: + """Process SDL sensor update event into IMU samples.""" + ctx = contexts.get(event.csensor.which) + if ctx is None or not ctx.sensors_enabled or config.no_imu: + return + + sensor_type = event.csensor.sensor + data = event.csensor.data + + if sensor_type == SENSOR_ACCEL: + ctx.last_accel = (float(data[0]), float(data[1]), float(data[2])) + return + + if sensor_type != SENSOR_GYRO: + return + + gx, gy, gz = float(data[0]), float(data[1]), float(data[2]) + + if not ctx.gyro_bias_locked: + if ctx.gyro_bias_samples < GYRO_BIAS_SAMPLES: + ctx.gyro_bias_x += gx + ctx.gyro_bias_y += gy + ctx.gyro_bias_z += gz + ctx.gyro_bias_samples += 1 + if ctx.gyro_bias_samples >= GYRO_BIAS_SAMPLES: + n = ctx.gyro_bias_samples + ctx.gyro_bias_x /= n + ctx.gyro_bias_y /= n + ctx.gyro_bias_z /= n + ctx.gyro_bias_locked = True + + if not ctx.gyro_bias_locked: + bx, by, bz = 0.0, 0.0, 0.0 + else: + bx, by, bz = ctx.gyro_bias_x, ctx.gyro_bias_y, ctx.gyro_bias_z + + ux, uy, uz = gx, gy, gz + ux -= bx + uy -= by + uz -= bz + + ax, ay, az = ctx.last_accel + sample = IMUSample( + accel_x=convert_accel_to_raw(ax), + accel_y=convert_accel_to_raw(ay), + accel_z=convert_accel_to_raw(az), + gyro_x=convert_gyro_to_raw(ux, config.gyro_scale), + gyro_y=convert_gyro_to_raw(uy, config.gyro_scale), + gyro_z=convert_gyro_to_raw(uz, config.gyro_scale), + ) + + ctx.imu_samples.append(sample) + if len(ctx.imu_samples) > IMU_BUFFER_SIZE: + ctx.imu_samples = ctx.imu_samples[-IMU_BUFFER_SIZE:] + + if config.debug_imu: + now = time.monotonic() + if now - ctx.last_debug_imu_print > 0.2: + ctx.last_debug_imu_print = now + print( + f"[IMU idx={ctx.controller_index}] " + f"accel_m_s2=({ax:.3f},{ay:.3f},{az:.3f}) " + f"gyro_rad_s=({gx:.3f},{gy:.3f},{gz:.3f}) " + f"bias_rad_s=({bx:.4f},{by:.4f},{bz:.4f}) " + f"raw=({sample.accel_x},{sample.accel_y},{sample.accel_z};" + f"{sample.gyro_x},{sample.gyro_y},{sample.gyro_z})" + ) + + def handle_button_event( event: sdl2.SDL_Event, config: BridgeConfig, @@ -1093,7 +1397,9 @@ def handle_device_added( if idx >= sdl2.SDL_NumJoysticks() or not sdl2.SDL_IsGameController(idx): name = sdl2.SDL_JoystickNameForIndex(idx) name_str = name.decode() if isinstance(name, bytes) else str(name) - console.print(f"[yellow]Index {idx} is not a GameController ({name_str}). Trying raw open failed.[/yellow]") + console.print( + f"[yellow]Index {idx} is not a GameController ({name_str}). Trying raw open failed.[/yellow]" + ) return try: controller, instance_id, guid = open_controller(idx) @@ -1106,9 +1412,13 @@ def handle_device_added( uart = open_uart_or_warn(port, args.baud, console) if port else None if uart: uarts.append(uart) - console.print(f"[green]Controller {idx} (id {stable_id}, inst {instance_id}) paired to {port}[/green]") + console.print( + f"[green]Controller {idx} (id {stable_id}, inst {instance_id}) paired to {port}[/green]" + ) elif port: - console.print(f"[yellow]Controller {idx} (id {stable_id}, inst {instance_id}) waiting for UART {port}[/yellow]") + console.print( + f"[yellow]Controller {idx} (id {stable_id}, inst {instance_id}) waiting for UART {port}[/yellow]" + ) else: console.print( f"[yellow]Controller {idx} (id {stable_id}, inst {instance_id}) connected; waiting for an available UART[/yellow]" @@ -1122,6 +1432,8 @@ def handle_device_added( uart=uart, swap_abxy=should_swap, ) + if not config.no_imu: + initialize_controller_sensors(ctx, console) if config.zero_sticks: zero_context_sticks(ctx, console) contexts[instance_id] = ctx @@ -1138,7 +1450,9 @@ def handle_device_removed( ctx = contexts.pop(instance_id, None) if not ctx: return - console.print(f"[yellow]Controller {instance_id} (id {ctx.stable_id}) removed[/yellow]") + console.print( + f"[yellow]Controller {instance_id} (id {ctx.stable_id}) removed[/yellow]" + ) if ctx.controller_index in pairing.auto_assigned_indices: # Return auto-paired UART back to the pool so a future device can use it. freed = pairing.mapping_by_index.pop(ctx.controller_index, None) @@ -1171,12 +1485,23 @@ def service_contexts( uart = open_uart_or_warn(ctx.port, args.baud, console) if uart: uarts.append(uart) - console.print(f"[green]Reconnected UART {ctx.port} for controller {ctx.controller_index}[/green]") + console.print( + f"[green]Reconnected UART {ctx.port} for controller {ctx.controller_index}[/green]" + ) ctx.uart = uart if ctx.uart is None: continue try: if now - ctx.last_send >= config.interval: + if ctx.sensors_enabled and not config.no_imu: + count = min(len(ctx.imu_samples), IMU_SAMPLES_PER_REPORT) + if count > 0: + ctx.report.imu_samples = ctx.imu_samples[:count] + ctx.imu_samples = ctx.imu_samples[count:] + else: + ctx.report.imu_samples = [] + else: + ctx.report.imu_samples = [] ctx.uart.send_report(ctx.report) ctx.last_send = now @@ -1199,7 +1524,10 @@ def service_contexts( sdl2.SDL_GameControllerRumble(ctx.controller, 0, 0, 0) ctx.rumble_active = False ctx.last_rumble_energy = 0.0 - elif ctx.rumble_active and (now - ctx.last_rumble_change) > RUMBLE_STUCK_TIMEOUT: + elif ( + ctx.rumble_active + and (now - ctx.last_rumble_change) > RUMBLE_STUCK_TIMEOUT + ): sdl2.SDL_GameControllerRumble(ctx.controller, 0, 0, 0) ctx.rumble_active = False ctx.last_rumble_energy = 0.0 @@ -1240,10 +1568,17 @@ def run_bridge_loop( break if event.type == sdl2.SDL_CONTROLLERAXISMOTION: handle_axis_motion(event, contexts, config) - elif event.type in (sdl2.SDL_CONTROLLERBUTTONDOWN, sdl2.SDL_CONTROLLERBUTTONUP): + elif event.type in ( + sdl2.SDL_CONTROLLERBUTTONDOWN, + sdl2.SDL_CONTROLLERBUTTONUP, + ): handle_button_event(event, config, contexts) + elif event.type == SDL_CONTROLLERSENSORUPDATE: + handle_sensor_update(event, contexts, config) elif event.type == sdl2.SDL_CONTROLLERDEVICEADDED: - handle_device_added(event, args, pairing, contexts, uarts, console, config) + handle_device_added( + event, args, pairing, contexts, uarts, console, config + ) elif event.type == sdl2.SDL_CONTROLLERDEVICEREMOVED: handle_device_removed(event, pairing, contexts, console) @@ -1289,7 +1624,9 @@ def main() -> None: list_controllers_with_guids(console, parser) return controller_indices, controller_names = detect_controllers(console, args, parser) - pairing = prepare_pairing_state(args, console, parser, controller_indices, controller_names) + pairing = prepare_pairing_state( + args, console, parser, controller_indices, controller_names + ) hotkey_messages: Dict[str, str] = {} if config.zero_hotkey: hotkey_messages[config.zero_hotkey] = "re-zero controller sticks" @@ -1299,14 +1636,20 @@ def main() -> None: hotkey_messages[config.swap_hotkey] + "; toggle ABXY layout" ) else: - hotkey_messages[config.swap_hotkey] = "toggle ABXY layout for a controller" + hotkey_messages[config.swap_hotkey] = ( + "toggle ABXY layout for a controller" + ) if hotkey_messages: candidate = HotkeyMonitor(console, hotkey_messages) if candidate.start(): hotkey_monitor = candidate - contexts, uarts = open_initial_contexts(args, pairing, controller_indices, console, config) + contexts, uarts = open_initial_contexts( + args, pairing, controller_indices, console, config + ) if not contexts: - console.print("[yellow]No controllers opened; waiting for hotplug events...[/yellow]") + console.print( + "[yellow]No controllers opened; waiting for hotplug events...[/yellow]" + ) run_bridge_loop(args, console, config, pairing, contexts, uarts, hotkey_monitor) finally: if hotkey_monitor: