use std::collections::BTreeSet;
use crate::bytes::LeCursor;
use crate::error::{Error, Result};
use super::header::HiCIndexItem;
use super::matrix::{Loc2D, MatrixMetadata};
use super::{HiCMode, Normalizations};
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct ContactRecord {
pub x_bin: i64,
pub y_bin: i64,
pub value: f32,
}
fn block_reaches_diagonal(
x_block: i64,
y_block: i64,
loc: &Loc2D,
block_bin_count: i64,
max_distance: Option<i64>,
) -> bool {
let Some(max_distance) = max_distance else {
return true;
};
let span = block_bin_count * loc.bin_size;
let (x0, y0) = (x_block * span, y_block * span);
let (x1, y1) = (x0 + span, y0 + span);
[(x0, y0), (x0, y1), (x1, y0), (x1, y1)]
.into_iter()
.map(|(x, y)| loc.distance_from_diagonal(x, y))
.min()
.is_some_and(|nearest| nearest <= max_distance)
}
pub fn block_numbers(
loc: &Loc2D,
meta: &MatrixMetadata,
max_distance: Option<i64>,
version: i64,
triangle: bool,
) -> Result<Vec<i64>> {
let (block_bin_count, block_column_count) = (meta.block_bin_count, meta.block_column_count);
if block_bin_count <= 0 || block_column_count <= 0 {
return Err(Error::invalid(format!(
"matrix declares {block_bin_count} bins and {block_column_count} columns per block"
)));
}
let mut blocks = BTreeSet::new();
let x_bin_end = loc.x.bin_end.min(loc.x.chr.size / loc.bin_size + 1);
let y_bin_end = loc.y.bin_end.min(loc.y.chr.size / loc.bin_size + 1);
if version > 8 && loc.is_intra() {
let lower_pad = (loc.x.bin_start + loc.y.bin_start) / 2 / block_bin_count;
let higher_pad = (x_bin_end + y_bin_end) / 2 / block_bin_count + 1;
let depth_of = |a: i64, b: i64| -> i64 {
(1.0 + (a - b).abs() as f64 / std::f64::consts::SQRT_2 / block_bin_count as f64).log2()
as i64
};
let nearer = depth_of(loc.x.bin_start, y_bin_end);
let further = depth_of(x_bin_end, loc.y.bin_start);
let mut nearer_depth = nearer.min(further);
if (loc.x.bin_start > y_bin_end && x_bin_end < loc.y.bin_start)
|| (x_bin_end > loc.y.bin_start && loc.x.bin_start < y_bin_end)
{
nearer_depth = 0;
}
let further_depth = nearer.max(further) + 1;
for depth in nearer_depth..=further_depth {
for pad in lower_pad..=higher_pad {
blocks.insert(depth * block_column_count + pad);
}
}
} else {
let col1 = loc.x.bin_start / block_bin_count;
let col2 = (x_bin_end - 1).max(col1 * block_bin_count) / block_bin_count;
let row1 = loc.y.bin_start / block_bin_count;
let row2 = (y_bin_end - 1).max(row1 * block_bin_count) / block_bin_count;
for row in row1..=row2 {
for col in col1..=col2 {
if !block_reaches_diagonal(col, row, loc, block_bin_count, max_distance) {
continue;
}
blocks.insert(row * block_column_count + col);
}
}
if loc.is_intra() && !triangle {
for row in col1..=col2 {
for col in row1..=row2 {
if !block_reaches_diagonal(col, row, loc, block_bin_count, max_distance) {
continue;
}
blocks.insert(row * block_column_count + col);
}
}
}
}
Ok(blocks
.into_iter()
.filter(|n| meta.blocks.contains_key(n))
.collect())
}
pub struct RecordContext<'a> {
pub loc: &'a Loc2D,
pub normalization: &'a str,
pub mode: HiCMode,
pub vectors: &'a Normalizations,
pub average_value: f32,
pub min_distance: Option<i64>,
pub max_distance: Option<i64>,
}
pub fn process_record(record: &mut ContactRecord, ctx: &RecordContext<'_>) -> bool {
let loc = ctx.loc;
let x = record.x_bin * loc.bin_size;
let y = record.y_bin * loc.bin_size;
if ctx.min_distance.is_some() || ctx.max_distance.is_some() {
let distance = loc.distance_from_diagonal(x, y);
if ctx.min_distance.is_some_and(|min| distance < min) {
return false;
}
if ctx.max_distance.is_some_and(|max| distance > max) {
return false;
}
}
let inside = (x >= loc.x.binned_start
&& x <= loc.x.binned_end
&& y >= loc.y.binned_start
&& y <= loc.y.binned_end)
|| (loc.is_intra()
&& y >= loc.x.binned_start
&& y <= loc.x.binned_end
&& x >= loc.y.binned_start
&& x <= loc.y.binned_end);
if !inside {
return false;
}
if ctx.normalization != "none" {
let x_norm = ctx.vectors.x.get(record.x_bin.max(0) as usize).copied();
let y_norm = ctx.vectors.y.get(record.y_bin.max(0) as usize).copied();
match (x_norm, y_norm) {
(Some(a), Some(b)) if record.x_bin >= 0 && record.y_bin >= 0 => record.value /= a * b,
_ => {
record.value = f32::NAN;
return true;
}
}
}
if matches!(ctx.mode, HiCMode::Oe | HiCMode::Expected) {
let expected = if loc.is_intra() {
if ctx.vectors.expected.is_empty() {
record.value = f32::NAN;
return true;
}
let i = ((y - x).abs() / loc.bin_size).max(0) as usize;
ctx.vectors.expected[i.min(ctx.vectors.expected.len() - 1)]
} else {
ctx.average_value
};
record.value = match ctx.mode {
HiCMode::Oe => record.value / expected,
_ => expected,
};
}
if !record.value.is_finite() {
record.value = f32::NAN;
}
true
}
pub fn read_block(
raw: bytes::Bytes,
version: i64,
block: HiCIndexItem,
ctx: &RecordContext<'_>,
path: &str,
) -> Result<Vec<ContactRecord>> {
let buffer = decompress(raw, path, block.position)?;
let mut c = LeCursor::new(&buffer, block.position, path);
let record_count = c.read_i32()? as i64;
if record_count < 0 {
return Err(Error::corrupt(
path,
block.position,
"hic block declares a negative record count",
));
}
let mut records = Vec::new();
let keep = |record: &mut ContactRecord, out: &mut Vec<ContactRecord>| {
if process_record(record, ctx) {
out.push(*record);
}
};
if version < 7 {
for _ in 0..record_count {
let mut record = ContactRecord {
x_bin: c.read_i32()? as i64,
y_bin: c.read_i32()? as i64,
value: c.read_f32()?,
};
keep(&mut record, &mut records);
}
records.shrink_to_fit();
return Ok(records);
}
let bin_column_offset = c.read_i32()? as i64;
let bin_row_offset = c.read_i32()? as i64;
let use_float = c.read_u8()? == 1;
let (use_int_x, use_int_y) = if version > 8 {
let x = c.read_u8()? == 1;
let y = c.read_u8()? == 1;
(x, y)
} else {
(false, false)
};
let matrix_type = c.read_u8()?;
let x_width = if use_int_x { 4 } else { 2 };
records.reserve((record_count as usize).min(c.remaining() / if use_float { 4 } else { 2 }));
let read_x = |c: &mut LeCursor<'_>| -> Result<i64> {
Ok(if use_int_x {
c.read_i32()? as i64
} else {
c.read_i16()? as i64
})
};
let read_y = |c: &mut LeCursor<'_>| -> Result<i64> {
Ok(if use_int_y {
c.read_i32()? as i64
} else {
c.read_i16()? as i64
})
};
match matrix_type {
1 => {
let row_count = read_y(&mut c)?;
for _ in 0..row_count.max(0) {
let row_number = read_y(&mut c)?;
let col_count = read_x(&mut c)?;
let y_bin = bin_row_offset + row_number;
let needed = (col_count.max(0) as usize)
.saturating_mul(x_width + if use_float { 4 } else { 2 });
if needed > c.remaining() {
return Err(Error::corrupt(
path,
block.position,
format!(
"hic block row declares {col_count} columns, which do not fit \
the {} bytes left in the block",
c.remaining()
),
));
}
for _ in 0..col_count.max(0) {
let col_number = read_x(&mut c)?;
let value = if use_float {
c.read_f32()?
} else {
c.read_i16()? as f32
};
let mut record = ContactRecord {
x_bin: bin_column_offset + col_number,
y_bin,
value,
};
keep(&mut record, &mut records);
}
}
}
2 => {
let count = c.read_i32()? as i64;
let width = c.read_i16()? as i64;
if count < 0 || width <= 0 {
return Err(Error::corrupt(
path,
block.position,
format!("hic dense block declares {count} values of width {width}"),
));
}
for i in 0..count {
let row = i / width;
let col = i - row * width;
let value = if use_float {
let v = c.read_f32()?;
if v.is_nan() {
continue;
}
v
} else {
let v = c.read_i16()?;
if v == -32768 {
continue;
}
v as f32
};
let mut record = ContactRecord {
x_bin: bin_column_offset + col,
y_bin: bin_row_offset + row,
value,
};
keep(&mut record, &mut records);
}
}
other => {
return Err(Error::corrupt(
path,
block.position,
format!("matrix type {other} invalid"),
))
}
}
records.shrink_to_fit();
Ok(records)
}
const MAX_INFLATED_SIZE: usize = 1 << 30;
const MAX_INFLATE_RESERVE: usize = 1 << 20;
fn decompress(raw: bytes::Bytes, path: &str, at: u64) -> Result<bytes::Bytes> {
use std::io::Read;
let mut out = Vec::with_capacity(raw.len().saturating_mul(4).min(MAX_INFLATE_RESERVE));
flate2::read::ZlibDecoder::new(&raw[..])
.take(MAX_INFLATED_SIZE as u64 + 1)
.read_to_end(&mut out)
.map_err(|e| Error::corrupt(path, at, format!("could not inflate the hic block: {e}")))?;
if out.len() > MAX_INFLATED_SIZE {
return Err(Error::corrupt(
path,
at,
format!("the inflated hic block exceeds the limit ({MAX_INFLATED_SIZE})"),
));
}
Ok(bytes::Bytes::from(out))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::genomic::ChrMap;
use crate::hic::matrix::parse_loc2d;
use std::io::Write;
fn map() -> ChrMap {
ChrMap::from_indexed_entries([
("chr1".to_string(), 1_000_000, 0),
("chr2".to_string(), 500_000, 1),
])
}
fn loc(ids: &[&str], starts: &[i64], ends: &[i64]) -> Loc2D {
let ids: Vec<String> = ids.iter().map(|s| s.to_string()).collect();
parse_loc2d(&map(), &[5000], &ids, starts, ends, Some(5000), None, false).unwrap()
}
fn ctx<'a>(loc: &'a Loc2D, vectors: &'a Normalizations) -> RecordContext<'a> {
RecordContext {
loc,
normalization: "none",
mode: HiCMode::Observed,
vectors,
average_value: f32::NAN,
min_distance: None,
max_distance: None,
}
}
fn zlib(payload: &[u8]) -> bytes::Bytes {
let mut encoder = flate2::write::ZlibEncoder::new(Vec::new(), flate2::Compression::new(6));
encoder.write_all(payload).unwrap();
bytes::Bytes::from(encoder.finish().unwrap())
}
fn sparse_block() -> Vec<u8> {
let mut b = 2i32.to_le_bytes().to_vec(); b.extend_from_slice(&0i32.to_le_bytes()); b.extend_from_slice(&0i32.to_le_bytes()); b.push(1); b.push(1); b.extend_from_slice(&1i16.to_le_bytes()); b.extend_from_slice(&2i16.to_le_bytes()); b.extend_from_slice(&2i16.to_le_bytes()); for (col, value) in [(3i16, 1.5f32), (4, 2.5)] {
b.extend_from_slice(&col.to_le_bytes());
b.extend_from_slice(&value.to_le_bytes());
}
b
}
fn item() -> HiCIndexItem {
HiCIndexItem {
position: 0,
size: 0,
}
}
#[test]
fn a_v8_sparse_block_decodes_to_its_contacts() {
let l = loc(&["chr1"], &[0], &[100_000]);
let v = Normalizations::default();
let records = read_block(zlib(&sparse_block()), 8, item(), &ctx(&l, &v), "test").unwrap();
assert_eq!(
records,
[
ContactRecord {
x_bin: 3,
y_bin: 2,
value: 1.5
},
ContactRecord {
x_bin: 4,
y_bin: 2,
value: 2.5
},
]
);
}
#[test]
fn a_v9_block_carries_two_more_flag_bytes() {
let mut b = 1i32.to_le_bytes().to_vec();
b.extend_from_slice(&0i32.to_le_bytes());
b.extend_from_slice(&0i32.to_le_bytes());
b.push(1); b.push(1); b.push(0); b.push(1); b.extend_from_slice(&1i16.to_le_bytes()); b.extend_from_slice(&5i16.to_le_bytes()); b.extend_from_slice(&1i32.to_le_bytes()); b.extend_from_slice(&7i32.to_le_bytes()); b.extend_from_slice(&9.0f32.to_le_bytes());
let l = loc(&["chr1"], &[0], &[100_000]);
let v = Normalizations::default();
let records = read_block(zlib(&b), 9, item(), &ctx(&l, &v), "test").unwrap();
assert_eq!(
records,
[ContactRecord {
x_bin: 7,
y_bin: 5,
value: 9.0
}]
);
}
#[test]
fn a_dense_block_skips_its_sentinels() {
let mut b = 4i32.to_le_bytes().to_vec();
b.extend_from_slice(&0i32.to_le_bytes());
b.extend_from_slice(&0i32.to_le_bytes());
b.push(0); b.push(2); b.extend_from_slice(&4i32.to_le_bytes()); b.extend_from_slice(&2i16.to_le_bytes()); for value in [1i16, -32768, 3, 4] {
b.extend_from_slice(&value.to_le_bytes());
}
let l = loc(&["chr1"], &[0], &[100_000]);
let v = Normalizations::default();
let records = read_block(zlib(&b), 8, item(), &ctx(&l, &v), "test").unwrap();
assert_eq!(records.len(), 3, "the sentinel is not a contact");
assert_eq!(
records[0],
ContactRecord {
x_bin: 0,
y_bin: 0,
value: 1.0
}
);
assert_eq!(
records[1],
ContactRecord {
x_bin: 0,
y_bin: 1,
value: 3.0
}
);
}
#[test]
fn an_unknown_matrix_type_is_refused() {
let mut b = sparse_block();
b[13] = 7; let l = loc(&["chr1"], &[0], &[100_000]);
let v = Normalizations::default();
let err = read_block(zlib(&b), 8, item(), &ctx(&l, &v), "test")
.unwrap_err()
.to_string();
assert!(err.contains("matrix type 7 invalid"), "{err}");
}
#[test]
fn a_row_declaring_more_columns_than_the_block_holds_is_refused() {
let mut b = sparse_block();
let at = 4 + 4 + 4 + 1 + 1 + 2 + 2;
b[at..at + 2].copy_from_slice(&1000i16.to_le_bytes());
let l = loc(&["chr1"], &[0], &[100_000]);
let v = Normalizations::default();
let err = read_block(zlib(&b), 8, item(), &ctx(&l, &v), "test")
.unwrap_err()
.to_string();
assert!(err.contains("do not fit"), "{err}");
}
#[test]
fn a_contact_outside_the_window_is_dropped() {
let l = loc(&["chr1"], &[0], &[10_000]);
let v = Normalizations::default();
let records = read_block(zlib(&sparse_block()), 8, item(), &ctx(&l, &v), "test").unwrap();
assert!(records.is_empty(), "{records:?}");
}
#[test]
fn a_contact_the_file_cannot_value_comes_back_nan() {
let l = loc(&["chr1"], &[0], &[100_000]);
let vectors = Normalizations {
x: std::sync::Arc::new(vec![1.0, 1.0]),
y: std::sync::Arc::new(vec![1.0, 1.0]),
expected: std::sync::Arc::new(Vec::new()),
};
let mut context = ctx(&l, &vectors);
context.normalization = "kr";
let records = read_block(zlib(&sparse_block()), 8, item(), &context, "test").unwrap();
assert_eq!(records.len(), 2, "kept, not dropped");
assert!(records.iter().all(|r| r.value.is_nan()));
}
#[test]
fn a_zero_normalization_factor_also_gives_nan_rather_than_infinity() {
let l = loc(&["chr1"], &[0], &[100_000]);
let vectors = Normalizations {
x: std::sync::Arc::new(vec![1.0; 10]),
y: std::sync::Arc::new(vec![0.0; 10]),
expected: std::sync::Arc::new(Vec::new()),
};
let mut context = ctx(&l, &vectors);
context.normalization = "kr";
let records = read_block(zlib(&sparse_block()), 8, item(), &context, "test").unwrap();
assert!(records.iter().all(|r| r.value.is_nan()));
}
#[test]
fn oe_divides_by_the_expected_value_at_that_distance() {
let l = loc(&["chr1"], &[0], &[100_000]);
let vectors = Normalizations {
x: std::sync::Arc::new(Vec::new()),
y: std::sync::Arc::new(Vec::new()),
expected: std::sync::Arc::new(vec![10.0, 5.0, 2.0]),
};
let mut context = ctx(&l, &vectors);
context.mode = HiCMode::Oe;
let records = read_block(zlib(&sparse_block()), 8, item(), &context, "test").unwrap();
assert_eq!(records[0].value, 1.5 / 5.0);
assert_eq!(records[1].value, 2.5 / 2.0);
}
#[test]
fn distance_filters_drop_what_falls_outside_them() {
let l = loc(&["chr1"], &[0], &[100_000]);
let v = Normalizations::default();
let mut context = ctx(&l, &v);
context.min_distance = Some(7_000);
let records = read_block(zlib(&sparse_block()), 8, item(), &context, "test").unwrap();
assert_eq!(records.len(), 1);
assert_eq!(records[0].x_bin, 4);
let mut context = ctx(&l, &v);
context.max_distance = Some(7_000);
let records = read_block(zlib(&sparse_block()), 8, item(), &context, "test").unwrap();
assert_eq!(records.len(), 1);
assert_eq!(records[0].x_bin, 3);
}
#[test]
fn garbage_where_a_deflate_stream_should_be_is_corrupt() {
let l = loc(&["chr1"], &[0], &[100_000]);
let v = Normalizations::default();
let err = read_block(
bytes::Bytes::from(vec![9u8; 64]),
8,
item(),
&ctx(&l, &v),
"test",
)
.unwrap_err()
.to_string();
assert!(err.contains("could not inflate"), "{err}");
}
}