use serde::{Serialize, Deserialize};
#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)]
pub struct Matrix4 {
pub m: [[f64; 4]; 4],
}
impl Matrix4 {
pub fn new(
m00: f64, m01: f64, m02: f64, m03: f64,
m10: f64, m11: f64, m12: f64, m13: f64,
m20: f64, m21: f64, m22: f64, m23: f64,
m30: f64, m31: f64, m32: f64, m33: f64,
) -> Self {
Self {
m: [
[m00, m01, m02, m03],
[m10, m11, m12, m13],
[m20, m21, m22, m23],
[m30, m31, m32, m33],
],
}
}
pub fn identity() -> Self {
Self::new(
1.0, 0.0, 0.0, 0.0,
0.0, 1.0, 0.0, 0.0,
0.0, 0.0, 1.0, 0.0,
0.0, 0.0, 0.0, 1.0,
)
}
pub fn from_translation(x: f64, y: f64, z: f64) -> Self {
Self::new(
1.0, 0.0, 0.0, x,
0.0, 1.0, 0.0, y,
0.0, 0.0, 1.0, z,
0.0, 0.0, 0.0, 1.0,
)
}
pub fn translation_2d(x: f64, y: f64) -> Self {
Self::from_translation(x, y, 0.0)
}
pub fn rotation_2d(angle: f64) -> Self {
let (s, c) = angle.sin_cos();
Self::new(
c, -s, 0.0, 0.0,
s, c, 0.0, 0.0,
0.0, 0.0, 1.0, 0.0,
0.0, 0.0, 0.0, 1.0,
)
}
pub fn scale_2d(sx: f64, sy: f64) -> Self {
Self::new(
sx, 0.0, 0.0, 0.0,
0.0, sy, 0.0, 0.0,
0.0, 0.0, 1.0, 0.0,
0.0, 0.0, 0.0, 1.0,
)
}
pub fn multiply(&self, other: &Matrix4) -> Matrix4 {
let mut result = [[0.0f64; 4]; 4];
for i in 0..4 {
for j in 0..4 {
let mut sum = 0.0;
for k in 0..4 {
sum += self.m[i][k] * other.m[k][j];
}
result[i][j] = sum;
}
}
Self { m: result }
}
pub fn transform_point(&self, x: f64, y: f64, z: f64) -> (f64, f64, f64) {
let w = self.m[3][0] * x + self.m[3][1] * y + self.m[3][2] * z + self.m[3][3];
if w.abs() < 1e-12 {
return (0.0, 0.0, 0.0);
}
(
(self.m[0][0] * x + self.m[0][1] * y + self.m[0][2] * z + self.m[0][3]) / w,
(self.m[1][0] * x + self.m[1][1] * y + self.m[1][2] * z + self.m[1][3]) / w,
(self.m[2][0] * x + self.m[2][1] * y + self.m[2][2] * z + self.m[2][3]) / w,
)
}
pub fn determinant(&self) -> f64 {
let m = &self.m;
let mut det = 0.0;
for c in 0..4 {
let mut sub = [[0.0f64; 3]; 3];
for i in 1..4 {
let mut sub_col = 0;
for j in 0..4 {
if j == c {
continue;
}
sub[i - 1][sub_col] = m[i][j];
sub_col += 1;
}
}
let minor = sub[0][0] * (sub[1][1] * sub[2][2] - sub[1][2] * sub[2][1])
- sub[0][1] * (sub[1][0] * sub[2][2] - sub[1][2] * sub[2][0])
+ sub[0][2] * (sub[1][0] * sub[2][1] - sub[1][1] * sub[2][0]);
det += if c % 2 == 0 { 1.0 } else { -1.0 } * m[0][c] * minor;
}
det
}
}
impl Default for Matrix4 {
fn default() -> Self {
Self::identity()
}
}
impl std::ops::Mul for Matrix4 {
type Output = Matrix4;
fn mul(self, other: Matrix4) -> Matrix4 {
self.multiply(&other)
}
}
impl std::ops::Mul<&Matrix4> for Matrix4 {
type Output = Matrix4;
fn mul(self, other: &Matrix4) -> Matrix4 {
self.multiply(other)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_identity() {
let m = Matrix4::identity();
assert_eq!(m, Matrix4::default());
assert!((m.determinant() - 1.0).abs() < 1e-10);
}
#[test]
fn test_translation() {
let m = Matrix4::from_translation(1.0, 2.0, 3.0);
let (x, y, z) = m.transform_point(0.0, 0.0, 0.0);
assert!((x - 1.0).abs() < 1e-10);
assert!((y - 2.0).abs() < 1e-10);
assert!((z - 3.0).abs() < 1e-10);
}
}