use crate::coo::CooMatrix;
use crate::csr::CsrMatrix;
use num_traits::ToPrimitive;
use oxiblas_core::scalar::{Field, Real, Scalar};
use std::io::{BufRead, BufReader, Write};
use std::path::Path;
const MAX_PREALLOC_ENTRIES: usize = 1 << 20;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum MtxError {
InvalidHeader(String),
InvalidData(String),
UnsupportedType(String),
IoError(String),
ParseError(String),
MissingSizeLine,
IndexOutOfBounds {
row: usize,
col: usize,
nrows: usize,
ncols: usize,
},
}
impl core::fmt::Display for MtxError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::InvalidHeader(s) => write!(f, "Invalid Matrix Market header: {s}"),
Self::InvalidData(s) => write!(f, "Invalid data: {s}"),
Self::UnsupportedType(s) => write!(f, "Unsupported matrix type: {s}"),
Self::IoError(s) => write!(f, "I/O error: {s}"),
Self::ParseError(s) => write!(f, "Parse error: {s}"),
Self::MissingSizeLine => write!(f, "Missing size line"),
Self::IndexOutOfBounds {
row,
col,
nrows,
ncols,
} => {
write!(
f,
"Index ({row}, {col}) out of bounds for {nrows}×{ncols} matrix"
)
}
}
}
}
impl std::error::Error for MtxError {}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MtxObject {
Matrix,
Vector,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MtxFormat {
Coordinate,
Array,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MtxField {
Real,
Complex,
Pattern,
Integer,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MtxSymmetry {
General,
Symmetric,
SkewSymmetric,
Hermitian,
}
#[derive(Debug, Clone)]
pub struct MtxHeader {
pub object: MtxObject,
pub format: MtxFormat,
pub field: MtxField,
pub symmetry: MtxSymmetry,
pub nrows: usize,
pub ncols: usize,
pub nnz: usize,
pub comments: Vec<String>,
}
fn parse_header_line(
line: &str,
) -> Result<(MtxObject, MtxFormat, MtxField, MtxSymmetry), MtxError> {
let line = line.to_lowercase();
if !line.starts_with("%%matrixmarket") {
return Err(MtxError::InvalidHeader(
"Header must start with %%MatrixMarket".to_string(),
));
}
let parts: Vec<&str> = line.split_whitespace().collect();
if parts.len() < 5 {
return Err(MtxError::InvalidHeader(
"Header must have 5 parts".to_string(),
));
}
let object = match parts[1] {
"matrix" => MtxObject::Matrix,
"vector" => MtxObject::Vector,
other => {
return Err(MtxError::UnsupportedType(format!(
"Unknown object type: {other}"
)));
}
};
let format = match parts[2] {
"coordinate" => MtxFormat::Coordinate,
"array" => MtxFormat::Array,
other => {
return Err(MtxError::UnsupportedType(format!(
"Unknown format: {other}"
)));
}
};
let field = match parts[3] {
"real" => MtxField::Real,
"double" => MtxField::Real,
"complex" => MtxField::Complex,
"pattern" => MtxField::Pattern,
"integer" => MtxField::Integer,
other => {
return Err(MtxError::UnsupportedType(format!(
"Unknown field type: {other}"
)));
}
};
let symmetry = match parts[4] {
"general" => MtxSymmetry::General,
"symmetric" => MtxSymmetry::Symmetric,
"skew-symmetric" => MtxSymmetry::SkewSymmetric,
"hermitian" => MtxSymmetry::Hermitian,
other => {
return Err(MtxError::UnsupportedType(format!(
"Unknown symmetry: {other}"
)));
}
};
Ok((object, format, field, symmetry))
}
pub fn read_header<R: BufRead>(reader: &mut R) -> Result<MtxHeader, MtxError> {
let mut line = String::new();
reader
.read_line(&mut line)
.map_err(|e| MtxError::IoError(e.to_string()))?;
let (object, format, field, symmetry) = parse_header_line(line.trim())?;
let mut comments = Vec::new();
loop {
line.clear();
reader
.read_line(&mut line)
.map_err(|e| MtxError::IoError(e.to_string()))?;
if line.is_empty() {
return Err(MtxError::MissingSizeLine);
}
let trimmed = line.trim();
if trimmed.starts_with('%') {
comments.push(trimmed[1..].trim().to_string());
} else {
break;
}
}
let size_parts: Vec<&str> = line.split_whitespace().collect();
let (nrows, ncols, nnz) = match format {
MtxFormat::Coordinate => {
if size_parts.len() < 3 {
return Err(MtxError::InvalidData(
"Coordinate size line must have 3 values".to_string(),
));
}
let nrows = size_parts[0]
.parse::<usize>()
.map_err(|_| MtxError::ParseError("Invalid nrows".to_string()))?;
let ncols = size_parts[1]
.parse::<usize>()
.map_err(|_| MtxError::ParseError("Invalid ncols".to_string()))?;
let nnz = size_parts[2]
.parse::<usize>()
.map_err(|_| MtxError::ParseError("Invalid nnz".to_string()))?;
let max_nnz = nrows.saturating_mul(ncols);
if nnz > max_nnz {
return Err(MtxError::InvalidData(format!(
"nnz ({nnz}) exceeds the number of cells in a {nrows}x{ncols} \
matrix ({max_nnz})"
)));
}
(nrows, ncols, nnz)
}
MtxFormat::Array => {
if size_parts.len() < 2 {
return Err(MtxError::InvalidData(
"Array size line must have 2 values".to_string(),
));
}
let nrows = size_parts[0]
.parse::<usize>()
.map_err(|_| MtxError::ParseError("Invalid nrows".to_string()))?;
let ncols = size_parts[1]
.parse::<usize>()
.map_err(|_| MtxError::ParseError("Invalid ncols".to_string()))?;
let nnz = nrows.checked_mul(ncols).ok_or_else(|| {
MtxError::InvalidData(format!("array dimensions overflow usize: {nrows}x{ncols}"))
})?;
(nrows, ncols, nnz)
}
};
Ok(MtxHeader {
object,
format,
field,
symmetry,
nrows,
ncols,
nnz,
comments,
})
}
pub fn read_matrix_market<T: Scalar<Real = T> + Clone + Field + Real, P: AsRef<Path>>(
path: P,
) -> Result<CsrMatrix<T>, MtxError> {
let file = std::fs::File::open(path).map_err(|e| MtxError::IoError(e.to_string()))?;
let mut reader = BufReader::new(file);
read_matrix_market_from_reader(&mut reader)
}
pub fn read_matrix_market_from_reader<T: Scalar<Real = T> + Clone + Field + Real, R: BufRead>(
reader: &mut R,
) -> Result<CsrMatrix<T>, MtxError> {
let header = read_header(reader)?;
if header.format != MtxFormat::Coordinate {
return Err(MtxError::UnsupportedType(
"Only coordinate format is supported".to_string(),
));
}
if header.field == MtxField::Complex {
return Err(MtxError::UnsupportedType(
"Complex matrices not supported for real type".to_string(),
));
}
let prealloc = header.nnz.min(MAX_PREALLOC_ENTRIES);
let mut rows = Vec::with_capacity(prealloc);
let mut cols = Vec::with_capacity(prealloc);
let mut vals = Vec::with_capacity(prealloc);
for line_result in reader.lines() {
let line = line_result.map_err(|e| MtxError::IoError(e.to_string()))?;
let trimmed = line.trim();
if trimmed.is_empty() || trimmed.starts_with('%') {
continue;
}
let parts: Vec<&str> = trimmed.split_whitespace().collect();
if parts.len() < 2 {
return Err(MtxError::InvalidData(format!(
"Invalid data line: {trimmed}"
)));
}
let row: usize = parts[0]
.parse()
.map_err(|_| MtxError::ParseError(format!("Invalid row: {}", parts[0])))?;
let col: usize = parts[1]
.parse()
.map_err(|_| MtxError::ParseError(format!("Invalid col: {}", parts[1])))?;
if row == 0 || col == 0 {
return Err(MtxError::IndexOutOfBounds {
row,
col,
nrows: header.nrows,
ncols: header.ncols,
});
}
let row = row - 1;
let col = col - 1;
if row >= header.nrows || col >= header.ncols {
return Err(MtxError::IndexOutOfBounds {
row: row + 1,
col: col + 1,
nrows: header.nrows,
ncols: header.ncols,
});
}
let val = if header.field == MtxField::Pattern {
T::one()
} else {
if parts.len() < 3 {
return Err(MtxError::InvalidData(format!(
"Missing value on line: {trimmed}"
)));
}
parts[2]
.parse::<f64>()
.map_err(|_| MtxError::ParseError(format!("Invalid value: {}", parts[2])))
.and_then(|v| {
T::from_f64(v)
.ok_or_else(|| MtxError::ParseError(format!("Cannot convert value: {v}")))
})?
};
rows.push(row);
cols.push(col);
vals.push(val.clone());
if row != col {
match header.symmetry {
MtxSymmetry::Symmetric => {
rows.push(col);
cols.push(row);
vals.push(val);
}
MtxSymmetry::SkewSymmetric => {
rows.push(col);
cols.push(row);
vals.push(T::zero() - val);
}
MtxSymmetry::Hermitian => {
rows.push(col);
cols.push(row);
vals.push(val);
}
MtxSymmetry::General => {}
}
}
}
let coo = CooMatrix::new(header.nrows, header.ncols, rows, cols, vals)
.map_err(|e| MtxError::InvalidData(format!("Failed to create COO matrix: {e:?}")))?;
Ok(crate::convert::coo_to_csr(&coo))
}
pub fn read_matrix_market_coo<T: Scalar<Real = T> + Clone + Field + Real, P: AsRef<Path>>(
path: P,
) -> Result<CooMatrix<T>, MtxError> {
let csr: CsrMatrix<T> = read_matrix_market(path)?;
Ok(crate::convert::csr_to_coo(&csr))
}
pub fn write_matrix_market<T: Scalar + Clone + Field + ToPrimitive, P: AsRef<Path>>(
csr: &CsrMatrix<T>,
path: P,
comment: Option<&str>,
) -> Result<(), MtxError> {
let file = std::fs::File::create(path).map_err(|e| MtxError::IoError(e.to_string()))?;
let mut writer = std::io::BufWriter::new(file);
write_matrix_market_to_writer(csr, &mut writer, comment)
}
pub fn write_matrix_market_to_writer<T: Scalar + Clone + Field + ToPrimitive, W: Write>(
csr: &CsrMatrix<T>,
writer: &mut W,
comment: Option<&str>,
) -> Result<(), MtxError> {
let eps = <T as Scalar>::epsilon();
let mut nnz = 0;
for (_, _, val) in csr.iter() {
if Scalar::abs(val.clone()) > eps {
nnz += 1;
}
}
writeln!(writer, "%%MatrixMarket matrix coordinate real general")
.map_err(|e| MtxError::IoError(e.to_string()))?;
if let Some(c) = comment {
for line in c.lines() {
writeln!(writer, "% {line}").map_err(|e| MtxError::IoError(e.to_string()))?;
}
}
writeln!(writer, "{} {} {}", csr.nrows(), csr.ncols(), nnz)
.map_err(|e| MtxError::IoError(e.to_string()))?;
for (row, col, val) in csr.iter() {
if Scalar::abs(val.clone()) > eps {
let f = val.to_f64().unwrap_or(0.0);
writeln!(writer, "{} {} {}", row + 1, col + 1, f)
.map_err(|e| MtxError::IoError(e.to_string()))?;
}
}
Ok(())
}
pub fn write_matrix_market_symmetric<T: Scalar + Clone + Field + ToPrimitive, P: AsRef<Path>>(
csr: &CsrMatrix<T>,
path: P,
comment: Option<&str>,
) -> Result<(), MtxError> {
let file = std::fs::File::create(path).map_err(|e| MtxError::IoError(e.to_string()))?;
let mut writer = std::io::BufWriter::new(file);
let eps = <T as Scalar>::epsilon();
let mut nnz = 0;
for (row, col, val) in csr.iter() {
if row >= col && Scalar::abs(val.clone()) > eps {
nnz += 1;
}
}
writeln!(writer, "%%MatrixMarket matrix coordinate real symmetric")
.map_err(|e| MtxError::IoError(e.to_string()))?;
if let Some(c) = comment {
for line in c.lines() {
writeln!(writer, "% {line}").map_err(|e| MtxError::IoError(e.to_string()))?;
}
}
writeln!(writer, "{} {} {}", csr.nrows(), csr.ncols(), nnz)
.map_err(|e| MtxError::IoError(e.to_string()))?;
for (row, col, val) in csr.iter() {
if row >= col && Scalar::abs(val.clone()) > eps {
let f = val.to_f64().unwrap_or(0.0);
writeln!(writer, "{} {} {}", row + 1, col + 1, f)
.map_err(|e| MtxError::IoError(e.to_string()))?;
}
}
Ok(())
}
pub fn read_matrix_market_str<T: Scalar<Real = T> + Clone + Field + Real>(
s: &str,
) -> Result<CsrMatrix<T>, MtxError> {
let mut reader = BufReader::new(s.as_bytes());
read_matrix_market_from_reader(&mut reader)
}
pub fn write_matrix_market_str<T: Scalar + Clone + Field + ToPrimitive>(
csr: &CsrMatrix<T>,
comment: Option<&str>,
) -> Result<String, MtxError> {
let mut buf = Vec::new();
write_matrix_market_to_writer(csr, &mut buf, comment)?;
String::from_utf8(buf).map_err(|e| MtxError::IoError(e.to_string()))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_header() {
let (obj, fmt, field, sym) =
parse_header_line("%%MatrixMarket matrix coordinate real general").unwrap();
assert_eq!(obj, MtxObject::Matrix);
assert_eq!(fmt, MtxFormat::Coordinate);
assert_eq!(field, MtxField::Real);
assert_eq!(sym, MtxSymmetry::General);
}
#[test]
fn test_parse_header_symmetric() {
let (_, _, _, sym) =
parse_header_line("%%MatrixMarket matrix coordinate real symmetric").unwrap();
assert_eq!(sym, MtxSymmetry::Symmetric);
}
#[test]
fn test_read_simple_matrix() {
let mtx = r#"%%MatrixMarket matrix coordinate real general
% A simple test matrix
3 3 5
1 1 1.0
1 3 2.0
2 2 3.0
3 1 4.0
3 3 5.0
"#;
let csr: CsrMatrix<f64> = read_matrix_market_str(mtx).unwrap();
assert_eq!(csr.nrows(), 3);
assert_eq!(csr.ncols(), 3);
assert_eq!(csr.nnz(), 5);
assert_eq!(csr.get(0, 0), Some(&1.0));
assert_eq!(csr.get(0, 2), Some(&2.0));
assert_eq!(csr.get(1, 1), Some(&3.0));
assert_eq!(csr.get(2, 0), Some(&4.0));
assert_eq!(csr.get(2, 2), Some(&5.0));
}
#[test]
fn test_read_symmetric_matrix() {
let mtx = r#"%%MatrixMarket matrix coordinate real symmetric
3 3 4
1 1 1.0
2 1 2.0
2 2 3.0
3 3 4.0
"#;
let csr: CsrMatrix<f64> = read_matrix_market_str(mtx).unwrap();
assert_eq!(csr.nrows(), 3);
assert_eq!(csr.ncols(), 3);
assert_eq!(csr.get(0, 0), Some(&1.0));
assert_eq!(csr.get(1, 0), Some(&2.0));
assert_eq!(csr.get(0, 1), Some(&2.0)); assert_eq!(csr.get(1, 1), Some(&3.0));
assert_eq!(csr.get(2, 2), Some(&4.0));
}
#[test]
fn test_read_pattern_matrix() {
let mtx = r#"%%MatrixMarket matrix coordinate pattern general
2 2 2
1 1
2 2
"#;
let csr: CsrMatrix<f64> = read_matrix_market_str(mtx).unwrap();
assert_eq!(csr.get(0, 0), Some(&1.0));
assert_eq!(csr.get(1, 1), Some(&1.0));
}
#[test]
fn test_write_read_roundtrip() {
let values = vec![1.0f64, 2.0, 3.0, 4.0, 5.0];
let col_indices = vec![0, 2, 1, 0, 2];
let row_ptrs = vec![0, 2, 3, 5];
let csr = CsrMatrix::new(3, 3, row_ptrs, col_indices, values).unwrap();
let mtx_str = write_matrix_market_str(&csr, Some("Test matrix")).unwrap();
let csr2: CsrMatrix<f64> = read_matrix_market_str(&mtx_str).unwrap();
assert_eq!(csr.nrows(), csr2.nrows());
assert_eq!(csr.ncols(), csr2.ncols());
assert_eq!(csr.nnz(), csr2.nnz());
for row in 0..3 {
for col in 0..3 {
let v1 = csr.get(row, col).cloned().unwrap_or(0.0);
let v2 = csr2.get(row, col).cloned().unwrap_or(0.0);
assert!((v1 - v2).abs() < 1e-10);
}
}
}
#[test]
fn test_header_parsing_error() {
let result = parse_header_line("invalid header");
assert!(result.is_err());
}
#[test]
fn test_index_error() {
let mtx = r#"%%MatrixMarket matrix coordinate real general
2 2 1
3 1 1.0
"#;
let result: Result<CsrMatrix<f64>, _> = read_matrix_market_str(mtx);
assert!(result.is_err());
}
#[test]
fn test_absurd_nnz_is_rejected_not_preallocated() {
let mtx = "%%MatrixMarket matrix coordinate real general\n3 3 18446744073709551615\n";
let result: Result<CsrMatrix<f64>, _> = read_matrix_market_str(mtx);
match result {
Err(MtxError::InvalidData(msg)) => assert!(msg.contains("nnz")),
Err(other) => panic!("expected InvalidData about nnz, got {other:?}"),
Ok(_) => panic!("a 3x3 matrix claiming 2^64-1 non-zeros was accepted"),
}
}
#[test]
fn test_nnz_above_cell_count_is_rejected() {
let mtx = "%%MatrixMarket matrix coordinate real general\n2 2 5\n";
let result: Result<CsrMatrix<f64>, _> = read_matrix_market_str(mtx);
assert!(matches!(result, Err(MtxError::InvalidData(_))));
}
#[test]
fn test_large_but_legal_nnz_still_parses() {
let mtx = "%%MatrixMarket matrix coordinate real general\n3 3 9\n1 1 1.0\n2 2 2.0\n";
let csr: CsrMatrix<f64> = read_matrix_market_str(mtx).expect("valid file");
assert_eq!(csr.nnz(), 2);
assert_eq!(csr.get(0, 0), Some(&1.0));
assert_eq!(csr.get(1, 1), Some(&2.0));
}
#[test]
fn test_array_header_dimension_overflow_is_rejected() {
let mtx = "%%MatrixMarket matrix array real general\n4294967296 4294967296\n";
let mut reader = std::io::Cursor::new(mtx.as_bytes());
match read_header(&mut reader) {
Err(MtxError::InvalidData(msg)) => assert!(msg.contains("overflow")),
Err(other) => panic!("expected InvalidData about overflow, got {other:?}"),
Ok(h) => panic!("overflowing array dimensions accepted: nnz={}", h.nnz),
}
}
#[test]
fn test_array_header_sane_dimensions_still_parse() {
let mtx = "%%MatrixMarket matrix array real general\n4 5\n";
let mut reader = std::io::Cursor::new(mtx.as_bytes());
let header = read_header(&mut reader).expect("sane array header");
assert_eq!((header.nrows, header.ncols, header.nnz), (4, 5, 20));
}
}