use std::ops::Range;
use bytes::Bytes;
use crate::bbi::header::{DATA_TREE_HEADER_SIZE, DATA_TREE_MAGIC};
use crate::bytes::LeCursor;
use crate::error::{Error, Result};
use crate::genomic::{IndexedLoc, LocBatch};
use crate::progress::ProgressTracker;
use crate::source::ByteSource;
const SPECULATIVE_ITEM_COUNT: usize = 256;
const MAX_DEPTH: usize = 64;
const LEAF_ITEM_SIZE: usize = 32;
const INTERNAL_ITEM_SIZE: usize = 24;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Leaf {
pub start_chr: u32,
pub start_base: i64,
pub end_chr: u32,
pub end_base: i64,
pub offset: u64,
pub size: u64,
}
#[derive(Debug)]
struct NodeState {
is_leaf: bool,
buf: Bytes,
base: u64,
item_size: usize,
count: usize,
item_index: usize,
}
pub struct LeafWalk<'a> {
source: &'a dyn ByteSource,
locs: &'a [IndexedLoc],
range: Range<usize>,
tracker: &'a ProgressTracker,
stack: Vec<NodeState>,
done: bool,
}
impl<'a> LeafWalk<'a> {
pub fn new(
source: &'a dyn ByteSource,
root_offset: u64,
locs: &'a [IndexedLoc],
batch: LocBatch,
tracker: &'a ProgressTracker,
) -> Result<Self> {
let mut walk = Self {
source,
locs,
range: batch.start..batch.end,
tracker,
stack: Vec::new(),
done: batch.is_empty(),
};
if !walk.done {
let root = walk.read_node(root_offset)?;
walk.stack.push(root);
}
Ok(walk)
}
fn read_node(&self, offset: u64) -> Result<NodeState> {
let speculative = 4 + LEAF_ITEM_SIZE * SPECULATIVE_ITEM_COUNT;
let buf = self.source.read_at(offset, speculative)?;
if buf.len() < 4 {
return Err(Error::corrupt(
self.source.path(),
offset,
"truncated data tree node header",
));
}
let is_leaf = buf[0] != 0;
let count = u16::from_le_bytes([buf[2], buf[3]]) as usize;
let item_size = if is_leaf {
LEAF_ITEM_SIZE
} else {
INTERNAL_ITEM_SIZE
};
let body = count * item_size;
let buf = if buf.len() >= 4 + body {
buf.slice(4..4 + body)
} else {
self.source.read_exact_at(offset + 4, body)?
};
Ok(NodeState {
is_leaf,
buf,
base: offset + 4,
item_size,
count,
item_index: 0,
})
}
fn match_loci(
&mut self,
start_chr: u32,
start_base: i64,
end_chr: u32,
end_base: i64,
) -> usize {
let mut index = self.range.start;
while index < self.range.end {
let loc = &self.locs[index];
let chr = loc.chr_index as u32;
if chr < start_chr
|| (chr == start_chr && loc.binned_end <= start_base && index == self.range.start)
{
self.tracker
.add((loc.binned_end - loc.binned_start).max(0) as u64);
self.range.start += 1;
index += 1;
continue;
}
if chr > end_chr || (chr == end_chr && loc.binned_start > end_base) {
break;
}
index += 1;
}
index
}
fn drain_remaining(&mut self) {
while self.range.start < self.range.end {
let loc = &self.locs[self.range.start];
self.tracker
.add((loc.binned_end - loc.binned_start).max(0) as u64);
self.range.start += 1;
}
}
}
impl Iterator for LeafWalk<'_> {
type Item = Result<(Leaf, Range<usize>)>;
fn next(&mut self) -> Option<Self::Item> {
if self.done {
return None;
}
while let Some(state) = self.stack.last() {
if self.range.start >= self.range.end {
break;
}
if state.item_index >= state.count {
self.stack.pop();
continue;
}
let (
item_start_chr,
item_start_base,
item_end_chr,
item_end_base,
item_offset,
is_leaf,
base,
item_index,
item_size,
) = {
let state = self.stack.last().unwrap();
let at = state.item_index * state.item_size;
let mut c = LeCursor::new(&state.buf, state.base, self.source.path());
if let Err(e) = c.seek(at) {
return Some(Err(e));
}
let read = (|| -> Result<(u32, i64, u32, i64, u64)> {
Ok((
c.read_u32()?,
c.read_u32()? as i64,
c.read_u32()?,
c.read_u32()? as i64,
c.read_u64()?,
))
})();
match read {
Ok((a, b, cc, d, e)) => (
a,
b,
cc,
d,
e,
state.is_leaf,
state.base,
state.item_index,
state.item_size,
),
Err(e) => return Some(Err(e)),
}
};
let matched =
self.match_loci(item_start_chr, item_start_base, item_end_chr, item_end_base);
let state = self.stack.last_mut().unwrap();
if matched == self.range.start {
state.item_index += 1;
continue;
}
state.item_index += 1;
if is_leaf {
let size = {
let state = self.stack.last().unwrap();
let mut c = LeCursor::new(&state.buf, base, self.source.path());
if let Err(e) = c.seek(item_index * item_size + 24) {
return Some(Err(e));
}
match c.read_u64() {
Ok(size) => size,
Err(e) => return Some(Err(e)),
}
};
return Some(Ok((
Leaf {
start_chr: item_start_chr,
start_base: item_start_base,
end_chr: item_end_chr,
end_base: item_end_base,
offset: item_offset,
size,
},
self.range.start..matched,
)));
}
if self.stack.len() >= MAX_DEPTH {
self.done = true;
return Some(Err(Error::corrupt(
self.source.path(),
item_offset,
format!("data tree is deeper than {MAX_DEPTH} levels"),
)));
}
match self.read_node(item_offset) {
Ok(child) => self.stack.push(child),
Err(e) => {
self.done = true;
return Some(Err(e));
}
}
}
self.done = true;
self.drain_remaining();
None
}
}
#[derive(Debug, Clone, Copy)]
pub struct LeafItem {
pub start_chr: u32,
pub start_base: u32,
pub end_chr: u32,
pub end_base: u32,
pub offset: u64,
pub size: u64,
}
pub const TREE_BLOCK_SIZE: u32 = 256;
pub const TREE_NODE_HEADER_SIZE: usize = 4;
#[derive(Debug, Clone)]
pub struct TreeShape {
pub counts: Vec<u64>,
pub spans: Vec<u64>,
}
impl TreeShape {
pub fn new(item_count: u64, block_size: u32) -> Result<Self> {
if block_size < 2 {
return Err(Error::invalid(format!(
"tree block size {block_size} invalid (>= 2)"
)));
}
let block = block_size as u64;
let mut levels = 1usize;
let mut root_span = block;
while root_span < item_count {
root_span = root_span.saturating_mul(block);
levels += 1;
}
let mut counts = vec![0u64; levels];
let mut spans = vec![0u64; levels];
let mut span = block;
for level in (0..levels).rev() {
spans[level] = span;
counts[level] = item_count.div_ceil(span).max(1);
span = span.saturating_mul(block);
}
Ok(Self { counts, spans })
}
pub fn levels(&self) -> usize {
self.counts.len()
}
pub fn is_leaf_level(&self, level: usize) -> bool {
level == self.levels() - 1
}
}
#[derive(Debug, Clone, Copy, Default)]
struct Bounds {
start_chr: u32,
start_base: u32,
end_chr: u32,
end_base: u32,
}
fn position_less(chr_a: u32, base_a: u32, chr_b: u32, base_b: u32) -> bool {
(chr_a, base_a) < (chr_b, base_b)
}
pub fn write_tree(
items: &[LeafItem],
tree_offset: u64,
block_size: u32,
items_per_slot: u32,
end_file_offset: u64,
sink: &mut dyn FnMut(&[u8]),
) -> Result<()> {
let item_count = items.len() as u64;
let shape = TreeShape::new(item_count, block_size)?;
let levels = shape.levels();
let block = block_size as u64;
let mut bounds: Vec<Vec<Bounds>> = vec![Vec::new(); levels];
for level in (0..levels).rev() {
let mut level_bounds = vec![Bounds::default(); shape.counts[level] as usize];
for node in 0..shape.counts[level] {
let child_count = child_count(&shape, level, node, item_count, block);
let mut box_ = Bounds::default();
for i in 0..child_count {
let child = if shape.is_leaf_level(level) {
let item = &items[(node * block + i) as usize];
Bounds {
start_chr: item.start_chr,
start_base: item.start_base,
end_chr: item.end_chr,
end_base: item.end_base,
}
} else {
bounds[level + 1][(node * block + i) as usize]
};
if i == 0 {
box_ = child;
continue;
}
if position_less(
child.start_chr,
child.start_base,
box_.start_chr,
box_.start_base,
) {
box_.start_chr = child.start_chr;
box_.start_base = child.start_base;
}
if position_less(box_.end_chr, box_.end_base, child.end_chr, child.end_base) {
box_.end_chr = child.end_chr;
box_.end_base = child.end_base;
}
}
level_bounds[node as usize] = box_;
}
bounds[level] = level_bounds;
}
let root = bounds[0][0];
let mut header = Vec::with_capacity(DATA_TREE_HEADER_SIZE as usize);
header.extend_from_slice(&DATA_TREE_MAGIC.to_le_bytes());
header.extend_from_slice(&block_size.to_le_bytes());
header.extend_from_slice(&item_count.to_le_bytes());
header.extend_from_slice(&root.start_chr.to_le_bytes());
header.extend_from_slice(&root.start_base.to_le_bytes());
header.extend_from_slice(&root.end_chr.to_le_bytes());
header.extend_from_slice(&root.end_base.to_le_bytes());
header.extend_from_slice(&end_file_offset.to_le_bytes());
header.extend_from_slice(&items_per_slot.to_le_bytes());
header.extend_from_slice(&[0u8; 4]); debug_assert_eq!(header.len(), DATA_TREE_HEADER_SIZE as usize);
sink(&header);
let node_size = |level: usize| -> u64 {
let item_size = if shape.is_leaf_level(level) {
LEAF_ITEM_SIZE
} else {
INTERNAL_ITEM_SIZE
};
TREE_NODE_HEADER_SIZE as u64 + block * item_size as u64
};
let mut level_offsets = vec![0u64; levels];
let mut offset = tree_offset + DATA_TREE_HEADER_SIZE;
for (level, slot) in level_offsets.iter_mut().enumerate() {
*slot = offset;
offset += shape.counts[level] * node_size(level);
}
let mut node = Vec::new();
for level in 0..levels {
let is_leaf = shape.is_leaf_level(level);
let item_size = if is_leaf {
LEAF_ITEM_SIZE
} else {
INTERNAL_ITEM_SIZE
};
let child_node_size = if is_leaf { 0 } else { node_size(level + 1) };
for index in 0..shape.counts[level] {
let count = child_count(&shape, level, index, item_count, block);
node.clear();
node.reserve(TREE_NODE_HEADER_SIZE + block as usize * item_size);
node.push(u8::from(is_leaf));
node.push(0); node.extend_from_slice(&(count as u16).to_le_bytes());
for i in 0..count {
if is_leaf {
let item = &items[(index * block + i) as usize];
node.extend_from_slice(&item.start_chr.to_le_bytes());
node.extend_from_slice(&item.start_base.to_le_bytes());
node.extend_from_slice(&item.end_chr.to_le_bytes());
node.extend_from_slice(&item.end_base.to_le_bytes());
node.extend_from_slice(&item.offset.to_le_bytes());
node.extend_from_slice(&item.size.to_le_bytes());
} else {
let child = index * block + i;
let box_ = bounds[level + 1][child as usize];
node.extend_from_slice(&box_.start_chr.to_le_bytes());
node.extend_from_slice(&box_.start_base.to_le_bytes());
node.extend_from_slice(&box_.end_chr.to_le_bytes());
node.extend_from_slice(&box_.end_base.to_le_bytes());
node.extend_from_slice(
&(level_offsets[level + 1] + child * child_node_size).to_le_bytes(),
);
}
}
node.resize(TREE_NODE_HEADER_SIZE + block as usize * item_size, 0);
sink(&node);
}
}
Ok(())
}
fn child_count(shape: &TreeShape, level: usize, index: u64, item_count: u64, block: u64) -> u64 {
let total = if shape.is_leaf_level(level) {
item_count
} else {
shape.counts[level + 1]
};
block.min(total.saturating_sub(index * block))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::source::testing::MemorySource;
fn loc(chr: usize, start: i64, end: i64) -> IndexedLoc {
IndexedLoc {
chr_index: chr,
start,
end,
binned_start: start,
binned_end: end,
bin_size: 1.0,
reverse: false,
output_start: 0,
output_end: 1,
}
}
fn leaf_node(items: &[(u32, u32, u32, u32, u64, u64)]) -> Vec<u8> {
let mut b = vec![1u8, 0];
b.extend_from_slice(&(items.len() as u16).to_le_bytes());
for (sc, sb, ec, eb, off, size) in items {
b.extend_from_slice(&sc.to_le_bytes());
b.extend_from_slice(&sb.to_le_bytes());
b.extend_from_slice(&ec.to_le_bytes());
b.extend_from_slice(&eb.to_le_bytes());
b.extend_from_slice(&off.to_le_bytes());
b.extend_from_slice(&size.to_le_bytes());
}
b
}
fn internal_node(items: &[(u32, u32, u32, u32, u64)]) -> Vec<u8> {
let mut b = vec![0u8, 0];
b.extend_from_slice(&(items.len() as u16).to_le_bytes());
for (sc, sb, ec, eb, off) in items {
b.extend_from_slice(&sc.to_le_bytes());
b.extend_from_slice(&sb.to_le_bytes());
b.extend_from_slice(&ec.to_le_bytes());
b.extend_from_slice(&eb.to_le_bytes());
b.extend_from_slice(&off.to_le_bytes());
}
b
}
fn walk(bytes: Vec<u8>, root: u64, locs: &[IndexedLoc]) -> (Vec<(Leaf, Range<usize>)>, u64) {
let source = MemorySource::new(bytes);
let tracker = ProgressTracker::new(
locs.iter()
.map(|l| (l.binned_end - l.binned_start) as u64)
.sum(),
);
let batch = LocBatch {
start: 0,
end: locs.len(),
};
let out: Vec<_> = LeafWalk::new(&source, root, locs, batch, &tracker)
.unwrap()
.map(|r| r.unwrap())
.collect();
(out, tracker.done())
}
#[test]
fn a_flat_tree_yields_the_overlapping_leaves_with_their_loci() {
let bytes = leaf_node(&[
(0, 0, 0, 100, 1000, 10),
(0, 100, 0, 200, 2000, 20),
(0, 200, 0, 300, 3000, 30),
]);
let locs = [loc(0, 120, 180)];
let (got, _) = walk(bytes, 0, &locs);
assert_eq!(got.len(), 1);
assert_eq!(got[0].0.offset, 2000);
assert_eq!(got[0].0.size, 20);
assert_eq!(got[0].1, 0..1);
}
#[test]
fn a_locus_spanning_several_leaves_is_reported_by_each() {
let bytes = leaf_node(&[
(0, 0, 0, 100, 1000, 10),
(0, 100, 0, 200, 2000, 20),
(0, 200, 0, 300, 3000, 30),
]);
let locs = [loc(0, 50, 250)];
let (got, _) = walk(bytes, 0, &locs);
assert_eq!(
got.iter().map(|(l, _)| l.offset).collect::<Vec<_>>(),
[1000, 2000, 3000]
);
assert!(got.iter().all(|(_, r)| *r == (0..1)));
}
#[test]
fn passed_loci_are_dropped_and_reported_to_the_tracker() {
let bytes = leaf_node(&[(0, 100, 0, 200, 2000, 20)]);
let locs = [loc(0, 0, 10), loc(0, 150, 160)];
let (got, done) = walk(bytes, 0, &locs);
assert_eq!(got.len(), 1);
assert_eq!(got[0].1, 1..2);
assert_eq!(done, 20);
}
#[test]
fn every_locus_reaches_the_tracker_even_with_no_matching_leaf() {
let bytes = leaf_node(&[(5, 0, 5, 100, 2000, 20)]);
let locs = [loc(0, 0, 30), loc(0, 100, 140)];
let (got, done) = walk(bytes, 0, &locs);
assert!(got.is_empty());
assert_eq!(done, 70);
}
#[test]
fn descends_internal_nodes_depth_first_in_file_order() {
let root = internal_node(&[(0, 0, 0, 100, 100), (0, 100, 0, 200, 200)]);
let mut bytes = root;
bytes.resize(100, 0);
bytes.extend_from_slice(&leaf_node(&[(0, 0, 0, 100, 1000, 11)]));
bytes.resize(200, 0);
bytes.extend_from_slice(&leaf_node(&[(0, 100, 0, 200, 2000, 22)]));
let locs = [loc(0, 0, 200)];
let (got, _) = walk(bytes, 0, &locs);
assert_eq!(
got.iter().map(|(l, _)| l.offset).collect::<Vec<_>>(),
[1000, 2000]
);
}
#[test]
fn a_subtree_no_locus_reaches_is_never_descended() {
let root = internal_node(&[(0, 0, 0, 100, 100), (9, 0, 9, 100, 99_000)]);
let mut bytes = root;
bytes.resize(100, 0);
bytes.extend_from_slice(&leaf_node(&[(0, 0, 0, 100, 1000, 11)]));
let locs = [loc(0, 0, 50)];
let (got, _) = walk(bytes, 0, &locs);
assert_eq!(got.len(), 1);
assert_eq!(got[0].0.offset, 1000);
}
#[test]
fn an_item_spanning_a_chromosome_boundary_matches_loci_on_both() {
let bytes = leaf_node(&[(0, 900, 1, 100, 5000, 50)]);
let locs = [loc(0, 950, 960), loc(1, 10, 20)];
let (got, _) = walk(bytes, 0, &locs);
assert_eq!(got.len(), 1);
assert_eq!(got[0].1, 0..2);
}
#[test]
fn an_empty_batch_yields_nothing() {
let source = MemorySource::new(leaf_node(&[(0, 0, 0, 100, 1000, 10)]));
let tracker = ProgressTracker::new(0);
let locs: [IndexedLoc; 0] = [];
let mut walk =
LeafWalk::new(&source, 0, &locs, LocBatch { start: 0, end: 0 }, &tracker).unwrap();
assert!(walk.next().is_none());
}
#[test]
fn a_cycle_in_the_tree_stops_at_the_depth_limit() {
let bytes = internal_node(&[(0, 0, 0, 1000, 0)]);
let source = MemorySource::new(bytes);
let tracker = ProgressTracker::new(100);
let locs = [loc(0, 0, 100)];
let walk =
LeafWalk::new(&source, 0, &locs, LocBatch { start: 0, end: 1 }, &tracker).unwrap();
let err = walk.filter_map(|r| r.err()).next().unwrap();
assert!(err.to_string().contains("deeper than 64 levels"), "{err}");
}
#[test]
fn a_truncated_node_is_corrupt_not_a_panic() {
let mut bytes = leaf_node(&[(0, 0, 0, 100, 1000, 10), (0, 100, 0, 200, 2000, 20)]);
bytes.truncate(bytes.len() - 6);
let source = MemorySource::new(bytes);
let tracker = ProgressTracker::new(200);
let locs = [loc(0, 0, 200)];
let walk = LeafWalk::new(&source, 0, &locs, LocBatch { start: 0, end: 1 }, &tracker);
let failed = match walk {
Err(_) => true,
Ok(w) => w.filter_map(|r| r.err()).next().is_some(),
};
assert!(failed);
}
fn write_and_walk(
items: &[LeafItem],
block_size: u32,
locs: &[IndexedLoc],
) -> Vec<(Leaf, Range<usize>)> {
let mut bytes = Vec::new();
write_tree(items, 0, block_size, 1024, 0, &mut |b| {
bytes.extend_from_slice(b)
})
.unwrap();
let source = MemorySource::new(bytes);
crate::bbi::header::check_data_tree_magic(&source, 0).unwrap();
let tracker = ProgressTracker::new(
locs.iter()
.map(|l| (l.binned_end - l.binned_start) as u64)
.sum(),
);
let batch = LocBatch {
start: 0,
end: locs.len(),
};
LeafWalk::new(&source, DATA_TREE_HEADER_SIZE, locs, batch, &tracker)
.unwrap()
.map(|r| r.unwrap())
.collect()
}
fn item(chr: u32, start: u32, end: u32, offset: u64) -> LeafItem {
LeafItem {
start_chr: chr,
start_base: start,
end_chr: chr,
end_base: end,
offset,
size: 16,
}
}
#[test]
fn a_written_flat_tree_walks_back_to_the_overlapping_leaves() {
let items: Vec<LeafItem> = (0..5)
.map(|i| item(0, i * 100, (i + 1) * 100, 1000 + i as u64 * 10))
.collect();
let got = write_and_walk(&items, 256, &[loc(0, 150, 320)]);
assert_eq!(
got.iter().map(|(l, _)| l.offset).collect::<Vec<_>>(),
[1010, 1020, 1030]
);
}
#[test]
fn a_written_deep_tree_walks_back_to_every_leaf() {
let items: Vec<LeafItem> = (0..17)
.map(|i| item(0, i * 10, (i + 1) * 10, 5000 + i as u64))
.collect();
let got = write_and_walk(&items, 2, &[loc(0, 0, 170)]);
assert_eq!(got.len(), 17);
assert_eq!(
got.iter().map(|(l, _)| l.offset).collect::<Vec<_>>(),
(5000..5017).collect::<Vec<u64>>()
);
}
#[test]
fn a_written_tree_spanning_chromosomes_keeps_them_apart() {
let items = vec![
item(0, 0, 100, 10),
item(0, 100, 200, 20),
item(1, 0, 100, 30),
item(2, 0, 100, 40),
];
let got = write_and_walk(&items, 2, &[loc(1, 0, 100)]);
assert_eq!(got.iter().map(|(l, _)| l.offset).collect::<Vec<_>>(), [30]);
}
#[test]
fn the_written_header_carries_the_root_bounds_and_the_leaf_count() {
let items = vec![item(0, 40, 100, 10), item(3, 0, 900, 20)];
let mut bytes = Vec::new();
write_tree(&items, 0, 256, 512, 4096, &mut |b| {
bytes.extend_from_slice(b)
})
.unwrap();
let u32_at = |o: usize| u32::from_le_bytes(bytes[o..o + 4].try_into().unwrap());
let u64_at = |o: usize| u64::from_le_bytes(bytes[o..o + 8].try_into().unwrap());
assert_eq!(u32_at(0), DATA_TREE_MAGIC);
assert_eq!(u32_at(4), 256); assert_eq!(u64_at(8), 2); assert_eq!((u32_at(16), u32_at(20)), (0, 40)); assert_eq!((u32_at(24), u32_at(28)), (3, 900)); assert_eq!(u64_at(32), 4096); assert_eq!(u32_at(40), 512); }
#[test]
fn every_node_is_padded_to_block_size_so_the_offsets_are_arithmetic() {
let items: Vec<LeafItem> = (0..3)
.map(|i| item(0, i * 10, i * 10 + 5, i as u64))
.collect();
let mut bytes = Vec::new();
write_tree(&items, 0, 4, 1, 0, &mut |b| bytes.extend_from_slice(b)).unwrap();
let expected = DATA_TREE_HEADER_SIZE as usize + TREE_NODE_HEADER_SIZE + 4 * LEAF_ITEM_SIZE;
assert_eq!(bytes.len(), expected);
assert_eq!(
u16::from_le_bytes(
bytes[DATA_TREE_HEADER_SIZE as usize + 2..DATA_TREE_HEADER_SIZE as usize + 4]
.try_into()
.unwrap()
),
3
);
assert!(bytes[expected - LEAF_ITEM_SIZE..].iter().all(|b| *b == 0));
}
#[test]
fn a_block_size_below_two_is_refused() {
let err = write_tree(&[item(0, 0, 1, 0)], 0, 1, 1, 0, &mut |_| {})
.unwrap_err()
.to_string();
assert!(err.contains("tree block size 1 invalid"), "{err}");
}
}