use std::fs;
use std::io::Read;
use std::path::Path;
use flate2::read::ZlibDecoder;
use matrw::{MatVariable, MatlabType, load_matfile};
use ndarray::{Array1, Array2};
use sprs::{CsMat, TriMat};
use crate::{Laplacian, error::GspError};
pub(super) fn load_mat_laplacian(path: &Path) -> Result<Laplacian, GspError> {
let primary = (|| {
let path_str = path
.to_str()
.ok_or_else(|| GspError::Parse("invalid path encoding".to_string()))?;
let mat = load_matfile(path_str)?;
let (_, var) = mat.iter().next().ok_or_else(|| {
GspError::Parse(format!(
"MAT file contains no variables: {}",
path.display()
))
})?;
match var {
MatVariable::SparseArray(sa) => {
let data = matlab_type_to_f64(&sa.value);
let nrows = sa.dim.first().copied().unwrap_or(0);
let ncols = sa.dim.get(1).copied().unwrap_or(0);
Ok(CsMat::new_csc(
(nrows, ncols),
sa.jc.clone(),
sa.ir.clone(),
data,
))
}
MatVariable::NumericArray(na) => {
let dense = numeric_array_to_dense(na)?;
Ok(dense_to_sparse(dense.view()))
}
_ => Err(GspError::Parse(format!(
"unsupported MAT variable type for Laplacian in {}",
path.display()
))),
}
})();
match primary {
Ok(l) => Ok(l),
Err(_) => load_mat_laplacian_fallback(path),
}
}
pub(super) fn load_mat_dense(path: &Path) -> Result<Array2<f64>, GspError> {
let primary = (|| {
let path_str = path
.to_str()
.ok_or_else(|| GspError::Parse("invalid path encoding".to_string()))?;
let mat = load_matfile(path_str)?;
let vars = mat.iter().collect::<Vec<_>>();
if vars.is_empty() {
return Err(GspError::Parse(format!(
"MAT file contains no variables: {}",
path.display()
)));
}
if vars.len() == 1 {
let dense = variable_to_dense(vars[0].1)?;
if dense.nrows() == 1 {
return Ok(dense.t().to_owned());
}
return Ok(dense);
}
let mut cols = Vec::<Array1<f64>>::with_capacity(vars.len());
for (_, var) in vars {
let d = variable_to_dense(var)?;
let (raw, _offset) = d.into_raw_vec_and_offset();
cols.push(Array1::from(raw));
}
let n = cols[0].len();
let mut out = Array2::<f64>::zeros((n, cols.len()));
for (j, c) in cols.into_iter().enumerate() {
if c.len() != n {
return Err(GspError::Dimensions(
"MAT variables have mismatched flattened lengths".to_string(),
));
}
out.column_mut(j).assign(&c);
}
Ok(out)
})();
match primary {
Ok(v) => Ok(v),
Err(_) => load_mat_dense_fallback(path),
}
}
#[derive(Debug, Clone)]
enum MatParsedValue {
Dense {
nrows: usize,
ncols: usize,
col_major_data: Vec<f64>,
},
Sparse {
nrows: usize,
ncols: usize,
ir: Vec<usize>,
jc: Vec<usize>,
data: Vec<f64>,
},
}
#[derive(Debug, Clone)]
struct MatParsedVar {
value: MatParsedValue,
}
fn load_mat_laplacian_fallback(path: &Path) -> Result<Laplacian, GspError> {
let vars = parse_mat_v5_file(path)?;
let first = vars.first().ok_or_else(|| {
GspError::Parse(format!(
"MAT file contains no parseable variables: {}",
path.display()
))
})?;
match &first.value {
MatParsedValue::Sparse {
nrows,
ncols,
ir,
jc,
data,
} => Ok(CsMat::new_csc(
(*nrows, *ncols),
jc.clone(),
ir.clone(),
data.clone(),
)),
MatParsedValue::Dense {
nrows,
ncols,
col_major_data,
} => {
let d = dense_from_col_major(*nrows, *ncols, col_major_data)?;
Ok(dense_to_sparse(d.view()))
}
}
}
fn load_mat_dense_fallback(path: &Path) -> Result<Array2<f64>, GspError> {
let vars = parse_mat_v5_file(path)?;
if vars.is_empty() {
return Err(GspError::Parse(format!(
"MAT file contains no parseable variables: {}",
path.display()
)));
}
if vars.len() == 1 {
let dense = mat_parsed_to_dense(&vars[0].value)?;
if dense.nrows() == 1 {
return Ok(dense.t().to_owned());
}
return Ok(dense);
}
let mut cols = Vec::<Array1<f64>>::with_capacity(vars.len());
for var in &vars {
let d = mat_parsed_to_dense(&var.value)?;
let (raw, _offset) = d.into_raw_vec_and_offset();
cols.push(Array1::from(raw));
}
let n = cols[0].len();
let mut out = Array2::<f64>::zeros((n, cols.len()));
for (j, c) in cols.into_iter().enumerate() {
if c.len() != n {
return Err(GspError::Dimensions(
"MAT variables have mismatched flattened lengths".to_string(),
));
}
out.column_mut(j).assign(&c);
}
Ok(out)
}
fn parse_mat_v5_file(path: &Path) -> Result<Vec<MatParsedVar>, GspError> {
let bytes = fs::read(path)?;
if bytes.len() < 128 {
return Err(GspError::Parse(format!(
"MAT file too small: {}",
path.display()
)));
}
let endian = &bytes[126..128];
if endian != b"IM" {
return Err(GspError::UnsupportedFormat(format!(
"unsupported MAT endianness marker {:?} in {}",
endian,
path.display()
)));
}
let mut out = Vec::<MatParsedVar>::new();
parse_mat_elements(&bytes[128..], &mut out)?;
Ok(out)
}
fn parse_mat_elements(buf: &[u8], out: &mut Vec<MatParsedVar>) -> Result<(), GspError> {
let mut pos = 0usize;
while pos < buf.len() {
if buf[pos..].iter().all(|b| *b == 0) {
break;
}
let (dtype, nbytes, data_start, next_pos) = parse_mat_tag(buf, pos)?;
let data_end = data_start
.checked_add(nbytes)
.ok_or_else(|| GspError::Parse("MAT element overflow".to_string()))?;
if data_end > buf.len() {
return Err(GspError::Parse(
"MAT element exceeds file buffer".to_string(),
));
}
match dtype {
14 => {
if let Some(v) = parse_mat_matrix(&buf[data_start..data_end])? {
out.push(v);
}
}
15 => {
let mut z = ZlibDecoder::new(&buf[data_start..data_end]);
let mut decoded = Vec::<u8>::new();
z.read_to_end(&mut decoded)?;
parse_mat_elements(&decoded, out)?;
}
_ => {}
}
pos = next_pos;
}
Ok(())
}
fn parse_mat_matrix(data: &[u8]) -> Result<Option<MatParsedVar>, GspError> {
let mut cursor = 0usize;
let (flags_t, flags_n, flags_s, flags_next) = parse_mat_tag(data, cursor)?;
cursor = flags_next;
if flags_t != 6 || flags_n < 8 || flags_s + flags_n > data.len() {
return Ok(None);
}
let class_flags = u32::from_le_bytes(
data[flags_s..(flags_s + 4)]
.try_into()
.map_err(|_| GspError::Parse("invalid MAT flags payload".to_string()))?,
);
let class = class_flags & 0xFF;
let is_complex = (class_flags & 0x0800) != 0;
let (dims_t, dims_n, dims_s, dims_next) = parse_mat_tag(data, cursor)?;
cursor = dims_next;
if dims_t != 5 || dims_s + dims_n > data.len() || dims_n % 4 != 0 {
return Ok(None);
}
let mut dims = Vec::<usize>::new();
for c in data[dims_s..(dims_s + dims_n)].chunks_exact(4) {
let v = i32::from_le_bytes(
c.try_into()
.map_err(|_| GspError::Parse("invalid MAT dim chunk".to_string()))?,
);
if v > 0 {
dims.push(v as usize);
}
}
if dims.is_empty() {
return Ok(None);
}
let nrows = dims[0];
let ncols = if dims.len() >= 2 {
dims[1..].iter().product()
} else {
1
};
let (name_t, name_n, name_s, name_next) = parse_mat_tag(data, cursor)?;
cursor = name_next;
if (name_t != 1 && name_t != 2) || name_s + name_n > data.len() {
return Ok(None);
}
if class == 5 {
let (ir_t, ir_n, ir_s, ir_next) = parse_mat_tag(data, cursor)?;
cursor = ir_next;
let ir = parse_mat_indices(ir_t, &data[ir_s..(ir_s + ir_n)])?;
let (jc_t, jc_n, jc_s, jc_next) = parse_mat_tag(data, cursor)?;
cursor = jc_next;
let jc = parse_mat_indices(jc_t, &data[jc_s..(jc_s + jc_n)])?;
let (pr_t, pr_n, pr_s, pr_next) = parse_mat_tag(data, cursor)?;
cursor = pr_next;
let pr = parse_mat_numeric_f64(pr_t, &data[pr_s..(pr_s + pr_n)])?;
if is_complex {
let (_, _, _, pi_next) = parse_mat_tag(data, cursor)?;
let _ = pi_next;
}
return Ok(Some(MatParsedVar {
value: MatParsedValue::Sparse {
nrows,
ncols,
ir,
jc,
data: pr,
},
}));
}
let (real_t, real_n, real_s, real_next) = parse_mat_tag(data, cursor)?;
cursor = real_next;
let real = parse_mat_numeric_f64(real_t, &data[real_s..(real_s + real_n)])?;
if is_complex && cursor < data.len() {
let (_, _, _, imag_next) = parse_mat_tag(data, cursor)?;
let _ = imag_next;
}
Ok(Some(MatParsedVar {
value: MatParsedValue::Dense {
nrows,
ncols,
col_major_data: real,
},
}))
}
fn parse_mat_tag(buf: &[u8], pos: usize) -> Result<(u32, usize, usize, usize), GspError> {
if pos + 4 > buf.len() {
return Err(GspError::Parse(
"unexpected end while reading MAT tag".to_string(),
));
}
let word = u32::from_le_bytes(
buf[pos..(pos + 4)]
.try_into()
.map_err(|_| GspError::Parse("invalid MAT tag word".to_string()))?,
);
if (word >> 16) != 0 {
let dtype = word & 0xFFFF;
let nbytes = (word >> 16) as usize;
let data_start = pos + 4;
let next = data_start
.checked_add(align_up(nbytes, 4))
.ok_or_else(|| GspError::Parse("MAT small-tag overflow".to_string()))?;
return Ok((dtype, nbytes, data_start, next));
}
if pos + 8 > buf.len() {
return Err(GspError::Parse(
"unexpected end while reading MAT tag size".to_string(),
));
}
let dtype = word;
let nbytes = u32::from_le_bytes(
buf[(pos + 4)..(pos + 8)]
.try_into()
.map_err(|_| GspError::Parse("invalid MAT tag size".to_string()))?,
) as usize;
let data_start = pos + 8;
let next = data_start
.checked_add(align_up(nbytes, 8))
.ok_or_else(|| GspError::Parse("MAT tag overflow".to_string()))?;
Ok((dtype, nbytes, data_start, next))
}
fn parse_mat_numeric_f64(dtype: u32, bytes: &[u8]) -> Result<Vec<f64>, GspError> {
match dtype {
1 => Ok(bytes.iter().map(|b| (*b as i8) as f64).collect()),
2 => Ok(bytes.iter().map(|b| *b as f64).collect()),
3 => {
if bytes.len() % 2 != 0 {
return Err(GspError::Parse("miINT16 payload misaligned".to_string()));
}
Ok(bytes
.chunks_exact(2)
.map(|c| i16::from_le_bytes([c[0], c[1]]) as f64)
.collect())
}
4 => {
if bytes.len() % 2 != 0 {
return Err(GspError::Parse("miUINT16 payload misaligned".to_string()));
}
Ok(bytes
.chunks_exact(2)
.map(|c| u16::from_le_bytes([c[0], c[1]]) as f64)
.collect())
}
5 => {
if bytes.len() % 4 != 0 {
return Err(GspError::Parse("miINT32 payload misaligned".to_string()));
}
Ok(bytes
.chunks_exact(4)
.map(|c| i32::from_le_bytes([c[0], c[1], c[2], c[3]]) as f64)
.collect())
}
6 => {
if bytes.len() % 4 != 0 {
return Err(GspError::Parse("miUINT32 payload misaligned".to_string()));
}
Ok(bytes
.chunks_exact(4)
.map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]) as f64)
.collect())
}
7 => {
if bytes.len() % 4 != 0 {
return Err(GspError::Parse("miSINGLE payload misaligned".to_string()));
}
Ok(bytes
.chunks_exact(4)
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]) as f64)
.collect())
}
9 => {
if bytes.len() % 8 != 0 {
return Err(GspError::Parse("miDOUBLE payload misaligned".to_string()));
}
Ok(bytes
.chunks_exact(8)
.map(|c| f64::from_le_bytes([c[0], c[1], c[2], c[3], c[4], c[5], c[6], c[7]]))
.collect())
}
12 => {
if bytes.len() % 8 != 0 {
return Err(GspError::Parse("miINT64 payload misaligned".to_string()));
}
Ok(bytes
.chunks_exact(8)
.map(|c| {
i64::from_le_bytes([c[0], c[1], c[2], c[3], c[4], c[5], c[6], c[7]]) as f64
})
.collect())
}
13 => {
if bytes.len() % 8 != 0 {
return Err(GspError::Parse("miUINT64 payload misaligned".to_string()));
}
Ok(bytes
.chunks_exact(8)
.map(|c| {
u64::from_le_bytes([c[0], c[1], c[2], c[3], c[4], c[5], c[6], c[7]]) as f64
})
.collect())
}
_ => Err(GspError::UnsupportedFormat(format!(
"unsupported MAT numeric dtype {dtype}"
))),
}
}
fn parse_mat_indices(dtype: u32, bytes: &[u8]) -> Result<Vec<usize>, GspError> {
match dtype {
5 => {
if bytes.len() % 4 != 0 {
return Err(GspError::Parse("miINT32 indices misaligned".to_string()));
}
Ok(bytes
.chunks_exact(4)
.map(|c| i32::from_le_bytes([c[0], c[1], c[2], c[3]]) as usize)
.collect())
}
6 => {
if bytes.len() % 4 != 0 {
return Err(GspError::Parse("miUINT32 indices misaligned".to_string()));
}
Ok(bytes
.chunks_exact(4)
.map(|c| u32::from_le_bytes([c[0], c[1], c[2], c[3]]) as usize)
.collect())
}
_ => Err(GspError::UnsupportedFormat(format!(
"unsupported MAT index dtype {dtype}"
))),
}
}
fn dense_from_col_major(nrows: usize, ncols: usize, data: &[f64]) -> Result<Array2<f64>, GspError> {
if nrows * ncols != data.len() {
return Err(GspError::Dimensions(format!(
"dense MAT shape {nrows}x{ncols} incompatible with {} values",
data.len()
)));
}
let mut out = Array2::<f64>::zeros((nrows, ncols));
for c in 0..ncols {
for r in 0..nrows {
out[[r, c]] = data[c * nrows + r];
}
}
Ok(out)
}
fn mat_parsed_to_dense(value: &MatParsedValue) -> Result<Array2<f64>, GspError> {
match value {
MatParsedValue::Dense {
nrows,
ncols,
col_major_data,
} => dense_from_col_major(*nrows, *ncols, col_major_data),
MatParsedValue::Sparse {
nrows,
ncols,
ir,
jc,
data,
} => {
let mut out = Array2::<f64>::zeros((*nrows, *ncols));
if jc.len() != ncols + 1 {
return Err(GspError::Dimensions(format!(
"sparse jc length {} incompatible with ncols {}",
jc.len(),
ncols
)));
}
for c in 0..*ncols {
let start = jc[c];
let end = jc[c + 1];
if end > ir.len() || end > data.len() {
return Err(GspError::Parse(
"sparse MAT index/value arrays out of bounds".to_string(),
));
}
for idx in start..end {
out[[ir[idx], c]] = data[idx];
}
}
Ok(out)
}
}
}
fn align_up(value: usize, multiple: usize) -> usize {
if value == 0 || multiple == 0 {
return value;
}
value.div_ceil(multiple) * multiple
}
fn variable_to_dense(var: &MatVariable) -> Result<Array2<f64>, GspError> {
match var {
MatVariable::NumericArray(na) => numeric_array_to_dense(na),
MatVariable::SparseArray(sa) => sparse_array_to_dense(sa),
_ => Err(GspError::Parse(
"unsupported MAT variable type for dense conversion".to_string(),
)),
}
}
fn numeric_array_to_dense(na: &matrw::NumericArray) -> Result<Array2<f64>, GspError> {
let data = matlab_type_to_f64(&na.value);
let nrows = na.dim.first().copied().unwrap_or(0);
let ncols = if na.dim.len() >= 2 {
na.dim[1..].iter().product()
} else {
1
};
if nrows * ncols != data.len() {
return Err(GspError::Dimensions(format!(
"numeric array dimensions {:?} incompatible with {} values",
na.dim,
data.len()
)));
}
let mut out = Array2::<f64>::zeros((nrows, ncols));
for c in 0..ncols {
for r in 0..nrows {
out[[r, c]] = data[c * nrows + r];
}
}
Ok(out)
}
fn sparse_array_to_dense(sa: &matrw::SparseArray) -> Result<Array2<f64>, GspError> {
let nrows = sa.dim.first().copied().unwrap_or(0);
let ncols = sa.dim.get(1).copied().unwrap_or(0);
let mut out = Array2::<f64>::zeros((nrows, ncols));
let data = matlab_type_to_f64(&sa.value);
for c in 0..ncols {
let start = sa.jc[c];
let end = sa.jc[c + 1];
for (&r, &val) in sa.ir[start..end].iter().zip(data[start..end].iter()) {
out[[r, c]] = val;
}
}
Ok(out)
}
fn matlab_type_to_f64(t: &MatlabType) -> Vec<f64> {
match t {
MatlabType::U8(v) => v.iter().map(|x| *x as f64).collect(),
MatlabType::I8(v) => v.iter().map(|x| *x as f64).collect(),
MatlabType::U16(v) => v.iter().map(|x| *x as f64).collect(),
MatlabType::I16(v) => v.iter().map(|x| *x as f64).collect(),
MatlabType::U32(v) => v.iter().map(|x| *x as f64).collect(),
MatlabType::I32(v) => v.iter().map(|x| *x as f64).collect(),
MatlabType::U64(v) => v.iter().map(|x| *x as f64).collect(),
MatlabType::I64(v) => v.iter().map(|x| *x as f64).collect(),
MatlabType::F32(v) => v.iter().map(|x| *x as f64).collect(),
MatlabType::F64(v) => v.clone(),
MatlabType::BOOL(v) => v.iter().map(|x| if *x { 1.0 } else { 0.0 }).collect(),
MatlabType::UTF8(v) => v.iter().map(|x| *x as u32 as f64).collect(),
MatlabType::UTF16(v) => v.iter().map(|x| *x as u32 as f64).collect(),
}
}
fn dense_to_sparse(d: ndarray::ArrayView2<'_, f64>) -> CsMat<f64> {
let mut tri = TriMat::<f64>::new((d.nrows(), d.ncols()));
for r in 0..d.nrows() {
for c in 0..d.ncols() {
let v = d[[r, c]];
if v != 0.0 {
tri.add_triplet(r, c, v);
}
}
}
tri.to_csc()
}