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);
    }
}