From ee547c3721a0bb5b478b7956fda0448a2d955b43 Mon Sep 17 00:00:00 2001 From: Kyle Vedder Date: Sun, 20 Sep 2026 02:52:34 -0700 Subject: [PATCH] Avoid blocking and lost updates in shared CAN commands Snapshot commands before synchronous CAN I/O. Update selected motor velocities atomically so concurrent base and rail commands do not overwrite each other. --- i2rt/flow_base/flow_base_controller.py | 35 +-------- i2rt/flow_base/linear_rail_controller.py | 19 +---- .../tests/test_caster_steering_shutdown.py | 5 ++ i2rt/motor_drivers/dm_driver.py | 78 +++++++++++-------- .../tests/test_command_concurrency.py | 57 ++++++++++++++ 5 files changed, 113 insertions(+), 81 deletions(-) create mode 100644 i2rt/motor_drivers/tests/test_command_concurrency.py diff --git a/i2rt/flow_base/flow_base_controller.py b/i2rt/flow_base/flow_base_controller.py index 7846b789..9d5f8dca 100644 --- a/i2rt/flow_base/flow_base_controller.py +++ b/i2rt/flow_base/flow_base_controller.py @@ -249,38 +249,11 @@ def get_velocities(self) -> List[float]: def set_velocities(self, input_dict: Dict[str, Any]) -> None: steer_vel, drive_vel = input_dict["steer_vel"], input_dict["drive_vel"] - num_motors_in_chain = len(self.motor_interface) - num_base_motors = 2 * self.num_casters - - # Build base motor velocities (steer and drive alternating) - vels = np.zeros(num_motors_in_chain) + updates = {} for i in range(self.num_casters): - vels[i * 2] = steer_vel[i] # Steer motor - vels[i * 2 + 1] = drive_vel[i] # Drive motor - - if num_motors_in_chain > num_base_motors: - with self.motor_interface.command_lock: - current_commands = self.motor_interface.commands - if current_commands and len(current_commands) == num_motors_in_chain: - vels[num_base_motors:] = [cmd.vel for cmd in current_commands[num_base_motors:]] - elif self.homing_check_callback is not None: - try: - if self.homing_check_callback(): - logger.warning( - "Linear rail homing in progress but current_commands unavailable. " - "Linear rail velocity may be set to zero." - ) - except Exception as e: - logger.warning(f"Error checking homing status: {e}") - - self.motor_interface.set_commands( - torques=np.zeros(num_motors_in_chain), - pos=np.zeros(num_motors_in_chain), - vel=vels, - kp=np.zeros(num_motors_in_chain), - kd=2.0 * np.ones(num_motors_in_chain), - get_state=False, - ) + updates[i * 2] = steer_vel[i] + updates[i * 2 + 1] = drive_vel[i] + self.motor_interface.update_command_velocities(updates) def set_neutral(self) -> None: num_motors_in_chain = len(self.motor_interface) diff --git a/i2rt/flow_base/linear_rail_controller.py b/i2rt/flow_base/linear_rail_controller.py index 37e81ec9..0aae5921 100644 --- a/i2rt/flow_base/linear_rail_controller.py +++ b/i2rt/flow_base/linear_rail_controller.py @@ -4,8 +4,6 @@ import time from typing import Any, Dict, Literal, Optional -import numpy as np - from i2rt.motor_drivers.dm_driver import DMChainCanInterface from i2rt.motor_drivers.utils import MotorInfo from i2rt.utils.usb_gpio_driver import get_gpio_backend, is_raspberry_pi @@ -106,22 +104,7 @@ def __init__(self, motor_chain: DMChainCanInterface, target_motor_idx: int = -1) def set_velocity(self, vel: float) -> None: """Set motor velocity""" - num_motors = len(self.motor_chain) - - velocities = np.zeros(num_motors) - velocities[self.target_motor_idx] = vel - - # Preserve velocities of other motors (e.g., base motors) by reading current commands - with self.motor_chain.command_lock: - current_commands = self.motor_chain.commands - if current_commands and len(current_commands) == num_motors: - # Preserve velocities of other motors - for idx in range(num_motors): - if idx != self.target_motor_idx: - velocities[idx] = current_commands[idx].vel - - torques = np.zeros(num_motors) - self.motor_chain.set_commands(torques=torques, vel=velocities, pos=None, kp=None, kd=None, get_state=False) + self.motor_chain.update_command_velocities({self.target_motor_idx: vel}) def get_state(self) -> MotorInfo: """Get motor state""" diff --git a/i2rt/flow_base/tests/test_caster_steering_shutdown.py b/i2rt/flow_base/tests/test_caster_steering_shutdown.py index 1af106ed..d8729a1d 100644 --- a/i2rt/flow_base/tests/test_caster_steering_shutdown.py +++ b/i2rt/flow_base/tests/test_caster_steering_shutdown.py @@ -77,6 +77,11 @@ def set_commands(self, torques, pos=None, vel=None, kp=None, kd=None, get_state= self.calls.append("set_commands") return None if not get_state else self.read_states() + def update_command_velocities(self, updates: dict[int, float]) -> None: + with self._lock: + for index, velocity in updates.items(): + self._vel[index] = velocity + def close(self) -> None: self.calls.append("close") self.running = False diff --git a/i2rt/motor_drivers/dm_driver.py b/i2rt/motor_drivers/dm_driver.py index d9174672..973f0401 100644 --- a/i2rt/motor_drivers/dm_driver.py +++ b/i2rt/motor_drivers/dm_driver.py @@ -1,3 +1,4 @@ +import copy import logging import os import struct @@ -606,39 +607,41 @@ def _set_torques_and_update_state(self) -> None: max_step_time = 0.0 report_start_time = curr_time - # Update state + # Snapshot commands without holding the lock across synchronous CAN I/O. + # set_commands replaces the command list, so this reference remains stable. with self.command_lock: - try: - motor_feedback = self._set_commands(self.commands) - except RuntimeError as e: - if self.enable_auto_recovery and "Motor error detected" in str(e): - logging.warning(f"Motor error in control loop, attempting recovery: {e}") - if self._try_recover_motors(): - logging.warning("Motor recovery successful, continuing control loop") - continue - self.running = False - raise - raise - - errors = np.array([motor_feedback[i].error_code != "0x1" for i in range(len(motor_feedback))]) - if np.any(errors): - if self.enable_auto_recovery: - logging.warning(f"Motor errors detected in feedback: {errors}, attempting recovery") - if self._try_recover_motors(motor_feedback): - logging.warning("Motor recovery successful, continuing control loop") - continue + commands = self.commands + try: + motor_feedback = self._set_commands(commands) + except RuntimeError as e: + if self.enable_auto_recovery and "Motor error detected" in str(e): + logging.warning(f"Motor error in control loop, attempting recovery: {e}") + if self._try_recover_motors(): + logging.warning("Motor recovery successful, continuing control loop") + continue self.running = False - logging.error(f"motor errors: {errors}") - raise Exception(f"motor errors detected: {errors}, stopping control loop") + raise + raise + + errors = np.array([motor_feedback[i].error_code != "0x1" for i in range(len(motor_feedback))]) + if np.any(errors): + if self.enable_auto_recovery: + logging.warning(f"Motor errors detected in feedback: {errors}, attempting recovery") + if self._try_recover_motors(motor_feedback): + logging.warning("Motor recovery successful, continuing control loop") + continue + self.running = False + logging.error(f"motor errors: {errors}") + raise Exception(f"motor errors detected: {errors}, stopping control loop") with self.state_lock: self.state = motor_feedback self._update_absolute_positions(motor_feedback) if self.same_bus_device_driver is not None: time.sleep(0.001) + states = self.same_bus_device_driver.read_states() with self.same_bus_device_lock: - # assume the same bus device is a passive input device (no commands to send) for now. - self.same_bus_device_states = self.same_bus_device_driver.read_states() + self.same_bus_device_states = states time.sleep(0.0005) # yield GIL so other threads can acquire locks self._rate_recorder.track() except Exception as e: @@ -678,13 +681,14 @@ def _try_recover_motors(self, motor_feedback: Optional[List[MotorInfo]] = None, time.sleep(0.01) try: with self.command_lock: - motor_feedback = self._set_commands(self.commands) - if all(fb.error_code == "0x1" for fb in motor_feedback): - logging.warning("All motors recovered successfully") - with self.state_lock: - self.state = motor_feedback - self._update_absolute_positions(motor_feedback) - return True + commands = self.commands + motor_feedback = self._set_commands(commands) + if all(fb.error_code == "0x1" for fb in motor_feedback): + logging.warning("All motors recovered successfully") + with self.state_lock: + self.state = motor_feedback + self._update_absolute_positions(motor_feedback) + return True except RuntimeError: continue @@ -771,9 +775,19 @@ def set_commands( if get_state: return self.read_states(torques=torques) + def update_command_velocities(self, updates: Dict[int, float]) -> None: + """Atomically update selected motor velocities without replacing unrelated commands.""" + with self.command_lock: + commands = [copy.copy(command) for command in self.commands] + for idx, velocity in updates.items(): + if idx < 0 or idx >= len(commands): + raise IndexError(f"Motor command index {idx} out of range [0, {len(commands)})") + commands[idx].vel = float(velocity) + self.commands = commands + def get_same_bus_device_states(self) -> Any: with self.same_bus_device_lock: - return self.same_bus_device_states + return copy.deepcopy(self.same_bus_device_states) def close(self) -> None: self.running = False diff --git a/i2rt/motor_drivers/tests/test_command_concurrency.py b/i2rt/motor_drivers/tests/test_command_concurrency.py new file mode 100644 index 00000000..8dd82193 --- /dev/null +++ b/i2rt/motor_drivers/tests/test_command_concurrency.py @@ -0,0 +1,57 @@ +import threading +from concurrent.futures import ThreadPoolExecutor +from types import SimpleNamespace + +import numpy as np + +from i2rt.flow_base.flow_base_controller import VehicleMotorController +from i2rt.flow_base.linear_rail_controller import SingleMotorControlInterface +from i2rt.motor_drivers.dm_driver import DMChainCanInterface, MotorCmd +from i2rt.utils.utils import RateRecorder + + +def test_can_io_does_not_block_command_updates() -> None: + chain = DMChainCanInterface.__new__(DMChainCanInterface) + chain.running = True + chain.command_lock, chain.state_lock = threading.RLock(), threading.Lock() + chain.commands, chain.motor_list = [MotorCmd()], [(1, "fake")] + chain._rate_recorder, chain._report_interval = RateRecorder(), 30.0 + chain.same_bus_device_driver = None + chain._update_absolute_positions = lambda _feedback: None + with ThreadPoolExecutor(1) as pool: + + def motor_io(_commands: list[MotorCmd]) -> list[SimpleNamespace]: + pool.submit(chain.set_commands, np.ones(1), get_state=False).result(timeout=0.2) + chain.running = False + return [SimpleNamespace(error_code="0x1")] + + chain._set_commands = motor_io + chain._set_torques_and_update_state() + assert chain.commands[0].torque == 1.0 + + +def test_base_and_rail_updates_preserve_each_other_and_inflight_snapshot() -> None: + chain = DMChainCanInterface.__new__(DMChainCanInterface) + chain.command_lock = threading.RLock() + chain.commands = [MotorCmd() for _ in range(3)] + chain.motor_list = [(index, "fake") for index in range(3)] + snapshot = chain.commands + base = VehicleMotorController.__new__(VehicleMotorController) + base.num_casters, base.motor_interface = 1, chain + rail = SingleMotorControlInterface(chain, target_motor_idx=2) + barrier = threading.Barrier(2) + original = chain.set_commands + + def interleaved_full_update(*args: object, **kwargs: object) -> object: + # The old read/modify/write paths both read the old list before replacing it. + barrier.wait(timeout=1.0) + return original(*args, **kwargs) + + chain.set_commands = interleaved_full_update + with ThreadPoolExecutor(2) as pool: + first = pool.submit(base.set_velocities, {"steer_vel": [1.0], "drive_vel": [2.0]}) + second = pool.submit(rail.set_velocity, 3.0) + first.result(timeout=2.0) + second.result(timeout=2.0) + assert [command.vel for command in chain.commands] == [1.0, 2.0, 3.0] + assert [command.vel for command in snapshot] == [0.0, 0.0, 0.0]