use bevy_ecs::entity::Entity;
use bevy_ecs::prelude::{Component, World};
use nalgebra::{DMatrix, Matrix3, Rotation3, UnitQuaternion, Vector3};
use serde::{Deserialize, Serialize};
use super::joints::{swing_twist, JointLink};
use super::{Body, Joint, JointKind, JointMotor, JointSpring};
use crate::runtime::{sim_math, RigidBody, RigidBodyKind};
#[derive(
Component, Clone, Copy, Debug, Default, PartialEq, Serialize, Deserialize,
)]
pub struct Articulation {}
const ARMATURE: f32 = 1e-4;
const LIMIT_BIAS: f32 = 0.2;
#[derive(Clone, Copy, PartialEq)]
enum Dof {
Root,
Fixed,
Revolute,
Prismatic,
Spherical,
}
impl Dof {
fn count(self) -> usize {
match self {
Self::Revolute | Self::Prismatic => 1,
Self::Spherical => 3,
Self::Root | Self::Fixed => 0,
}
}
}
struct Link {
body: usize,
parent: usize,
joint: Joint,
kind: Dof,
dof: usize,
q: f32,
rel: UnitQuaternion<f32>,
columns: Vec<(usize, Vector3<f32>, Vector3<f32>)>,
gravity_scale: f32,
}
#[derive(Clone, Copy)]
struct Axis {
rotation: bool,
axis: Vector3<f32>,
point: Vector3<f32>,
}
struct Row {
dof: usize,
mass: f32,
bias: f32,
softness: f32,
lower: f32,
upper: f32,
impulse: f32,
}
struct Tree {
links: Vec<Link>,
floating: bool,
axes: Vec<Axis>,
velocity: Vec<f32>,
inverse: DMatrix<f32>,
rows: Vec<Row>,
}
#[derive(Default)]
pub(super) struct Articulations {
trees: Vec<Tree>,
of_body: Vec<Option<(usize, usize)>>,
}
fn euler(angles: [f32; 3]) -> Rotation3<f32> {
sim_math::rotation_from_euler(angles[0], angles[1], angles[2])
}
fn parent_frame(body: &Body, joint: &Joint) -> (Rotation3<f32>, Vector3<f32>) {
(
body.rotation * euler(joint.target_frame),
body.position + body.rotation * Vector3::from(joint.target_anchor),
)
}
fn child_frame(body: &Body, joint: &Joint) -> (Rotation3<f32>, Vector3<f32>) {
(
body.rotation * euler(joint.frame),
body.position + body.rotation * Vector3::from(joint.anchor),
)
}
pub(super) fn build(
world: &mut World,
bodies: &mut [Body],
links: &[JointLink],
) -> (Articulations, Vec<bool>) {
let mut taken = vec![false; links.len()];
let mut articulations = Articulations {
trees: Vec::new(),
of_body: vec![None; bodies.len()],
};
let find = |entity: Entity| {
bodies
.binary_search_by_key(&entity, |body| body.entity)
.ok()
};
let mut roots: Vec<_> = world
.query::<(Entity, &Articulation)>()
.iter(world)
.filter_map(|(entity, _)| find(entity))
.collect();
roots.sort_unstable();
let mut children = vec![Vec::new(); bodies.len()];
for (index, link) in links.iter().enumerate() {
if let Some(a) = link.a {
children[a].push(index);
}
}
let mut used = vec![false; bodies.len()];
for &root in &roots {
used[root] = true;
}
for root in roots {
let floating = bodies[root].kind == RigidBodyKind::Dynamic
&& bodies[root].inverse_mass > 0.0;
let mut tree = Tree {
links: vec![Link {
body: root,
parent: 0,
joint: Joint::default(),
kind: Dof::Root,
dof: 0,
q: 0.0,
rel: UnitQuaternion::identity(),
columns: Vec::new(),
gravity_scale: 1.0,
}],
floating,
axes: Vec::new(),
velocity: Vec::new(),
inverse: DMatrix::zeros(0, 0),
rows: Vec::new(),
};
let mut dofs = if floating { 6 } else { 0 };
let mut next = 0;
while next < tree.links.len() {
let parent = tree.links[next].body;
for &index in &children[parent] {
let link = &links[index];
let kind = match link.joint.kind {
JointKind::Fixed => Dof::Fixed,
JointKind::Hinge { .. } => Dof::Revolute,
JointKind::Slider { .. } => Dof::Prismatic,
JointKind::BallSocket => Dof::Spherical,
_ => continue,
};
let body = &bodies[link.b];
if used[link.b]
|| body.kind != RigidBodyKind::Dynamic
|| body.inverse_mass == 0.0
{
continue;
}
used[link.b] = true;
taken[index] = true;
tree.links.push(Link {
body: link.b,
parent: next,
joint: link.joint,
kind,
dof: dofs,
q: 0.0,
rel: UnitQuaternion::identity(),
columns: Vec::new(),
gravity_scale: 1.0,
});
dofs += match kind {
Dof::Revolute | Dof::Prismatic => 1,
Dof::Spherical => 3,
_ => 0,
};
}
next += 1;
}
if tree.links.len() == 1 {
continue;
}
tree.velocity = vec![0.0; dofs];
tree.axes = vec![
Axis {
rotation: false,
axis: Vector3::zeros(),
point: Vector3::zeros(),
};
dofs
];
let number = articulations.trees.len();
for (index, link) in tree.links.iter_mut().enumerate() {
let body = &mut bodies[link.body];
if index > 0 || floating {
articulations.of_body[link.body] = Some((number, index));
body.articulated = true;
body.asleep = false;
}
link.gravity_scale = world
.get::<RigidBody>(body.entity)
.map_or(1.0, |rigid| rigid.gravity_scale);
}
tree.read(bodies);
tree.place(bodies);
articulations.trees.push(tree);
}
(articulations, taken)
}
impl Tree {
fn read(&mut self, bodies: &[Body]) {
if self.floating {
let root = &bodies[self.links[0].body];
self.velocity[..3]
.copy_from_slice(root.angular_velocity.as_slice());
self.velocity[3..6].copy_from_slice(root.velocity.as_slice());
}
for index in 1..self.links.len() {
let link = &self.links[index];
let (parent, child) =
(&bodies[self.links[link.parent].body], &bodies[link.body]);
let (frame_a, point_a) = parent_frame(parent, &link.joint);
let (frame_b, point_b) = child_frame(child, &link.joint);
let relative = UnitQuaternion::from_rotation_matrix(
&(frame_a.inverse() * frame_b),
);
let axis = frame_a * Vector3::x();
let spin = child.angular_velocity - parent.angular_velocity;
let (kind, dof) = (link.kind, link.dof);
let link = &mut self.links[index];
match kind {
Dof::Revolute => {
link.q = swing_twist(&relative).0;
self.velocity[dof] = axis.dot(&spin);
}
Dof::Prismatic => {
link.q = axis.dot(&(point_b - point_a));
let carried = parent.velocity
+ parent
.angular_velocity
.cross(&(child.position - parent.position));
self.velocity[dof] = axis.dot(&(child.velocity - carried));
}
Dof::Spherical => {
link.rel = relative;
let local = frame_a.inverse() * spin;
self.velocity[dof..dof + 3]
.copy_from_slice(local.as_slice());
}
Dof::Root | Dof::Fixed => {}
}
}
}
fn place(&mut self, bodies: &mut [Body]) {
let root = &bodies[self.links[0].body];
let origin = root.position;
if self.floating {
for k in 0..3 {
let axis = Vector3::ith(k, 1.0);
self.axes[k] = Axis {
rotation: true,
axis,
point: origin,
};
self.axes[k + 3] = Axis {
rotation: false,
axis,
point: origin,
};
}
}
for index in 1..self.links.len() {
let link = &self.links[index];
let joint = link.joint;
let (frame_a, point_a) =
parent_frame(&bodies[self.links[link.parent].body], &joint);
let axis = frame_a * Vector3::x();
let (relative, point_b) = match link.kind {
Dof::Revolute => (
Rotation3::from_axis_angle(&Vector3::x_axis(), link.q),
point_a,
),
Dof::Prismatic => {
(Rotation3::identity(), point_a + axis * link.q)
}
Dof::Spherical => (link.rel.to_rotation_matrix(), point_a),
Dof::Root | Dof::Fixed => (Rotation3::identity(), point_a),
};
let rotation = frame_a * relative * euler(joint.frame).inverse();
let body = &mut bodies[link.body];
body.rotation = rotation;
body.position = point_b - rotation * Vector3::from(joint.anchor);
let dof = link.dof;
match link.kind {
Dof::Revolute | Dof::Prismatic => {
self.axes[dof] = Axis {
rotation: link.kind == Dof::Revolute,
axis,
point: point_a,
};
}
Dof::Spherical => {
for k in 0..3 {
self.axes[dof + k] = Axis {
rotation: true,
axis: frame_a * Vector3::ith(k, 1.0),
point: point_a,
};
}
}
Dof::Root | Dof::Fixed => {}
}
}
for index in 0..self.links.len() {
let mut dofs: Vec<usize> = Vec::new();
let mut at = index;
while at != 0 {
let link = &self.links[at];
let count = match link.kind {
Dof::Revolute | Dof::Prismatic => 1,
Dof::Spherical => 3,
_ => 0,
};
dofs.extend((link.dof..link.dof + count).rev());
at = link.parent;
}
if self.floating {
dofs.extend((0..6).rev());
}
dofs.reverse();
let position = bodies[self.links[index].body].position;
self.links[index].columns = dofs
.into_iter()
.map(|k| {
let axis = self.axes[k];
if axis.rotation {
(
k,
axis.axis,
axis.axis.cross(&(position - axis.point)),
)
} else {
(k, Vector3::zeros(), axis.axis)
}
})
.collect();
}
self.refresh(bodies);
}
fn refresh(&self, bodies: &mut [Body]) {
let root = &bodies[self.links[0].body];
let (spin, velocity, origin) = if self.floating {
(Vector3::zeros(), Vector3::zeros(), root.position)
} else {
(root.angular_velocity, root.velocity, root.position)
};
let start = usize::from(!self.floating);
for link in &self.links[start..] {
let body = &mut bodies[link.body];
let mut angular = spin;
let mut linear = velocity + spin.cross(&(body.position - origin));
for (k, column_angular, column_linear) in &link.columns {
angular += column_angular * self.velocity[*k];
linear += column_linear * self.velocity[*k];
}
body.angular_velocity = angular;
body.velocity = linear;
}
}
fn inertia(body: &Body) -> (f32, Matrix3<f32>) {
let rotation = body.rotation.matrix();
let moments = body.inverse_inertia.map(|inverse| {
if inverse > 0.0 {
1.0 / inverse
} else {
0.0
}
});
(
1.0 / body.inverse_mass,
rotation * Matrix3::from_diagonal(&moments) * rotation.transpose(),
)
}
fn generalized(
&self,
index: usize,
torque: &Vector3<f32>,
force: &Vector3<f32>,
into: &mut [f32],
) {
for (k, angular, linear) in &self.links[index].columns {
into[*k] += angular.dot(torque) + linear.dot(force);
}
}
fn forces(&self, bodies: &[Body], gravity: Vector3<f32>) -> Vec<f32> {
let mut force = vec![0.0; self.velocity.len()];
let mut accelerations =
vec![
(Vector3::<f32>::zeros(), Vector3::<f32>::zeros());
self.links.len()
];
for index in 1..self.links.len() {
let link = &self.links[index];
let (parent, child) =
(&bodies[self.links[link.parent].body], &bodies[link.body]);
let (parent_angular, parent_linear) = accelerations[link.parent];
let (mut spin, mut slide) = (Vector3::zeros(), Vector3::zeros());
let mut point = child.position;
for k in link.dof..link.dof + link.kind.count() {
let axis = self.axes[k];
if axis.rotation {
spin += axis.axis * self.velocity[k];
point = axis.point;
} else {
slide += axis.axis * self.velocity[k];
}
}
let w = parent.angular_velocity;
let point_velocity =
parent.velocity + w.cross(&(point - parent.position));
let angular = parent_angular + w.cross(&spin);
let linear = parent_linear
+ parent_angular.cross(&(child.position - parent.position))
+ w.cross(&(child.velocity - parent.velocity))
+ w.cross(&slide)
+ w.cross(&spin).cross(&(child.position - point))
+ spin.cross(&(child.velocity - point_velocity));
accelerations[index] = (angular, linear);
}
for index in usize::from(!self.floating)..self.links.len() {
let body = &bodies[self.links[index].body];
let (link_mass, inertia) = Self::inertia(body);
let (angular, linear) = accelerations[index];
let spin = body.angular_velocity;
let torque = -(inertia * angular + spin.cross(&(inertia * spin)));
let push = (gravity * self.links[index].gravity_scale - linear)
* link_mass;
self.generalized(index, &torque, &push, &mut force);
}
force
}
fn free_motion(
&mut self,
bodies: &mut [Body],
gravity: Vector3<f32>,
dt: f32,
) {
let dofs = self.velocity.len();
let mut mass = DMatrix::<f32>::identity(dofs, dofs) * ARMATURE;
for link in &self.links[usize::from(!self.floating)..] {
let (link_mass, inertia) = Self::inertia(&bodies[link.body]);
for (row, (k, angular_k, linear_k)) in
link.columns.iter().enumerate()
{
for (l, angular_l, linear_l) in &link.columns[row..] {
let value = angular_k.dot(&(inertia * angular_l))
+ link_mass * linear_k.dot(linear_l);
mass[(*k, *l)] += value;
if k != l {
mass[(*l, *k)] += value;
}
}
}
}
self.inverse = mass.cholesky().map_or_else(
|| DMatrix::zeros(dofs, dofs),
|cholesky| cholesky.inverse(),
);
let start = self.velocity.clone();
for share in [0.5, 1.0] {
let force = self.forces(bodies, gravity);
self.velocity.copy_from_slice(&start);
self.change_speed(&force, share * dt);
self.refresh(bodies);
}
}
fn build_rows(&mut self, dt: f32) {
self.rows.clear();
for link in &self.links {
let (limit, spring, motor): (
Option<[f32; 2]>,
Option<JointSpring>,
Option<JointMotor>,
) = match link.joint.kind {
JointKind::Hinge {
limit,
spring,
motor,
}
| JointKind::Slider {
limit,
spring,
motor,
} if matches!(link.kind, Dof::Revolute | Dof::Prismatic) => {
(limit, spring, motor)
}
_ => continue,
};
let (dof, value) = (link.dof, link.q);
let diagonal = self.inverse[(dof, dof)];
let mut push =
|bias: f32, softness: f32, lower: f32, upper: f32| {
if diagonal + softness > 0.0 {
self.rows.push(Row {
dof,
mass: 1.0 / (diagonal + softness),
bias,
softness,
lower,
upper,
impulse: 0.0,
});
}
};
if let Some([min, max]) = limit {
for (side, bound) in [(1.0_f32, min), (-1.0, max)] {
let gap = side * (value - bound);
let bias = if gap > 0.0 {
gap / dt
} else {
LIMIT_BIAS / dt * gap
};
let (lower, upper) = if side > 0.0 {
(0.0, f32::INFINITY)
} else {
(f32::NEG_INFINITY, 0.0)
};
push(side * bias, 0.0, lower, upper);
}
}
if let Some(JointSpring {
target,
stiffness,
damping,
}) = spring
{
let denominator = dt * (damping + dt * stiffness);
if denominator > 0.0 {
let factor = dt * stiffness / (damping + dt * stiffness);
push(
factor / dt * (value - target),
1.0 / denominator,
f32::NEG_INFINITY,
f32::INFINITY,
);
}
}
if let Some(motor) = motor {
let limit = motor.max_force.max(0.0) * dt;
push(-motor.speed, 0.0, -limit, limit);
}
}
}
fn change_speed(&mut self, generalized: &[f32], impulse: f32) {
let dofs = self.velocity.len();
for k in 0..dofs {
let change: f32 = (0..dofs)
.map(|l| self.inverse[(k, l)] * generalized[l])
.sum();
self.velocity[k] += change * impulse;
}
}
fn solve_rows(&mut self, bodies: &mut [Body]) {
if self.rows.is_empty() {
return;
}
for index in 0..self.rows.len() {
let row = &self.rows[index];
let speed = self.velocity[row.dof];
let change =
-row.mass * (speed + row.bias + row.softness * row.impulse);
let total = (row.impulse + change).clamp(row.lower, row.upper);
let applied = total - row.impulse;
let dof = row.dof;
self.rows[index].impulse = total;
let mut unit = vec![0.0; self.velocity.len()];
unit[dof] = 1.0;
self.change_speed(&unit, applied);
}
self.refresh(bodies);
}
fn integrate(&mut self, bodies: &mut [Body], dt: f32) {
if self.floating {
let spin = Vector3::from_column_slice(&self.velocity[..3]);
let root = &mut bodies[self.links[0].body];
root.position +=
Vector3::from_column_slice(&self.velocity[3..6]) * dt;
if spin != Vector3::zeros() {
root.rotation = sim_math::rotation_from_scaled_axis(spin * dt)
* root.rotation;
}
}
for link in &mut self.links[1..] {
let dof = link.dof;
match link.kind {
Dof::Revolute | Dof::Prismatic => {
link.q += self.velocity[dof] * dt;
let limit = match link.joint.kind {
JointKind::Hinge { limit, .. }
| JointKind::Slider { limit, .. } => limit,
_ => None,
};
if let Some([min, max]) = limit {
if link.q < min {
link.q = min;
self.velocity[dof] = self.velocity[dof].max(0.0);
} else if link.q > max {
link.q = max;
self.velocity[dof] = self.velocity[dof].min(0.0);
}
}
}
Dof::Spherical => {
let spin = Vector3::from_column_slice(
&self.velocity[dof..dof + 3],
);
link.rel =
UnitQuaternion::from_scaled_axis(spin * dt) * link.rel;
}
Dof::Root | Dof::Fixed => {}
}
}
self.place(bodies);
}
}
pub(super) type Side = (usize, Vector3<f32>, Vector3<f32>);
impl Articulations {
pub(super) fn contains(&self, body: usize) -> bool {
self.of_body.get(body).is_some_and(Option::is_some)
}
pub(super) fn free_motion(
&mut self,
bodies: &mut [Body],
gravity: Vector3<f32>,
dt: f32,
) {
for tree in &mut self.trees {
tree.free_motion(bodies, gravity, dt);
tree.build_rows(dt);
}
}
pub(super) fn solve_rows(&mut self, bodies: &mut [Body]) {
for tree in &mut self.trees {
tree.solve_rows(bodies);
}
}
pub(super) fn integrate(&mut self, bodies: &mut [Body], dt: f32) {
for tree in &mut self.trees {
tree.integrate(bodies, dt);
}
}
fn gather(&self, sides: &[Side]) -> Vec<(usize, Vec<f32>)> {
let mut forces: Vec<(usize, Vec<f32>)> = Vec::new();
for (body, angular, linear) in sides {
let Some((tree, link)) = self.of_body[*body] else {
continue;
};
let at = forces.iter().position(|(other, _)| *other == tree);
let at = at.unwrap_or_else(|| {
forces.push((tree, vec![0.0; self.trees[tree].velocity.len()]));
forces.len() - 1
});
self.trees[tree].generalized(
link,
angular,
linear,
&mut forces[at].1,
);
}
forces
}
pub(super) fn inverse_mass(
&self,
bodies: &[Body],
inertia: &[Matrix3<f32>],
sides: &[Side],
) -> f32 {
let mut k = 0.0;
for (body, angular, linear) in sides {
if !self.contains(*body) {
k += bodies[*body].inverse_mass * linear.norm_squared()
+ angular.dot(&(inertia[*body] * angular));
}
}
for (tree, force) in self.gather(sides) {
let inverse = &self.trees[tree].inverse;
for (k_index, value) in force.iter().enumerate() {
let row: f32 = force
.iter()
.enumerate()
.map(|(l, other)| inverse[(k_index, l)] * other)
.sum();
k += value * row;
}
}
k
}
pub(super) fn apply(
&mut self,
bodies: &mut [Body],
inertia: &[Matrix3<f32>],
sides: &[Side],
impulse: f32,
) {
for (body, angular, linear) in sides {
if !self.contains(*body) {
let rigid = &mut bodies[*body];
rigid.velocity += linear * (rigid.inverse_mass * impulse);
rigid.angular_velocity += inertia[*body] * angular * impulse;
}
}
for (tree, force) in self.gather(sides) {
let tree = &mut self.trees[tree];
tree.change_speed(&force, impulse);
tree.refresh(bodies);
}
}
}