use std::sync::{Arc, Mutex};
use std::time::Duration;
use std::{marker::PhantomData, sync::mpsc::Sender};
use anyhow::Result;
use robot_behavior::behavior::{Arm, EndPoint, FlangeSpace, JointSpace, Joints};
use robot_behavior::utils::path_generate;
use robot_behavior::{
ArmState, ArmTorqueControl, ControlObservation, ControlObserver, ControlSpace,
JointPositionControl, JointState, JointVelocityControl, LoadState, MoveTo, MoveTraj, Pose,
Robot, RobotDescription, RobotException, RobotResult, TorqueControl, behavior::*,
};
use rsbullet_core::{
BulletError, BulletResult, ControlModeArray, JointType, LoadModelFlags, PhysicsClient,
UrdfOptions,
};
use crate::RsBullet;
use crate::types::{QueuedControl, RsBulletRobotState};
type RsBulletControlObservers = Arc<Mutex<Vec<ControlObserver<RsBulletRobotState>>>>;
fn control_observers() -> RsBulletControlObservers {
Arc::new(Mutex::new(Vec::new()))
}
fn notify_control_observers(
observers: &RsBulletControlObservers,
state: &RsBulletRobotState,
duration: Duration,
) {
let mut observers = observers
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
for observer in observers.iter_mut() {
observer(state, duration);
}
}
pub struct RsBulletRobot<R> {
pub body_id: i32,
pub joint_indices: Vec<i32>,
pub joint_names: Vec<String>,
pub(crate) command_sender: Sender<QueuedControl>,
pub end_effector_link: i32,
state_cache: Arc<Mutex<RsBulletRobotState>>,
before_observers: RsBulletControlObservers,
after_observers: RsBulletControlObservers,
pub control_period: Duration,
_marker: PhantomData<R>,
}
impl<R> RsBulletRobot<R> {
pub fn enqueue<CF>(&self, control: CF) -> anyhow::Result<()>
where
CF: FnMut(&mut PhysicsClient, Duration) -> BulletResult<bool> + Send + 'static,
{
self.command_sender
.send(Box::new(control))
.map_err(|_| anyhow::anyhow!("Failed to send control command: channel closed"))?;
Ok(())
}
pub fn control_with<S, F>(&self, mut controller: F) -> RobotResult<()>
where
S: RsBulletControlSpace<R> + Send + 'static,
S::Obs: Send + 'static,
S::Command: Send + 'static,
F: FnMut(S::Obs, Duration) -> (S::Command, bool) + Send + 'static,
{
let body_id = self.body_id;
let joint_indices = self.joint_indices.clone();
let end_effector_link = self.end_effector_link;
let state_cache = self.state_cache.clone();
let before_observers = self.before_observers.clone();
let after_observers = self.after_observers.clone();
let mut initialized = false;
self.enqueue(move |client, dt| {
if !initialized {
S::initialize(client, body_id, &joint_indices)?;
initialized = true;
}
let obs = S::observe(
client,
body_id,
&joint_indices,
end_effector_link,
&state_cache,
)?;
let full_state = state_cache
.lock()
.map_err(|_| bullet_error("state cache poisoned"))?
.clone();
notify_control_observers(&before_observers, &full_state, dt);
let (command, done) = controller(obs, dt);
notify_control_observers(&after_observers, &full_state, dt);
S::apply(client, body_id, &joint_indices, command)?;
Ok(done)
})
.map_err(|error| RobotException::CommandException(error.to_string()))
}
}
pub struct RsBulletRobotBuilder<'a, R> {
pub(crate) _marker: PhantomData<R>,
pub(crate) rsbullet: &'a mut RsBullet,
pub(crate) load_file: &'static str,
pub(crate) base: Option<nalgebra::Isometry3<f64>>,
pub(crate) base_fixed: bool,
pub(crate) use_maximal_coordinates: Option<bool>,
pub(crate) scaling: Option<f64>,
pub(crate) flags: Option<LoadModelFlags>,
pub(crate) end_effector_link: Option<i32>,
pub(crate) end_effector_link_name: Option<String>,
}
impl<'a, R: RobotDescription> RsBulletRobotBuilder<'a, R> {
pub fn new(rsbullet: &'a mut RsBullet) -> Self {
RsBulletRobotBuilder {
_marker: PhantomData,
rsbullet,
load_file: R::URDF.expect("robot description must provide a URDF path"),
base: None,
base_fixed: false,
scaling: None,
flags: None,
use_maximal_coordinates: None,
end_effector_link: None,
end_effector_link_name: None,
}
}
}
impl<R> RsBulletRobotBuilder<'_, R> {
pub fn use_maximal_coordinates(mut self, use_maximal: bool) -> Self {
self.use_maximal_coordinates = Some(use_maximal);
self
}
pub fn flags(mut self, flags: LoadModelFlags) -> Self {
self.flags = Some(flags);
self
}
pub fn end_effector_link(mut self, link_index: i32) -> Self {
self.end_effector_link = Some(link_index);
self
}
pub fn end_effector_link_name(mut self, link_name: impl Into<String>) -> Self {
self.end_effector_link_name = Some(link_name.into());
self
}
}
impl<'a, R> EntityBuilder<'a> for RsBulletRobotBuilder<'a, R> {
type Entity = RsBulletRobot<R>;
fn name(self, _: String) -> Self {
self
}
fn base(mut self, base: impl Into<nalgebra::Isometry3<f64>>) -> Self {
self.base = Some(base.into());
self
}
fn base_fixed(mut self, base_fixed: bool) -> Self {
self.base_fixed = base_fixed;
self
}
fn scaling(mut self, scaling: f64) -> Self {
self.scaling = Some(scaling);
self
}
fn load(self) -> Result<RsBulletRobot<R>> {
let body_id = self.rsbullet.client_mut().load_urdf(
self.load_file,
Some(UrdfOptions {
base: self.base,
use_fixed_base: self.base_fixed,
global_scaling: self.scaling,
flags: self.flags,
use_maximal_coordinates: self.use_maximal_coordinates,
}),
)?;
let joint_count = self.rsbullet.client_mut().get_num_joints(body_id);
let mut joint_indices = Vec::new();
let mut joint_names = Vec::new();
for i in 0..joint_count {
let info = self.rsbullet.client_mut().get_joint_info(body_id, i)?;
if matches!(info.joint_type, JointType::Revolute | JointType::Prismatic) {
joint_indices.push(i);
joint_names.push(info.joint_name.clone());
}
}
let end_effector_link = if let Some(idx) = self.end_effector_link {
idx
} else if let Some(name) = &self.end_effector_link_name {
let mut selected = None;
for i in 0..joint_count {
let info = self.rsbullet.client_mut().get_joint_info(body_id, i)?;
if info.link_name == *name {
selected = Some(i);
break;
}
}
selected.unwrap_or_else(|| {
if joint_count == 0 {
-1
} else {
joint_count - 1
}
})
} else {
if joint_count == 0 {
-1
} else {
joint_count - 1
}
};
let joint_states = self
.rsbullet
.client_mut()
.get_joint_states(body_id, &joint_indices)?;
let link_state = if end_effector_link >= 0 {
Some(self.rsbullet.client_mut().get_link_state(
body_id,
end_effector_link,
true,
true,
)?)
} else {
None
};
let state_cache = Arc::new(Mutex::new(RsBulletRobotState { joint_states, link_state }));
let cache_clone = state_cache.clone();
let joint_indices_clone = joint_indices.clone();
let end_effector_link_clone = end_effector_link;
let sender = self.rsbullet.command_sender();
let robot = RsBulletRobot::<R> {
body_id,
joint_indices,
joint_names,
command_sender: sender,
end_effector_link,
state_cache,
before_observers: control_observers(),
after_observers: control_observers(),
control_period: Duration::from_secs_f64(1.0 / 240.0),
_marker: PhantomData,
};
robot.enqueue(move |client, _| {
let joint_states = client.get_joint_states(body_id, &joint_indices_clone)?;
let link_state = if end_effector_link_clone >= 0 {
Some(client.get_link_state(body_id, end_effector_link_clone, true, true)?)
} else {
None
};
let mut cache = cache_clone.lock().map_err(|_| BulletError::CommandFailed {
message: "state cache poisoned",
code: -1,
})?;
cache.joint_states = joint_states;
cache.link_state = link_state;
Ok(false)
})?;
Ok(robot)
}
}
impl<const N: usize, R> Arm<N> for RsBulletRobot<R>
where
R: Joints<N> + EndPoint,
{
fn state(&mut self) -> RobotResult<ArmState<N>> {
let cache = self
.state_cache
.lock()
.map_err(|_| RobotException::NetworkError("state cache poisoned".to_string()))?;
Ok(cache.clone().into())
}
fn set_load(&mut self, _load: LoadState) -> RobotResult<()> {
Ok(())
}
fn get_joint(&self) -> [f64; N] {
self.state_cache
.lock()
.ok()
.map(|cache| Into::<ArmState<N>>::into(cache.clone()))
.and_then(|state| state.joint.meas.q)
.unwrap_or([0.; N])
}
fn get_endpoint(&self) -> Pose {
self.state_cache
.lock()
.ok()
.map(|cache| Into::<ArmState<N>>::into(cache.clone()))
.and_then(|state| state.flange.meas.pose)
.unwrap_or_default()
}
fn with_joint_vel(self, _vel_bound: [f64; N]) -> Self {
self
}
fn with_joint_acc(self, _acc_bound: [f64; N]) -> Self {
self
}
fn with_joint_jerk(self, _jerk_bound: [f64; N]) -> Self {
self
}
fn with_torque(self, _torque_bound: [f64; N]) -> Self {
self
}
fn with_torque_dot(self, _torque_dot_bound: [f64; N]) -> Self {
self
}
fn with_cartesian_vel(self, _vel_bound: f64) -> Self {
self
}
fn with_cartesian_acc(self, _acc_bound: f64) -> Self {
self
}
fn with_cartesian_jerk(self, _jerk_bound: f64) -> Self {
self
}
fn with_rotation_vel(self, _vel_bound: f64) -> Self {
self
}
fn with_rotation_acc(self, _acc_bound: f64) -> Self {
self
}
fn with_rotation_jerk(self, _jerk_bound: f64) -> Self {
self
}
}
impl<R> Robot for RsBulletRobot<R> {
type State = RsBulletRobotState;
const CONTROL_PERIOD: f64 = 1.0 / 240.0;
fn version() -> String {
"RsBulletRobot".to_string()
}
fn read_state(&mut self) -> RobotResult<Self::State> {
self.state_cache
.lock()
.map(|cache| cache.clone())
.map_err(|_| RobotException::NetworkError("state cache poisoned".to_string()))
}
}
impl<R> ControlObservation for RsBulletRobot<R> {
fn before<H>(&mut self, observer: H) -> &mut Self
where
H: FnMut(&Self::State, Duration) + Send + 'static,
{
self.before_observers
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.push(Box::new(observer));
self
}
fn after<H>(&mut self, observer: H) -> &mut Self
where
H: FnMut(&Self::State, Duration) + Send + 'static,
{
self.after_observers
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.push(Box::new(observer));
self
}
}
impl<const N: usize, R> Joints<N> for RsBulletRobot<R>
where
R: Joints<N>,
{
const JOINT_DEFAULT: [f64; N] = R::JOINT_DEFAULT;
const JOINT_PACKED: [f64; N] = R::JOINT_PACKED;
const JOINT_MIN: [f64; N] = R::JOINT_MIN;
const JOINT_MAX: [f64; N] = R::JOINT_MAX;
const JOINT_VEL_BOUND: [f64; N] = R::JOINT_VEL_BOUND;
const JOINT_ACC_BOUND: [f64; N] = R::JOINT_ACC_BOUND;
const JOINT_JERK_BOUND: [f64; N] = R::JOINT_JERK_BOUND;
const TORQUE_BOUND: [f64; N] = R::TORQUE_BOUND;
const TORQUE_DOT_BOUND: [f64; N] = R::TORQUE_DOT_BOUND;
}
impl<R> EndPoint for RsBulletRobot<R>
where
R: EndPoint,
{
const CARTESIAN_VEL_BOUND: f64 = R::CARTESIAN_VEL_BOUND;
const CARTESIAN_ACC_BOUND: f64 = R::CARTESIAN_ACC_BOUND;
const CARTESIAN_JERK_BOUND: f64 = R::CARTESIAN_JERK_BOUND;
const ROTATION_VEL_BOUND: f64 = R::ROTATION_VEL_BOUND;
const ROTATION_ACC_BOUND: f64 = R::ROTATION_ACC_BOUND;
const ROTATION_JERK_BOUND: f64 = R::ROTATION_JERK_BOUND;
}
impl<const N: usize, R> MoveTo<JointSpace<N>> for RsBulletRobot<R>
where
R: Joints<N>,
{
fn move_to(&mut self, target: [f64; N]) -> RobotResult<()> {
let body_id = self.body_id;
let joint_indices = self.joint_indices.clone();
let state: ArmState<N> = self
.state_cache
.lock()
.map(|cache| cache.clone().into())
.map_err(|_| RobotException::NetworkError("state cache poisoned".to_string()))?;
let (path_generate, t_max) = path_generate::joint_s_curve(
&state.joint.meas.q.unwrap_or([0.; N]),
&target,
&R::JOINT_VEL_BOUND,
&R::JOINT_ACC_BOUND,
&R::JOINT_JERK_BOUND,
);
let mut duration = Duration::from_secs(0);
self.enqueue(move |client, dt| {
duration += dt;
let target = path_generate(duration);
client.set_joint_motor_control_array(
body_id,
&joint_indices[0..N],
ControlModeArray::Position(&target),
None,
)?;
Ok(duration >= t_max)
})
.map_err(Into::into)
}
}
impl<R> MoveTo<FlangeSpace> for RsBulletRobot<R>
where
R: EndPoint,
{
fn move_to(&mut self, target: Pose) -> RobotResult<()> {
let _ = target;
Err(RobotException::UnprocessableInstructionError(
"RsBullet generic wrapper cannot infer joint count for cartesian IK; use JointSpace"
.into(),
))
}
}
impl<const N: usize, R> MoveTraj<JointSpace<N>> for RsBulletRobot<R>
where
R: Joints<N>,
{
fn move_traj(&mut self, path: Vec<[f64; N]>) -> RobotResult<()> {
let body_id = self.body_id;
let joint_indices = self.joint_indices.clone();
let mut path = path.into_iter();
self.enqueue(move |client, _| match path.next() {
Some(joint) => {
client.set_joint_motor_control_array(
body_id,
&joint_indices[0..N],
ControlModeArray::Position(&joint),
None,
)?;
Ok(false)
}
None => Ok(true),
})
.map_err(Into::into)
}
fn move_path<F>(&mut self, _path: F) -> RobotResult<()>
where
F: Fn(f64) -> Option<[f64; N]>,
{
Err(RobotException::UnprocessableInstructionError(
"RsBullet does not plan continuous joint paths; use move_traj or move_waypoints".into(),
))
}
fn move_waypoints(&mut self, waypoints: Vec<[f64; N]>) -> RobotResult<()> {
<Self as MoveTraj<JointSpace<N>>>::move_traj(self, waypoints)
}
}
impl<R> MoveTraj<FlangeSpace> for RsBulletRobot<R>
where
R: EndPoint,
{
fn move_traj(&mut self, path: Vec<Pose>) -> RobotResult<()> {
let _ = path;
Err(RobotException::UnprocessableInstructionError(
"RsBullet generic wrapper cannot infer joint count for cartesian IK; use JointSpace"
.into(),
))
}
fn move_path<F>(&mut self, _path: F) -> RobotResult<()>
where
F: Fn(f64) -> Option<Pose>,
{
Err(RobotException::UnprocessableInstructionError(
"RsBullet does not plan continuous cartesian paths; use move_traj or move_waypoints"
.into(),
))
}
fn move_waypoints(&mut self, waypoints: Vec<Pose>) -> RobotResult<()> {
<Self as MoveTraj<FlangeSpace>>::move_traj(self, waypoints)
}
}
pub trait RsBulletControlSpace<R>: ControlSpace<RsBulletRobot<R>> {
fn initialize(
_client: &mut PhysicsClient,
_body_id: i32,
_joint_indices: &[i32],
) -> BulletResult<()> {
Ok(())
}
fn observe(
client: &mut PhysicsClient,
body_id: i32,
joint_indices: &[i32],
end_effector_link: i32,
state_cache: &Arc<Mutex<RsBulletRobotState>>,
) -> BulletResult<Self::Obs>;
fn apply(
client: &mut PhysicsClient,
body_id: i32,
joint_indices: &[i32],
command: Self::Command,
) -> BulletResult<()>;
}
fn bullet_error(message: &'static str) -> BulletError {
BulletError::CommandFailed { message, code: -1 }
}
fn checked_joints<const N: usize>(joint_indices: &[i32]) -> BulletResult<&[i32]> {
if joint_indices.len() < N {
return Err(bullet_error(
"robot has fewer controllable joints than requested",
));
}
Ok(&joint_indices[..N])
}
fn read_robot_state(
client: &mut PhysicsClient,
body_id: i32,
joint_indices: &[i32],
end_effector_link: i32,
state_cache: &Arc<Mutex<RsBulletRobotState>>,
) -> BulletResult<RsBulletRobotState> {
let joint_states = client.get_joint_states(body_id, joint_indices)?;
let link_state = if end_effector_link >= 0 {
Some(client.get_link_state(body_id, end_effector_link, true, true)?)
} else {
None
};
let state = RsBulletRobotState { joint_states, link_state };
let mut cache = state_cache
.lock()
.map_err(|_| bullet_error("state cache poisoned"))?;
*cache = state.clone();
Ok(state)
}
fn read_arm_state<const N: usize>(
client: &mut PhysicsClient,
body_id: i32,
joint_indices: &[i32],
end_effector_link: i32,
cache: &Arc<Mutex<RsBulletRobotState>>,
) -> BulletResult<ArmState<N>> {
checked_joints::<N>(joint_indices)?;
Ok(read_robot_state(client, body_id, joint_indices, end_effector_link, cache)?.into())
}
fn clamp_by<const N: usize>(mut values: [f64; N], limit: [f64; N]) -> [f64; N] {
for i in 0..N {
let bound = limit[i].abs();
values[i] = values[i].clamp(-bound, bound);
}
values
}
fn clamp_position<const N: usize>(
mut values: [f64; N],
lower: [f64; N],
upper: [f64; N],
) -> [f64; N] {
for i in 0..N {
values[i] = values[i].clamp(lower[i], upper[i]);
}
values
}
fn disable_default_motor<const N: usize>(
client: &mut PhysicsClient,
body_id: i32,
joint_indices: &[i32],
) -> BulletResult<()> {
let joints = checked_joints::<N>(joint_indices)?;
let zero_velocity = [0.0; N];
let zero_force = [0.0; N];
client.set_joint_motor_control_array(
body_id,
joints,
ControlModeArray::Velocity(&zero_velocity),
Some(&zero_force),
)
}
impl<R, const N: usize> RsBulletControlSpace<R> for JointPositionControl<N>
where
R: Joints<N>,
{
fn observe(
client: &mut PhysicsClient,
body_id: i32,
joint_indices: &[i32],
end_effector_link: i32,
state_cache: &Arc<Mutex<RsBulletRobotState>>,
) -> BulletResult<JointState<N>> {
Ok(read_arm_state::<N>(
client,
body_id,
joint_indices,
end_effector_link,
state_cache,
)?
.joint)
}
fn apply(
client: &mut PhysicsClient,
body_id: i32,
joint_indices: &[i32],
command: [f64; N],
) -> BulletResult<()> {
let joints = checked_joints::<N>(joint_indices)?;
let command = clamp_position(command, R::JOINT_MIN, R::JOINT_MAX);
client.set_joint_motor_control_array(
body_id,
joints,
ControlModeArray::Position(&command),
None,
)
}
}
impl<R, const N: usize> RsBulletControlSpace<R> for JointVelocityControl<N>
where
R: Joints<N>,
{
fn observe(
client: &mut PhysicsClient,
body_id: i32,
joint_indices: &[i32],
end_effector_link: i32,
state_cache: &Arc<Mutex<RsBulletRobotState>>,
) -> BulletResult<JointState<N>> {
Ok(read_arm_state::<N>(
client,
body_id,
joint_indices,
end_effector_link,
state_cache,
)?
.joint)
}
fn apply(
client: &mut PhysicsClient,
body_id: i32,
joint_indices: &[i32],
command: [f64; N],
) -> BulletResult<()> {
let joints = checked_joints::<N>(joint_indices)?;
let command = clamp_by(command, R::JOINT_VEL_BOUND);
client.set_joint_motor_control_array(
body_id,
joints,
ControlModeArray::Velocity(&command),
None,
)
}
}
impl<R, const N: usize> RsBulletControlSpace<R> for TorqueControl<N>
where
R: Joints<N>,
{
fn initialize(
client: &mut PhysicsClient,
body_id: i32,
joint_indices: &[i32],
) -> BulletResult<()> {
disable_default_motor::<N>(client, body_id, joint_indices)
}
fn observe(
client: &mut PhysicsClient,
body_id: i32,
joint_indices: &[i32],
end_effector_link: i32,
state_cache: &Arc<Mutex<RsBulletRobotState>>,
) -> BulletResult<JointState<N>> {
Ok(read_arm_state::<N>(
client,
body_id,
joint_indices,
end_effector_link,
state_cache,
)?
.joint)
}
fn apply(
client: &mut PhysicsClient,
body_id: i32,
joint_indices: &[i32],
command: [f64; N],
) -> BulletResult<()> {
let joints = checked_joints::<N>(joint_indices)?;
let command = clamp_by(command, R::TORQUE_BOUND);
client.set_joint_motor_control_array(
body_id,
joints,
ControlModeArray::Torque(&command),
None,
)
}
}
impl<R, const N: usize> RsBulletControlSpace<R> for ArmTorqueControl<N>
where
R: Joints<N>,
{
fn initialize(
client: &mut PhysicsClient,
body_id: i32,
joint_indices: &[i32],
) -> BulletResult<()> {
disable_default_motor::<N>(client, body_id, joint_indices)
}
fn observe(
client: &mut PhysicsClient,
body_id: i32,
joint_indices: &[i32],
end_effector_link: i32,
state_cache: &Arc<Mutex<RsBulletRobotState>>,
) -> BulletResult<ArmState<N>> {
read_arm_state::<N>(
client,
body_id,
joint_indices,
end_effector_link,
state_cache,
)
}
fn apply(
client: &mut PhysicsClient,
body_id: i32,
joint_indices: &[i32],
command: [f64; N],
) -> BulletResult<()> {
let joints = checked_joints::<N>(joint_indices)?;
let command = clamp_by(command, R::TORQUE_BOUND);
client.set_joint_motor_control_array(
body_id,
joints,
ControlModeArray::Torque(&command),
None,
)
}
}