use symplex::prelude::*;
use symplex::quaternion::Quaternion;
fn assert_close(actual: f64, expected: f64, tol: f64, msg: &str) {
assert!(
(actual - expected).abs() < tol,
"{msg}: expected {expected}, got {actual}"
);
}
fn quat_to_f64(q: &Quaternion) -> (f64, f64, f64, f64) {
(
q.w.eval().eval_f64().unwrap(),
q.x.eval().eval_f64().unwrap(),
q.y.eval().eval_f64().unwrap(),
q.z.eval().eval_f64().unwrap(),
)
}
#[test]
fn quaternion_identity() {
let ctx = Context::new();
let id = Quaternion::identity(&ctx);
let (w, x, y, z) = quat_to_f64(&id);
assert_close(w, 1.0, 1e-12, "identity w");
assert_close(x, 0.0, 1e-12, "identity x");
assert_close(y, 0.0, 1e-12, "identity y");
assert_close(z, 0.0, 1e-12, "identity z");
}
#[test]
fn quaternion_mul_identity() {
let ctx = Context::new();
let q = Quaternion::new(ctx.int(1), ctx.int(2), ctx.int(3), ctx.int(4));
let id = Quaternion::identity(&ctx);
let result = q.mul(&id);
let (w, x, y, z) = quat_to_f64(&result.eval());
assert_close(w, 1.0, 1e-12, "q*id w");
assert_close(x, 2.0, 1e-12, "q*id x");
assert_close(y, 3.0, 1e-12, "q*id y");
assert_close(z, 4.0, 1e-12, "q*id z");
let result2 = id.mul(&q);
let (w2, x2, y2, z2) = quat_to_f64(&result2.eval());
assert_close(w2, 1.0, 1e-12, "id*q w");
assert_close(x2, 2.0, 1e-12, "id*q x");
assert_close(y2, 3.0, 1e-12, "id*q y");
assert_close(z2, 4.0, 1e-12, "id*q z");
}
#[test]
fn quaternion_mul_conjugate() {
let ctx = Context::new();
let q = Quaternion::new(ctx.int(1), ctx.int(2), ctx.int(3), ctx.int(4));
let qc = q.conjugate();
let product = q.mul(&qc).eval();
let (w, x, y, z) = quat_to_f64(&product);
assert_close(w, 30.0, 1e-12, "q*q* w should be |q|²");
assert_close(x, 0.0, 1e-12, "q*q* x should be 0");
assert_close(y, 0.0, 1e-12, "q*q* y should be 0");
assert_close(z, 0.0, 1e-12, "q*q* z should be 0");
}
#[test]
fn quaternion_i_squared() {
let ctx = Context::new();
let qi = Quaternion::new(ctx.int(0), ctx.int(1), ctx.int(0), ctx.int(0));
let result = qi.mul(&qi).eval();
let (w, x, y, z) = quat_to_f64(&result);
assert_close(w, -1.0, 1e-12, "i² w");
assert_close(x, 0.0, 1e-12, "i² x");
assert_close(y, 0.0, 1e-12, "i² y");
assert_close(z, 0.0, 1e-12, "i² z");
}
#[test]
fn quaternion_j_squared() {
let ctx = Context::new();
let qj = Quaternion::new(ctx.int(0), ctx.int(0), ctx.int(1), ctx.int(0));
let result = qj.mul(&qj).eval();
let (w, x, y, z) = quat_to_f64(&result);
assert_close(w, -1.0, 1e-12, "j² w");
assert_close(x, 0.0, 1e-12, "j² x");
assert_close(y, 0.0, 1e-12, "j² y");
assert_close(z, 0.0, 1e-12, "j² z");
}
#[test]
fn quaternion_k_squared() {
let ctx = Context::new();
let qk = Quaternion::new(ctx.int(0), ctx.int(0), ctx.int(0), ctx.int(1));
let result = qk.mul(&qk).eval();
let (w, x, y, z) = quat_to_f64(&result);
assert_close(w, -1.0, 1e-12, "k² w");
assert_close(x, 0.0, 1e-12, "k² x");
assert_close(y, 0.0, 1e-12, "k² y");
assert_close(z, 0.0, 1e-12, "k² z");
}
#[test]
fn quaternion_ij_equals_k() {
let ctx = Context::new();
let qi = Quaternion::new(ctx.int(0), ctx.int(1), ctx.int(0), ctx.int(0));
let qj = Quaternion::new(ctx.int(0), ctx.int(0), ctx.int(1), ctx.int(0));
let result = qi.mul(&qj).eval();
let (w, x, y, z) = quat_to_f64(&result);
assert_close(w, 0.0, 1e-12, "i*j w");
assert_close(x, 0.0, 1e-12, "i*j x");
assert_close(y, 0.0, 1e-12, "i*j y");
assert_close(z, 1.0, 1e-12, "i*j z (should be k)");
}
#[test]
fn quaternion_ji_equals_neg_k() {
let ctx = Context::new();
let qi = Quaternion::new(ctx.int(0), ctx.int(1), ctx.int(0), ctx.int(0));
let qj = Quaternion::new(ctx.int(0), ctx.int(0), ctx.int(1), ctx.int(0));
let result = qj.mul(&qi).eval();
let (w, x, y, z) = quat_to_f64(&result);
assert_close(w, 0.0, 1e-12, "j*i w");
assert_close(x, 0.0, 1e-12, "j*i x");
assert_close(y, 0.0, 1e-12, "j*i y");
assert_close(z, -1.0, 1e-12, "j*i z (should be -k)");
}
#[test]
fn quaternion_conjugate() {
let ctx = Context::new();
let q = Quaternion::new(ctx.int(5), ctx.int(3), ctx.int(-7), ctx.int(2));
let qc = q.conjugate().eval();
let (w, x, y, z) = quat_to_f64(&qc);
assert_close(w, 5.0, 1e-12, "conjugate w");
assert_close(x, -3.0, 1e-12, "conjugate x");
assert_close(y, 7.0, 1e-12, "conjugate y");
assert_close(z, -2.0, 1e-12, "conjugate z");
}
#[test]
fn quaternion_norm_squared() {
let ctx = Context::new();
let q = Quaternion::new(ctx.int(1), ctx.int(2), ctx.int(3), ctx.int(4));
let n2 = q.norm_squared().eval().eval_f64().unwrap();
assert_close(n2, 30.0, 1e-12, "|q|²");
}
#[test]
fn quaternion_inverse() {
let ctx = Context::new();
let q = Quaternion::new(ctx.int(1), ctx.int(2), ctx.int(3), ctx.int(4));
let qi = q.inverse();
let product = q.mul(&qi).eval();
let (w, x, y, z) = quat_to_f64(&product);
assert_close(w, 1.0, 1e-10, "q*q⁻¹ w");
assert_close(x, 0.0, 1e-10, "q*q⁻¹ x");
assert_close(y, 0.0, 1e-10, "q*q⁻¹ y");
assert_close(z, 0.0, 1e-10, "q*q⁻¹ z");
}
#[test]
fn quaternion_to_rotation_identity() {
let ctx = Context::new();
let id = Quaternion::identity(&ctx);
let r = id.to_rotation_matrix();
assert_eq!(r.nrows(), 3);
assert_eq!(r.ncols(), 3);
for i in 0..3 {
for j in 0..3 {
let val = r.get(i, j).eval().eval_f64().unwrap();
let expected = if i == j { 1.0 } else { 0.0 };
assert_close(val, expected, 1e-12, &format!("R_id[{i}][{j}]"));
}
}
}
#[test]
fn quaternion_to_rotation_180_z() {
let ctx = Context::new();
let q = Quaternion::new(ctx.int(0), ctx.int(0), ctx.int(0), ctx.int(1));
let r = q.to_rotation_matrix();
let expected = [[-1.0, 0.0, 0.0], [0.0, -1.0, 0.0], [0.0, 0.0, 1.0]];
for (i, expected_row) in expected.iter().enumerate() {
for (j, &exp_val) in expected_row.iter().enumerate() {
let val = r.get(i, j).eval().eval_f64().unwrap();
assert_close(val, exp_val, 1e-12, &format!("R_180z[{i}][{j}]"));
}
}
}
#[test]
fn quaternion_from_axis_angle_z_90() {
let ctx = Context::new();
let zero = ctx.int(0);
let one = ctx.int(1);
let angle = &ctx.pi() / &ctx.int(2);
let q = Quaternion::from_axis_angle(&zero, &zero, &one, &angle);
let (w, x, y, z) = quat_to_f64(&q.eval());
let sqrt2_2 = std::f64::consts::FRAC_1_SQRT_2;
assert_close(w, sqrt2_2, 1e-12, "axis-angle w");
assert_close(x, 0.0, 1e-12, "axis-angle x");
assert_close(y, 0.0, 1e-12, "axis-angle y");
assert_close(z, sqrt2_2, 1e-12, "axis-angle z");
let r = q.to_rotation_matrix();
let expected = [[0.0, -1.0, 0.0], [1.0, 0.0, 0.0], [0.0, 0.0, 1.0]];
for (i, expected_row) in expected.iter().enumerate() {
for (j, &exp_val) in expected_row.iter().enumerate() {
let val = r.get(i, j).eval().eval_f64().unwrap();
assert_close(val, exp_val, 1e-10, &format!("R_z90[{i}][{j}]"));
}
}
}
#[test]
fn quaternion_angular_velocity_derivative() {
let ctx = Context::new();
let q = Quaternion::identity(&ctx);
let zero = ctx.int(0);
let wz = ctx.int(1);
let qdot = q.angular_velocity_derivative(&zero, &zero, &wz);
let (w, x, y, z) = quat_to_f64(&qdot.eval());
assert_close(w, 0.0, 1e-12, "q̇ w");
assert_close(x, 0.0, 1e-12, "q̇ x");
assert_close(y, 0.0, 1e-12, "q̇ y");
assert_close(z, 0.5, 1e-12, "q̇ z (½ωz)");
}
#[test]
fn quaternion_display() {
let ctx = Context::new();
let q = Quaternion::new(ctx.int(1), ctx.int(2), ctx.int(3), ctx.int(4));
let s = format!("{q}");
assert_eq!(s, "(1 + 2i + 3j + 4k)");
}
#[test]
fn quaternion_zero() {
let ctx = Context::new();
let z = Quaternion::zero(&ctx);
let (w, x, y, z) = quat_to_f64(&z);
assert_close(w, 0.0, 1e-12, "zero w");
assert_close(x, 0.0, 1e-12, "zero x");
assert_close(y, 0.0, 1e-12, "zero y");
assert_close(z, 0.0, 1e-12, "zero z");
}
#[test]
fn quaternion_from_vector() {
let ctx = Context::new();
let vx = ctx.int(3);
let vy = ctx.int(4);
let vz = ctx.int(5);
let q = Quaternion::from_vector(&vx, &vy, &vz);
let (w, x, y, z) = quat_to_f64(&q);
assert_close(w, 0.0, 1e-12, "from_vector w");
assert_close(x, 3.0, 1e-12, "from_vector x");
assert_close(y, 4.0, 1e-12, "from_vector y");
assert_close(z, 5.0, 1e-12, "from_vector z");
}
#[test]
fn quaternion_ijk_equals_neg_one() {
let ctx = Context::new();
let qi = Quaternion::new(ctx.int(0), ctx.int(1), ctx.int(0), ctx.int(0));
let qj = Quaternion::new(ctx.int(0), ctx.int(0), ctx.int(1), ctx.int(0));
let qk = Quaternion::new(ctx.int(0), ctx.int(0), ctx.int(0), ctx.int(1));
let ij = qi.mul(&qj);
let ijk = ij.mul(&qk).eval();
let (w, x, y, z) = quat_to_f64(&ijk);
assert_close(w, -1.0, 1e-12, "ijk w");
assert_close(x, 0.0, 1e-12, "ijk x");
assert_close(y, 0.0, 1e-12, "ijk y");
assert_close(z, 0.0, 1e-12, "ijk z");
}
#[test]
fn quaternion_normalize() {
let ctx = Context::new();
let q = Quaternion::new(ctx.int(1), ctx.int(2), ctx.int(3), ctx.int(4));
let qn = q.normalize();
let norm_val = qn.norm_squared().eval().eval_f64().unwrap();
assert_close(norm_val, 1.0, 1e-10, "|normalize(q)|²");
}
#[test]
fn quaternion_subs() {
let ctx = Context::new();
let theta = ctx.symbol("theta");
let q = Quaternion::new(theta.cos(), theta.sin(), ctx.int(0), ctx.int(0));
let pi_half = &ctx.pi() / &ctx.int(2);
let q2 = q.subs(&theta, &pi_half).eval();
let (w, x, y, z) = quat_to_f64(&q2);
assert_close(w, 0.0, 1e-12, "subs w = cos(π/2)");
assert_close(x, 1.0, 1e-12, "subs x = sin(π/2)");
assert_close(y, 0.0, 1e-12, "subs y");
assert_close(z, 0.0, 1e-12, "subs z");
}
#[test]
fn quaternion_from_axis_angle_x_90() {
let ctx = Context::new();
let one = ctx.int(1);
let zero = ctx.int(0);
let angle = &ctx.pi() / &ctx.int(2);
let q = Quaternion::from_axis_angle(&one, &zero, &zero, &angle);
let r = q.to_rotation_matrix();
let expected = [[1.0, 0.0, 0.0], [0.0, 0.0, -1.0], [0.0, 1.0, 0.0]];
for (i, expected_row) in expected.iter().enumerate() {
for (j, &exp_val) in expected_row.iter().enumerate() {
let val = r.get(i, j).eval().eval_f64().unwrap();
assert_close(val, exp_val, 1e-10, &format!("R_x90[{i}][{j}]"));
}
}
}
#[test]
fn quaternion_rotation_matrix_orthogonal() {
let ctx = Context::new();
let angle = &ctx.pi() / &ctx.int(3); let inv_sqrt3 = ctx.int(3).sqrt();
let ax = &ctx.int(1) / &inv_sqrt3;
let ay = &ctx.int(1) / &inv_sqrt3;
let az = &ctx.int(1) / &inv_sqrt3;
let q = Quaternion::from_axis_angle(&ax, &ay, &az, &angle);
let r = q.to_rotation_matrix();
let rt = r.transpose();
let product = r.matmul(&rt).unwrap();
for i in 0..3 {
for j in 0..3 {
let val = product.get(i, j).eval().eval_f64().unwrap();
let expected = if i == j { 1.0 } else { 0.0 };
assert_close(val, expected, 1e-10, &format!("R*Rᵀ[{i}][{j}]"));
}
}
}
#[test]
fn quaternion_mul_associative() {
let ctx = Context::new();
let p = Quaternion::new(ctx.int(1), ctx.int(2), ctx.int(3), ctx.int(4));
let q = Quaternion::new(ctx.int(5), ctx.int(-1), ctx.int(2), ctx.int(-3));
let r = Quaternion::new(ctx.int(-2), ctx.int(1), ctx.int(0), ctx.int(7));
let lhs = p.mul(&q).mul(&r).eval();
let rhs = p.mul(&q.mul(&r)).eval();
let (lw, lx, ly, lz) = quat_to_f64(&lhs);
let (rw, rx, ry, rz) = quat_to_f64(&rhs);
assert_close(lw, rw, 1e-10, "associativity w");
assert_close(lx, rx, 1e-10, "associativity x");
assert_close(ly, ry, 1e-10, "associativity y");
assert_close(lz, rz, 1e-10, "associativity z");
}
#[test]
fn quaternion_norm_product() {
let ctx = Context::new();
let p = Quaternion::new(ctx.int(1), ctx.int(2), ctx.int(3), ctx.int(4));
let q = Quaternion::new(ctx.int(5), ctx.int(-1), ctx.int(2), ctx.int(-3));
let norm_p = p.norm().eval().eval_f64().unwrap();
let norm_q = q.norm().eval().eval_f64().unwrap();
let norm_pq = p.mul(&q).norm().eval().eval_f64().unwrap();
assert_close(norm_pq, norm_p * norm_q, 1e-10, "|p*q| = |p|*|q|");
}