use bevy_ecs::entity::Entity;
use bevy_ecs::prelude::{Component, World};
use nalgebra::{Matrix3, Quaternion, Rotation3, UnitQuaternion, Vector3};
use serde::{Deserialize, Serialize};
use super::articulation::{Articulations, Side};
use super::Body;
use crate::runtime::sim_math;
#[derive(Component, Clone, Copy, Debug, PartialEq, Serialize, Deserialize)]
#[serde(default)]
pub struct Joint {
pub target: Entity,
pub kind: JointKind,
pub anchor: [f32; 3],
pub frame: [f32; 3],
pub target_anchor: [f32; 3],
pub target_frame: [f32; 3],
pub collide_connected: bool,
pub break_force: f32,
pub break_torque: f32,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct JointBroken {
pub joint: Entity,
pub target: Entity,
pub force: f32,
pub torque: f32,
}
impl Default for Joint {
fn default() -> Self {
Self {
target: Entity::PLACEHOLDER,
kind: JointKind::default(),
anchor: [0.0; 3],
frame: [0.0; 3],
target_anchor: [0.0; 3],
target_frame: [0.0; 3],
collide_connected: false,
break_force: 0.0,
break_torque: 0.0,
}
}
}
impl Joint {
#[must_use]
pub fn new(
kind: JointKind,
target: Entity,
anchor: [f32; 3],
target_anchor: [f32; 3],
) -> Self {
Self {
target,
kind,
anchor,
target_anchor,
..Self::default()
}
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Serialize, Deserialize)]
pub enum JointKind {
#[default]
Fixed,
Hinge {
limit: Option<[f32; 2]>,
spring: Option<JointSpring>,
motor: Option<JointMotor>,
},
Slider {
limit: Option<[f32; 2]>,
spring: Option<JointSpring>,
motor: Option<JointMotor>,
},
BallSocket,
ConeTwist { swing: f32, twist: [f32; 2] },
Distance { min: f32, max: f32 },
Spring {
rest_length: f32,
stiffness: f32,
damping: f32,
},
Generic {
linear: [JointAxis; 3],
angular: [JointAxis; 3],
},
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Serialize, Deserialize)]
#[serde(default)]
pub struct JointAxis {
pub motion: AxisMotion,
pub spring: Option<JointSpring>,
pub motor: Option<JointMotor>,
}
impl JointAxis {
pub const LOCKED: Self = Self {
motion: AxisMotion::Locked,
spring: None,
motor: None,
};
pub const FREE: Self = Self {
motion: AxisMotion::Free,
spring: None,
motor: None,
};
fn free_or_limited(
limit: Option<[f32; 2]>,
spring: Option<JointSpring>,
motor: Option<JointMotor>,
) -> Self {
Self {
motion: limit.map_or(AxisMotion::Free, |[min, max]| {
AxisMotion::Limited { min, max }
}),
spring,
motor,
}
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Serialize, Deserialize)]
pub enum AxisMotion {
#[default]
Locked,
Free,
Limited {
min: f32,
max: f32,
},
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct JointSpring {
pub target: f32,
pub stiffness: f32,
pub damping: f32,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct JointMotor {
pub speed: f32,
pub max_force: f32,
}
struct Axes {
linear: [JointAxis; 3],
angular: [JointAxis; 3],
cone: Option<f32>,
distance: Option<(AxisMotion, Option<JointSpring>)>,
}
impl JointKind {
fn axes(&self) -> Axes {
use JointAxis as A;
let axes = |linear, angular| Axes {
linear,
angular,
cone: None,
distance: None,
};
match *self {
Self::Fixed => axes([A::LOCKED; 3], [A::LOCKED; 3]),
Self::Hinge {
limit,
spring,
motor,
} => axes(
[A::LOCKED; 3],
[
A::free_or_limited(limit, spring, motor),
A::LOCKED,
A::LOCKED,
],
),
Self::Slider {
limit,
spring,
motor,
} => axes(
[
A::free_or_limited(limit, spring, motor),
A::LOCKED,
A::LOCKED,
],
[A::LOCKED; 3],
),
Self::BallSocket => axes([A::LOCKED; 3], [A::FREE; 3]),
Self::ConeTwist { swing, twist } => Axes {
cone: Some(swing),
..axes(
[A::LOCKED; 3],
[
A::free_or_limited(Some(twist), None, None),
A::FREE,
A::FREE,
],
)
},
Self::Distance { min, max } => Axes {
distance: Some((AxisMotion::Limited { min, max }, None)),
..axes([A::FREE; 3], [A::FREE; 3])
},
Self::Spring {
rest_length,
stiffness,
damping,
} => Axes {
distance: Some((
AxisMotion::Free,
Some(JointSpring {
target: rest_length,
stiffness,
damping,
}),
)),
..axes([A::FREE; 3], [A::FREE; 3])
},
Self::Generic { linear, angular } => axes(linear, angular),
}
}
}
const JOINT_BIAS: f32 = 0.2;
pub(super) struct JointLink {
pub entity: Entity,
pub a: Option<usize>,
pub b: usize,
pub joint: Joint,
}
pub(super) struct JointRow {
a: Option<usize>,
b: usize,
linear: Vector3<f32>,
angular_a: Vector3<f32>,
angular_b: Vector3<f32>,
mass: f32,
bias: f32,
softness: f32,
lower: f32,
upper: f32,
pub impulse: f32,
pub key: u8,
}
#[derive(Clone, Copy)]
struct Jacobian {
linear: Vector3<f32>,
angular_a: Vector3<f32>,
angular_b: Vector3<f32>,
}
impl Jacobian {
fn negated(self) -> Self {
Self {
linear: -self.linear,
angular_a: -self.angular_a,
angular_b: -self.angular_b,
}
}
}
struct RowBuilder<'a> {
bodies: &'a [Body],
inertia: &'a [Matrix3<f32>],
articulations: &'a Articulations,
link: &'a JointLink,
dt: f32,
rows: Vec<JointRow>,
}
impl RowBuilder<'_> {
fn push(
&mut self,
jacobian: Jacobian,
key: u8,
bias: f32,
softness: f32,
[lower, upper]: [f32; 2],
) {
let b = &self.bodies[self.link.b];
let articulated = self.articulations.contains(self.link.b)
|| self.link.a.is_some_and(|a| self.articulations.contains(a));
let mut k = b.inverse_mass * jacobian.linear.norm_squared()
+ jacobian
.angular_b
.dot(&(self.inertia[self.link.b] * jacobian.angular_b));
if let Some(a) = self.link.a {
k += self.bodies[a].inverse_mass * jacobian.linear.norm_squared()
+ jacobian
.angular_a
.dot(&(self.inertia[a] * jacobian.angular_a));
}
if articulated {
k = self.articulations.inverse_mass(
self.bodies,
self.inertia,
&sides(self.link.a, self.link.b, &jacobian),
);
}
if k + softness <= 0.0 {
return;
}
self.rows.push(JointRow {
a: self.link.a,
b: self.link.b,
linear: jacobian.linear,
angular_a: jacobian.angular_a,
angular_b: jacobian.angular_b,
mass: 1.0 / (k + softness),
bias,
softness,
lower,
upper,
impulse: 0.0,
key,
});
}
fn axis(
&mut self,
axis: JointAxis,
motion: AxisMotion,
jacobian: Jacobian,
value: f32,
base: u8,
) {
let rigid = JOINT_BIAS / self.dt;
match motion {
AxisMotion::Locked => self.push(
jacobian,
base,
rigid * value,
0.0,
[f32::NEG_INFINITY, f32::INFINITY],
),
AxisMotion::Free => {}
AxisMotion::Limited { min, max } => {
for (side, bound, key) in [(1.0, min, 0), (-1.0, max, 1)] {
if !bound.is_finite() {
continue;
}
let gap = side * (value - bound);
let bias = if gap > 0.0 {
gap / self.dt
} else {
rigid * gap
};
let jacobian = if side > 0.0 {
jacobian
} else {
jacobian.negated()
};
self.push(
jacobian,
base + key,
bias,
0.0,
[0.0, f32::INFINITY],
);
}
}
}
if let Some(spring) = axis.spring {
let JointSpring {
target,
stiffness,
damping,
} = spring;
let denominator = self.dt * (damping + self.dt * stiffness);
if denominator > 0.0 {
let factor =
self.dt * stiffness / (damping + self.dt * stiffness);
self.push(
jacobian,
base + 2,
factor / self.dt * (value - target),
1.0 / denominator,
[f32::NEG_INFINITY, f32::INFINITY],
);
}
}
if let Some(motor) = axis.motor {
let limit = motor.max_force.max(0.0) * self.dt;
self.push(jacobian, base + 3, -motor.speed, 0.0, [-limit, limit]);
}
}
}
pub(super) fn gather_joints(
world: &mut World,
bodies: &[Body],
) -> Vec<JointLink> {
let find = |entity: Entity| {
bodies
.binary_search_by_key(&entity, |body| body.entity)
.ok()
};
let mut links: Vec<_> = world
.query::<(Entity, &Joint)>()
.iter(world)
.filter_map(|(entity, joint)| {
let a = if joint.target == Entity::PLACEHOLDER {
None
} else {
Some(find(joint.target)?)
};
Some(JointLink {
entity,
a,
b: find(entity)?,
joint: *joint,
})
})
.collect();
links.sort_unstable_by_key(|link| link.entity);
links
}
pub(super) fn excluded(links: &[JointLink], a: usize, b: usize) -> bool {
links.iter().any(|link| {
!link.joint.collide_connected
&& link.a.is_some_and(|target| {
(target, link.b) == (a, b) || (target, link.b) == (b, a)
})
})
}
pub(super) fn has_motor(joint: &Joint) -> bool {
let axes = joint.kind.axes();
axes.linear
.iter()
.chain(&axes.angular)
.any(|axis| axis.motor.is_some())
}
pub(super) fn joint_rows(
bodies: &[Body],
inertia: &[Matrix3<f32>],
articulations: &Articulations,
links: &[JointLink],
dt: f32,
) -> Vec<(usize, Vec<JointRow>)> {
links
.iter()
.enumerate()
.map(|(index, link)| {
let mut builder = RowBuilder {
bodies,
inertia,
articulations,
link,
dt,
rows: Vec::new(),
};
build(&mut builder);
(index, builder.rows)
})
.collect()
}
fn build(builder: &mut RowBuilder<'_>) {
let link = builder.link;
let joint = &link.joint;
let b = &builder.bodies[link.b];
let (a_position, a_rotation) = link
.a
.map_or((Vector3::zeros(), Rotation3::identity()), |a| {
(builder.bodies[a].position, builder.bodies[a].rotation)
});
let euler = |angles: [f32; 3]| {
sim_math::rotation_from_euler(angles[0], angles[1], angles[2])
};
let frame_a = a_rotation * euler(joint.target_frame);
let frame_b = b.rotation * euler(joint.frame);
let anchor_a = a_position + a_rotation * Vector3::from(joint.target_anchor);
let anchor_b = b.position + b.rotation * Vector3::from(joint.anchor);
let axes = joint.kind.axes();
let offset = anchor_b - anchor_a;
let (arm_a, arm_b) = (anchor_b - a_position, anchor_b - b.position);
for (index, axis) in axes.linear.into_iter().enumerate() {
let direction = frame_a * Vector3::ith(index, 1.0);
builder.axis(
axis,
axis.motion,
Jacobian {
linear: direction,
angular_a: arm_a.cross(&direction),
angular_b: arm_b.cross(&direction),
},
offset.dot(&direction),
index as u8 * 4,
);
}
let relative = UnitQuaternion::from_rotation_matrix(&frame_a).inverse()
* UnitQuaternion::from_rotation_matrix(&frame_b);
let (twist, swing) = swing_twist(&relative);
let angles = [twist, swing[0], swing[1]];
for (index, axis) in axes.angular.into_iter().enumerate() {
let direction = frame_a * Vector3::ith(index, 1.0);
builder.axis(
axis,
axis.motion,
Jacobian {
linear: Vector3::zeros(),
angular_a: direction,
angular_b: direction,
},
angles[index],
12 + index as u8 * 4,
);
}
if let Some(limit) = axes.cone {
let size = sim_math::sqrt(swing[0] * swing[0] + swing[1] * swing[1]);
if size > 1e-6 {
let direction =
frame_a * Vector3::new(0.0, swing[0] / size, swing[1] / size);
builder.axis(
JointAxis::FREE,
AxisMotion::Limited {
min: f32::NEG_INFINITY,
max: limit,
},
Jacobian {
linear: Vector3::zeros(),
angular_a: direction,
angular_b: direction,
},
size,
24,
);
}
}
if let Some((motion, spring)) = axes.distance {
let length = offset.norm();
let direction = if length > 1e-6 {
offset / length
} else {
frame_a * Vector3::x()
};
let arm_a = anchor_a - a_position;
builder.axis(
JointAxis {
motion,
spring,
motor: None,
},
motion,
Jacobian {
linear: direction,
angular_a: arm_a.cross(&direction),
angular_b: arm_b.cross(&direction),
},
length,
28,
);
}
}
pub(super) fn swing_twist(q: &UnitQuaternion<f32>) -> (f32, [f32; 2]) {
let sign = if q.w < 0.0 { -1.0 } else { 1.0 };
let (w, x) = (sign * q.w, sign * q.i);
let length = sim_math::sqrt(w * w + x * x);
let (tw, tx) = if length > 1e-6 {
(w / length, x / length)
} else {
(1.0, 0.0)
};
let twist = 2.0 * sim_math::atan2(tx, tw);
let swing = q.quaternion() * sign * Quaternion::new(tw, -tx, 0.0, 0.0);
let (sw, sy, sz) = if swing.w < 0.0 {
(-swing.w, -swing.j, -swing.k)
} else {
(swing.w, swing.j, swing.k)
};
let size = sim_math::sqrt(sy * sy + sz * sz);
if size < 1e-6 {
return (twist, [2.0 * sy, 2.0 * sz]);
}
let angle = 2.0 * sim_math::atan2(size, sw);
(twist, [sy / size * angle, sz / size * angle])
}
fn sides(a: Option<usize>, b: usize, jacobian: &Jacobian) -> Vec<Side> {
let mut sides = vec![(b, jacobian.angular_b, jacobian.linear)];
if let Some(a) = a {
sides.push((a, -jacobian.angular_a, -jacobian.linear));
}
sides
}
fn apply(
bodies: &mut [Body],
inertia: &[Matrix3<f32>],
articulations: &mut Articulations,
row: &JointRow,
impulse: f32,
) {
if articulations.contains(row.b)
|| row.a.is_some_and(|a| articulations.contains(a))
{
let jacobian = Jacobian {
linear: row.linear,
angular_a: row.angular_a,
angular_b: row.angular_b,
};
let sides = sides(row.a, row.b, &jacobian);
articulations.apply(bodies, inertia, &sides, impulse);
return;
}
if let Some(a) = row.a {
let body = &mut bodies[a];
body.velocity -= row.linear * (body.inverse_mass * impulse);
body.angular_velocity -= inertia[a] * row.angular_a * impulse;
}
let body = &mut bodies[row.b];
body.velocity += row.linear * (body.inverse_mass * impulse);
body.angular_velocity += inertia[row.b] * row.angular_b * impulse;
}
pub(super) fn broken(
links: &[JointLink],
rows: &[(usize, Vec<JointRow>)],
dt: f32,
) -> Vec<(usize, f32, f32)> {
rows.iter()
.filter_map(|(link, rows)| {
let joint = &links[*link].joint;
if joint.break_force <= 0.0 && joint.break_torque <= 0.0 {
return None;
}
let (force, torque) = rows.iter().fold(
(Vector3::zeros(), Vector3::zeros()),
|(force, torque), row| {
if row.linear == Vector3::zeros() {
(force, torque + row.angular_b * row.impulse)
} else {
(force + row.linear * row.impulse, torque)
}
},
);
let (force, torque) = (force.norm() / dt, torque.norm() / dt);
let over = |load: f32, limit: f32| limit > 0.0 && load > limit;
(over(force, joint.break_force) || over(torque, joint.break_torque))
.then_some((*link, force, torque))
})
.collect()
}
pub(super) fn warm_start(
bodies: &mut [Body],
inertia: &[Matrix3<f32>],
articulations: &mut Articulations,
rows: &[(usize, Vec<JointRow>)],
) {
for row in rows.iter().flat_map(|(_, rows)| rows) {
apply(bodies, inertia, articulations, row, row.impulse);
}
}
pub(super) fn solve(
bodies: &mut [Body],
inertia: &[Matrix3<f32>],
articulations: &mut Articulations,
rows: &mut [(usize, Vec<JointRow>)],
) {
for row in rows.iter_mut().flat_map(|(_, rows)| rows) {
let b = &bodies[row.b];
let mut speed = row.linear.dot(&b.velocity)
+ row.angular_b.dot(&b.angular_velocity);
if let Some(a) = row.a {
let a = &bodies[a];
speed -= row.linear.dot(&a.velocity)
+ row.angular_a.dot(&a.angular_velocity);
}
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;
row.impulse = total;
apply(bodies, inertia, articulations, row, applied);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn swing_twist_splits_rotations_about_each_axis() {
let about = |axis: Vector3<f32>, angle: f32| {
UnitQuaternion::from_axis_angle(
&nalgebra::Unit::new_normalize(axis),
angle,
)
};
let (twist, swing) = swing_twist(&about(Vector3::x(), 0.7));
assert!((twist - 0.7).abs() < 1e-5 && swing == [0.0, 0.0]);
let (twist, swing) = swing_twist(&about(Vector3::y(), -0.4));
assert!(twist.abs() < 1e-6, "{twist}");
assert!((swing[0] + 0.4).abs() < 1e-5 && swing[1].abs() < 1e-6);
let (twist, swing) = swing_twist(
&(about(Vector3::z(), 0.5) * about(Vector3::x(), -1.2)),
);
assert!((twist + 1.2).abs() < 1e-5, "{twist}");
assert!(swing[0].abs() < 1e-5 && (swing[1] - 0.5).abs() < 1e-5);
}
}