use std::io::{BufRead, BufReader, Read, Write};
use crate::error::IoError;
#[derive(Debug, Clone)]
pub struct Hat2Matrix {
pub names: Vec<String>,
pub distances: Vec<Vec<f64>>,
}
impl Hat2Matrix {
pub fn nseq(&self) -> usize {
self.names.len()
}
pub fn get(&self, i: usize, j: usize) -> f64 {
if i < j {
self.distances[i][j - i - 1]
} else if i > j {
self.distances[j][i - j - 1]
} else {
0.0
}
}
}
const VALS_PER_LINE: usize = 12;
pub fn read_hat2<R: Read>(reader: R) -> Result<Hat2Matrix, IoError> {
let reader = BufReader::new(reader);
let mut lines = reader.lines();
lines.next().ok_or(IoError::Hat2Format("empty file".into()))??;
let nseq_line = lines
.next()
.ok_or(IoError::Hat2Format("missing nseq line".into()))??;
let nseq: usize = nseq_line
.trim()
.parse()
.map_err(|_| IoError::Hat2Format(format!("invalid nseq: '{nseq_line}'")))?;
lines.next().ok_or(IoError::Hat2Format("missing max line".into()))??;
let mut names = Vec::with_capacity(nseq);
for _ in 0..nseq {
let line = lines
.next()
.ok_or(IoError::Hat2Format("truncated name section".into()))??;
let raw = if let Some(pos) = line.find(". ") {
&line[pos + 2..]
} else {
line.trim_start()
};
let name = raw.strip_prefix('=').unwrap_or(raw).to_string();
names.push(name);
}
let mut all_values: Vec<f64> = Vec::new();
for line_result in lines {
let line = line_result?;
if !line.trim().is_empty() {
for token in line.split_whitespace() {
let val: f64 = token
.parse()
.map_err(|_| IoError::Hat2Format(format!("invalid float: '{token}'")))?;
all_values.push(val);
}
}
}
let expected = nseq * (nseq - 1) / 2;
if all_values.len() != expected {
return Err(IoError::Hat2Format(format!(
"expected {} distance values, got {}",
expected,
all_values.len()
)));
}
let mut distances = Vec::with_capacity(nseq);
let mut offset = 0;
for i in 0..nseq {
let row_len = nseq - i - 1;
distances.push(all_values[offset..offset + row_len].to_vec());
offset += row_len;
}
Ok(Hat2Matrix { names, distances })
}
pub fn write_hat2<W: Write>(mat: &Hat2Matrix, writer: &mut W) -> Result<(), IoError> {
let nseq = mat.nseq();
let max_dist = mat
.distances
.iter()
.flat_map(|row| row.iter())
.cloned()
.fold(0.0_f64, f64::max);
writeln!(writer, " 1")?;
writeln!(writer, "{nseq:5}")?;
writeln!(writer, " {:>6.3}", max_dist * 2.5)?;
for (i, name) in mat.names.iter().enumerate() {
writeln!(writer, "{:>4}. ={name}", i + 1)?;
}
for i in 0..nseq.saturating_sub(1) {
let row = &mat.distances[i];
for (col, val) in row.iter().enumerate() {
write!(writer, "{val:6.3}")?;
if (col + 1) % VALS_PER_LINE == 0 || col == row.len() - 1 {
writeln!(writer)?;
}
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Cursor;
fn sample_matrix() -> Hat2Matrix {
Hat2Matrix {
names: vec!["seq1".into(), "seq2".into(), "seq3".into()],
distances: vec![
vec![0.123, 0.456], vec![0.789], ],
}
}
#[test]
fn roundtrip_hat2() {
let mat = sample_matrix();
let mut buf = Vec::new();
write_hat2(&mat, &mut buf).unwrap();
let parsed = read_hat2(Cursor::new(&buf)).unwrap();
assert_eq!(parsed.nseq(), 3);
assert_eq!(parsed.names, mat.names);
assert!((parsed.get(0, 1) - 0.123).abs() < 0.001);
assert!((parsed.get(0, 2) - 0.456).abs() < 0.001);
assert!((parsed.get(1, 2) - 0.789).abs() < 0.001);
assert!((parsed.get(2, 0) - 0.456).abs() < 0.001);
}
#[test]
fn self_distance_is_zero() {
let mat = sample_matrix();
assert_eq!(mat.get(0, 0), 0.0);
assert_eq!(mat.get(1, 1), 0.0);
}
}