Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 4 additions & 31 deletions i2rt/flow_base/flow_base_controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
19 changes: 1 addition & 18 deletions i2rt/flow_base/linear_rail_controller.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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"""
Expand Down
5 changes: 5 additions & 0 deletions i2rt/flow_base/tests/test_caster_steering_shutdown.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
78 changes: 46 additions & 32 deletions i2rt/motor_drivers/dm_driver.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import copy
import logging
import os
import struct
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand Down
57 changes: 57 additions & 0 deletions i2rt/motor_drivers/tests/test_command_concurrency.py
Original file line number Diff line number Diff line change
@@ -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]