diff --git a/modkit-core/src/entropy/methylation_entropy.rs b/modkit-core/src/entropy/methylation_entropy.rs index ce66982e..187dddd0 100644 --- a/modkit-core/src/entropy/methylation_entropy.rs +++ b/modkit-core/src/entropy/methylation_entropy.rs @@ -1,133 +1,147 @@ use derive_new::new; use itertools::Itertools; use log::debug; -use regex::Regex; -use rustc_hash::{FxHashMap, FxHashSet}; +use rustc_hash::FxHashSet; use std::collections::{BTreeSet, HashMap}; -use std::str::Chars; -use substring::Substring; + +#[derive(Debug, Copy, Clone, Eq, PartialEq, Hash, Ord, PartialOrd)] +#[repr(transparent)] +pub(super) struct EntropySymbol(u32); + +impl EntropySymbol { + pub(super) const CANONICAL: Self = Self(0); + pub(super) const FILTERED: Self = Self(u32::MAX); + + pub(super) fn called(id: u32) -> Self { + assert!( + id > Self::CANONICAL.0 && id < Self::FILTERED.0, + "modified state ID must not collide with canonical or wildcard" + ); + Self(id) + } + + fn is_filtered(self) -> bool { + self == Self::FILTERED + } +} + +pub(super) type EntropyPattern = Box<[EntropySymbol]>; #[derive(new)] struct AlphabetInfo { - columns: FxHashMap, + columns: Box<[Box<[EntropySymbol]>]>, } impl AlphabetInfo { - fn from_sequences(sequences: &[String], window_size: usize) -> Self { + fn from_sequences( + sequences: &[EntropyPattern], + window_size: usize, + ) -> Self { assert_eq!( sequences.iter().map(|x| x.len()).unique().count(), 1, "all sequences should be the same length {sequences:?}" ); - let columns = - (0..window_size).fold(FxHashMap::default(), |mut acc, idx| { - acc.insert(idx, BTreeSet::new()); - acc - }); - let columns = sequences - .iter() - .fold(columns, |mut acc, seq| { - for (i, c) in seq.chars().enumerate().filter(|(_, c)| *c != '*') - { - acc.entry(i).or_insert_with(BTreeSet::new).insert(c); + let mut columns = vec![BTreeSet::new(); window_size]; + for sequence in sequences { + for (position, symbol) in sequence.iter().enumerate() { + if !symbol.is_filtered() { + columns[position].insert(*symbol); } - acc - }) + } + } + let columns = columns .into_iter() - .inspect(|(_, elems)| { + .map(|elements| { debug_assert!( - elems.len() >= 1, - "column with zero coverage in {sequences:?}, {elems:?}" - ) - }) - .map(|(pos, elements)| { - (pos, elements.into_iter().collect::()) + !elements.is_empty(), + "column with zero coverage in {sequences:?}, \ + {elements:?}" + ); + elements.into_iter().collect::>().into_boxed_slice() }) - .collect::>(); + .collect::>() + .into_boxed_slice(); Self { columns } } - fn seq_to_regex(&self, seq: &str) -> Regex { - let pattern = seq - .chars() - .enumerate() - .map(|(pos, c)| match c { - '*' => { - format!("[{}]", self.columns.get(&pos).unwrap()) - } - _ => { - format!("{c}") - } - }) - .collect::(); - Regex::new(&pattern).unwrap() + fn get_column( + &self, + idx: usize, + ) -> impl Iterator + '_ { + self.columns[idx].iter().copied() } +} - fn get_column(&'_ self, idx: usize) -> Chars<'_> { - self.columns.get(&idx).unwrap().chars() - } +fn sequence_matches_pattern( + sequence: &[EntropySymbol], + pattern: &[EntropySymbol], +) -> bool { + debug_assert_eq!(sequence.len(), pattern.len()); + debug_assert!(pattern.iter().all(|symbol| !symbol.is_filtered())); + sequence + .iter() + .zip(pattern) + .all(|(observed, called)| observed.is_filtered() || observed == called) +} + +fn extend_pattern( + pattern: &[EntropySymbol], + symbol: EntropySymbol, +) -> EntropyPattern { + let mut extended = Vec::with_capacity(pattern.len() + 1); + extended.extend_from_slice(pattern); + extended.push(symbol); + extended.into_boxed_slice() } fn all_patterns_dp( - sequences: &[String], + sequences: &[EntropyPattern], window_size: usize, alphabet_info: &AlphabetInfo, -) -> Vec { - let sequences = sequences.iter().collect::>(); +) -> Vec { + let sequences = sequences.iter().collect::>(); debug_assert!( sequences.iter().all(|x| x.len() == window_size), "all sequences should be the same length, {sequences:?}" ); // easy case - if !sequences.iter().any(|x| x.contains('*')) { - return sequences.iter().map(|x| x.to_string()).collect(); + if !sequences.iter().any(|x| x.contains(&EntropySymbol::FILTERED)) { + return sequences.into_iter().cloned().collect(); } let basecase = alphabet_info .get_column(0) - .map(|c| format!("{c}")) - .collect::>(); + .map(|symbol| vec![symbol].into_boxed_slice()) + .collect::>(); debug_assert!(basecase.len() >= 1, "first column has zero valid coverage"); #[cfg(debug_assertions)] { let basecase_check = sequences .iter() - .map(|x| x.substring(0, 1).to_string()) - .filter(|x| x != "*") - .collect::>(); + .filter_map(|sequence| { + let symbol = sequence[0]; + (!symbol.is_filtered()).then(|| vec![symbol].into_boxed_slice()) + }) + .collect::>(); assert_eq!(basecase, basecase_check); } - let mut cache = FxHashMap::default(); let all_combs = (1..window_size).fold(basecase, |acc, idx| { let mut acc_patterns = FxHashSet::default(); - for seq in sequences.iter() { - let prefix = seq.substring(0, idx); - let re = if let Some(re) = cache.get(prefix) { - re - } else { - let re = alphabet_info.seq_to_regex(prefix); - cache.insert(prefix, re); - cache.get(prefix).unwrap() - }; + for sequence in &sequences { + let prefix = &sequence[..idx]; for pattern in acc.iter() { - if re.is_match(&pattern) { - let last_letter = seq - .chars() - .nth(idx) - .expect(&format!("should get last letter at {idx}")); - match last_letter { - '*' => { - for x in alphabet_info.get_column(idx) { - let new_pattern = format!("{pattern}{x}"); - acc_patterns.insert(new_pattern); - } - } - _ => { - let new_pattern = format!("{pattern}{last_letter}"); - acc_patterns.insert(new_pattern); + if sequence_matches_pattern(prefix, pattern) { + let observed = sequence[idx]; + if observed.is_filtered() { + for symbol in alphabet_info.get_column(idx) { + acc_patterns + .insert(extend_pattern(pattern, symbol)); } + } else { + acc_patterns.insert(extend_pattern(pattern, observed)); } } } @@ -136,29 +150,19 @@ fn all_patterns_dp( acc_patterns }); - all_combs.into_iter().sorted().collect::>() + all_combs.into_iter().sorted().collect::>() } -fn calc_entropy(sequences: &[String], window_size: usize) -> f32 { - let mut alphabet_info = - AlphabetInfo::from_sequences(sequences, window_size); - let patterns = all_patterns_dp(sequences, window_size, &mut alphabet_info); +fn calc_entropy(sequences: &[EntropyPattern], window_size: usize) -> f32 { + let alphabet_info = AlphabetInfo::from_sequences(sequences, window_size); + let patterns = all_patterns_dp(sequences, window_size, &alphabet_info); - let mut cache = FxHashMap::default(); let counts = sequences.iter().fold(HashMap::new(), |mut acc, seq| { - // let re = seq_to_regex(seq, &alphabet_info.wildcard_regex); - let re = if let Some(re) = cache.get(seq) { - re - } else { - let re = alphabet_info.seq_to_regex(seq); - cache.insert(seq, re); - cache.get(seq).unwrap() - }; let matches = patterns .iter() - .filter(|p| re.is_match(p)) - .collect::>(); - assert!(matches.len() > 0, "no matches for {seq} in {patterns:?}"); + .filter(|pattern| sequence_matches_pattern(seq, pattern)) + .collect::>(); + assert!(matches.len() > 0, "no matches for {seq:?} in {patterns:?}"); let factor = 1f32 / matches.len() as f32; for pattern in matches { *acc.entry(pattern).or_insert(0f32) += factor; @@ -193,7 +197,7 @@ fn calc_entropy(sequences: &[String], window_size: usize) -> f32 { } pub(super) fn calc_me_entropy( - sequences: &[String], + sequences: &[EntropyPattern], window_size: usize, constant: f32, ) -> f32 { @@ -208,151 +212,223 @@ pub(super) fn calc_me_entropy( #[cfg(test)] mod methylation_entropy_tests { - use crate::entropy::methylation_entropy::{ + use super::{ all_patterns_dp, calc_entropy, calc_me_entropy, AlphabetInfo, + EntropyPattern, EntropySymbol, }; use assert_approx_eq::assert_approx_eq; + use std::mem::{size_of, size_of_val}; + + // Keep the original 0/1/* fixtures readable while ensuring that the + // production entropy path has no text representation. + fn legacy_pattern(raw: &str) -> EntropyPattern { + raw.chars() + .map(|symbol| match symbol { + '0' => EntropySymbol::CANONICAL, + '1' => EntropySymbol::called(1), + '*' => EntropySymbol::FILTERED, + _ => panic!("unsupported legacy test symbol {symbol}"), + }) + .collect::>() + .into_boxed_slice() + } + + fn legacy_patterns(raw: &[&str]) -> Vec { + raw.iter().map(|pattern| legacy_pattern(pattern)).collect() + } + + fn pattern(states: &[Option]) -> EntropyPattern { + states + .iter() + .map(|state| match state { + Some(0) => EntropySymbol::CANONICAL, + Some(id) => EntropySymbol::called((*id).try_into().unwrap()), + None => EntropySymbol::FILTERED, + }) + .collect::>() + .into_boxed_slice() + } + + fn called(ids: &[usize]) -> EntropyPattern { + ids.iter() + .map(|id| match id { + 0 => EntropySymbol::CANONICAL, + _ => EntropySymbol::called((*id).try_into().unwrap()), + }) + .collect::>() + .into_boxed_slice() + } + + fn p_q_w_oracle(code_count: usize) -> Vec { + let p = called(&(1..=code_count).collect::>()); + let mut q = p.to_vec(); + q[0] = EntropySymbol::CANONICAL; + let mut wildcard = p.to_vec(); + wildcard[0] = EntropySymbol::FILTERED; + vec![p, q.into_boxed_slice(), wildcard.into_boxed_slice()] + } + + fn assert_per_read_mass_and_counts( + sequences: &[EntropyPattern], + concrete_patterns: &[EntropyPattern], + ) { + let mut counts = vec![0.0f32; concrete_patterns.len()]; + for sequence in sequences { + let matching_indices = concrete_patterns + .iter() + .enumerate() + .filter_map(|(idx, concrete)| { + super::sequence_matches_pattern(sequence, concrete) + .then_some(idx) + }) + .collect::>(); + assert!(!matching_indices.is_empty()); + let factor = 1.0 / matching_indices.len() as f32; + let read_mass = factor * matching_indices.len() as f32; + assert_eq!(read_mass, 1.0); + for idx in matching_indices { + counts[idx] += factor; + } + } + assert_eq!(counts, vec![1.5, 1.5]); + assert_eq!(counts.iter().sum::(), sequences.len() as f32); + } + + #[test] + fn entropy_symbols_are_compact_and_patterns_have_no_capacity_field() { + assert_eq!(size_of::(), size_of::()); + assert_eq!( + size_of::(), + size_of::>() + ); + let sequence = called(&[0, 1, 10, 100]); + assert_eq!( + size_of_val(sequence.as_ref()), + sequence.len() * size_of::() + ); + assert!(std::panic::catch_unwind(|| EntropySymbol::called(0)).is_err()); + assert!(std::panic::catch_unwind(|| EntropySymbol::called(u32::MAX)) + .is_err()); + } #[test] fn test_calc_entropy() { - let sequences = vec![ - "0000".to_string(), - "0000".to_string(), - "0000".to_string(), - "0000".to_string(), - ]; + let sequences = legacy_patterns(&["0000", "0000", "0000", "0000"]); assert_eq!(calc_me_entropy(&sequences, 4, 0.25), 0.0); - let sequences = vec![ - "1111".to_string(), - "1111".to_string(), - "1111".to_string(), - "1111".to_string(), - ]; + let sequences = legacy_patterns(&["1111", "1111", "1111", "1111"]); assert_eq!(calc_me_entropy(&sequences, 4, 0.25), 0.0); - let sequences = vec![ - "0010".to_string(), - "0010".to_string(), - "0010".to_string(), - "0010".to_string(), - ]; + let sequences = legacy_patterns(&["0010", "0010", "0010", "0010"]); assert_eq!(calc_me_entropy(&sequences, 4, 0.25), 0.0); - let sequences = vec![ - "1111".to_string(), - "1111".to_string(), - "1111".to_string(), - "1111".to_string(), - "0000".to_string(), - "0000".to_string(), - "0000".to_string(), - "0000".to_string(), - ]; + let sequences = legacy_patterns(&[ + "1111", "1111", "1111", "1111", "0000", "0000", "0000", "0000", + ]); assert_eq!(calc_me_entropy(&sequences, 4, 0.25), 0.25); - let sequences = vec![ - "1111".to_string(), - "1111".to_string(), - "0011".to_string(), - "0011".to_string(), - "1100".to_string(), - "1100".to_string(), - "0000".to_string(), - "0000".to_string(), - ]; + let sequences = legacy_patterns(&[ + "1111", "1111", "0011", "0011", "1100", "1100", "0000", "0000", + ]); assert_eq!(calc_me_entropy(&sequences, 4, 0.25), 0.50); - let sequences = vec![ - "0000".to_string(), - "1111".to_string(), - "0101".to_string(), - "0111".to_string(), - "0111".to_string(), - "0111".to_string(), - "0000".to_string(), - "1111".to_string(), - ]; + let sequences = legacy_patterns(&[ + "0000", "1111", "0101", "0111", "0111", "0111", "0000", "1111", + ]); assert_eq!(calc_me_entropy(&sequences, 4, 0.25), 0.47640976); } #[test] fn test_calc_entropy_wildcards() { - let sequences = vec!["1*01", "1111", "1011", "1111"] - .into_iter() - .map(|x| x.to_string()) - .collect::>(); + let sequences = legacy_patterns(&["1*01", "1111", "1011", "1111"]); let alphabet_info = AlphabetInfo::from_sequences(&sequences, 4); let patterns = all_patterns_dp(&sequences, 4, &alphabet_info); assert_eq!( patterns, - vec![ - "1001".to_string(), - "1011".to_string(), - "1101".to_string(), - "1111".to_string(), - ] + legacy_patterns(&["1001", "1011", "1101", "1111"]) ); let entropy = calc_entropy(&sequences, 4); assert_eq!(entropy, 1.75); - let sequences = vec!["1*11", "1111", "1011", "1111"] - .into_iter() - .map(|x| x.to_string()) - .collect::>(); + let sequences = legacy_patterns(&["1*11", "1111", "1011", "1111"]); let alphabet_info = AlphabetInfo::from_sequences(&sequences, 4); let patterns = all_patterns_dp(&sequences, 4, &alphabet_info); - assert_eq!(patterns, vec!["1011".to_string(), "1111".to_string(),]); + assert_eq!(patterns, legacy_patterns(&["1011", "1111"])); let entropy = calc_entropy(&sequences, 4); assert_eq!(entropy, 0.95443404); - let sequences = vec!["1*01", "1101", "1011", "1111"] - .into_iter() - .map(|x| x.to_string()) - .collect::>(); + let sequences = legacy_patterns(&["1*01", "1101", "1011", "1111"]); let alphabet_info = AlphabetInfo::from_sequences(&sequences, 4); let patterns = all_patterns_dp(&sequences, 4, &alphabet_info); assert_eq!( patterns, - vec![ - "1001".to_string(), - "1011".to_string(), - "1101".to_string(), - "1111".to_string(), - ] + legacy_patterns(&["1001", "1011", "1101", "1111"]) ); let entropy = calc_entropy(&sequences, 4); assert_approx_eq!(entropy, 1.9, 0.01); - let sequences = vec!["*010", "1010", "0010"] - .into_iter() - .map(|x| x.to_string()) - .collect::>(); + let sequences = legacy_patterns(&["*010", "1010", "0010"]); let alphabet_info = AlphabetInfo::from_sequences(&sequences, 4); let patterns = all_patterns_dp(&sequences, 4, &alphabet_info); - assert_eq!(patterns, vec!["0010".to_string(), "1010".to_string(),]); + assert_eq!(patterns, legacy_patterns(&["0010", "1010"])); let entropy = calc_entropy(&sequences, 4); assert_eq!(entropy, 1.0f32); - let sequences = vec!["1010", "1010", "1010", "1010"] - .into_iter() - .map(|x| x.to_string()) - .collect::>(); + let sequences = legacy_patterns(&["1010", "1010", "1010", "1010"]); let _alphabet_info = AlphabetInfo::from_sequences(&sequences, 4); let entropy = calc_entropy(&sequences, 4); assert_eq!(entropy, 0f32); } + #[test] + fn repeated_one_code_control_has_zero_entropy() { + let sequences = vec![called(&[12]); 8]; + assert_eq!(calc_entropy(&sequences, 1), 0.0); + } + + #[test] + fn p_q_w_oracles_cover_large_state_sets_and_read_order_permutations() { + for code_count in [9, 12, 17] { + let sequences = p_q_w_oracle(code_count); + let alphabet = AlphabetInfo::from_sequences(&sequences, code_count); + let concrete = all_patterns_dp(&sequences, code_count, &alphabet); + assert_eq!(concrete.len(), 2); + assert_per_read_mass_and_counts(&sequences, &concrete); + + let expected = 1.0 / code_count as f32; + assert_eq!( + calc_me_entropy( + &sequences, + code_count, + 1.0 / code_count as f32, + ), + expected + ); + + let mut reversed = sequences.clone(); + reversed.reverse(); + assert_eq!( + calc_me_entropy(&reversed, code_count, 1.0 / code_count as f32,), + expected + ); + } + } + + #[test] + fn twelve_code_wildcard_matches_uniform_hand_calculation() { + let mut sequences = + (1..=12).map(|id| called(&[id])).collect::>(); + sequences.push(pattern(&[None])); + + let entropy = calc_entropy(&sequences, 1); + assert_approx_eq!(entropy, 12f32.log2(), 0.000_01); + } + #[test] #[should_panic] fn test_alphabet_info() { // test that assert fires when columns are all * - let sequences = vec![ - "*111".to_string(), - "*111".to_string(), - "*111".to_string(), - "*111".to_string(), - ]; + let sequences = legacy_patterns(&["*111", "*111", "*111", "*111"]); AlphabetInfo::from_sequences(&sequences, 4); } } diff --git a/modkit-core/src/entropy/mod.rs b/modkit-core/src/entropy/mod.rs index 5de1ee4b..361051d8 100644 --- a/modkit-core/src/entropy/mod.rs +++ b/modkit-core/src/entropy/mod.rs @@ -1,4 +1,4 @@ -use std::collections::{BTreeSet, HashMap, HashSet, VecDeque}; +use std::collections::{BTreeMap, BTreeSet, HashMap, HashSet, VecDeque}; use std::fs::File; use std::io::{BufRead, BufReader}; use std::ops::Range; @@ -16,7 +16,9 @@ use rust_htslib::bam::ext::BamRecordExtensions; use rust_htslib::bam::{self, FetchDefinition, Read}; use rustc_hash::FxHashMap; -use crate::entropy::methylation_entropy::calc_me_entropy; +use crate::entropy::methylation_entropy::{ + calc_me_entropy, EntropyPattern, EntropySymbol, +}; use crate::errs::{MkError, MkResult}; use crate::mod_bam::{BaseModCall, ModBaseInfo}; use crate::mod_base_code::{DnaBase, ModCodeRepr}; @@ -33,6 +35,24 @@ mod writers; type BaseAndPosition = (DnaBase, u64); +fn half_open_interval(positions: I) -> Range +where + I: IntoIterator, +{ + match positions.into_iter().minmax() { + MinMaxResult::MinMax(start, end) => { + start..end.checked_add(1).expect("reference interval overflow") + } + MinMaxResult::OneElement(position) => { + position + ..position.checked_add(1).expect("reference interval overflow") + } + MinMaxResult::NoElements => { + unreachable!("cannot build an interval without positions") + } + } +} + #[derive(Debug)] pub(super) enum GenomeWindow { CombineStrands { @@ -76,22 +96,10 @@ impl GenomeWindow { num_positions: usize, ) -> Self { let pos_interval = pos_positions.as_ref().map(|positions| { - match positions.iter().map(|(_, p)| p).minmax() { - MinMaxResult::MinMax(s, t) => *s..*t, - MinMaxResult::OneElement(x) => *x..(*x + 1u64), - MinMaxResult::NoElements => { - unreachable!("should have >0 elements") - } - } + half_open_interval(positions.iter().map(|(_, position)| *position)) }); let neg_interval = neg_positions.as_ref().map(|positions| { - match positions.iter().map(|(_, p)| p).minmax() { - MinMaxResult::MinMax(s, t) => *s..*t, - MinMaxResult::OneElement(x) => *x..(*x + 1u64), - MinMaxResult::NoElements => { - unreachable!("should have >0 elements") - } - } + half_open_interval(positions.iter().map(|(_, position)| *position)) }); #[cfg(debug_assertions)] @@ -185,7 +193,7 @@ impl GenomeWindow { } } - fn rightmost(&self) -> u64 { + fn exclusive_end(&self) -> u64 { match (self.end(&Strand::Positive), self.end(&Strand::Negative)) { (Some(x), Some(y)) => std::cmp::max(x, y), (Some(x), None) => x, @@ -328,7 +336,7 @@ impl GenomeWindow { self.add_pattern(&strand, pattern); } - fn get_mod_code_lookup(&self) -> FxHashMap { + fn get_mod_code_lookup(&self) -> FxHashMap { // looks complicated, but it just iterates over either the positive and // negative read patterns or the positive-combined read patterns let read_patterns: Box>> = @@ -343,9 +351,6 @@ impl GenomeWindow { } }; - // todo this could be done more simply with a set, but the idea is to - // make a single char code (e.g. '1', '2', '3', etc. for each - // modification code read_patterns .flat_map(|pattern| { pattern.iter().filter_map(|call| match call { @@ -358,11 +363,17 @@ impl GenomeWindow { .enumerate() .map(|(id, code)| { // save 0 for canonical - let id = id.saturating_add(1); - let encoded = format!("{id}").parse::().unwrap(); - (code, encoded) + let id = id + .checked_add(1) + .and_then(|id| u32::try_from(id).ok()) + .filter(|id| *id < u32::MAX) + .expect( + "modification state count must fit below the reserved \ + wildcard symbol", + ); + (code, EntropySymbol::called(id)) }) - .collect::>() + .collect::>() } fn encode_patterns( @@ -370,10 +381,10 @@ impl GenomeWindow { chrom_id: u32, strand: Strand, patterns: &Vec>, - mod_code_lookup: &FxHashMap, + mod_code_lookup: &FxHashMap, position_valid_coverages: &[u32], min_coverage: u32, - ) -> MkResult> { + ) -> MkResult> { // todo remove these checks after testing assert!( self.start(&strand).is_some(), @@ -393,18 +404,21 @@ impl GenomeWindow { let pattern = pat .iter() .map(|call| match call { - BaseModCall::Canonical(_) => '0', + BaseModCall::Canonical(_) => { + EntropySymbol::CANONICAL + } BaseModCall::Modified(_, code) => { *mod_code_lookup.get(code).unwrap() } - BaseModCall::Filtered => '*', + BaseModCall::Filtered => EntropySymbol::FILTERED, }) - .collect::(); + .collect::>() + .into_boxed_slice(); // todo remove after testing assert_eq!( pattern.len(), position_valid_coverages.len(), - "pattern {pattern} is the wrong size? \ + "pattern {pattern:?} is the wrong size? \ {position_valid_coverages:?}" ); pattern @@ -524,7 +538,7 @@ impl GenomeWindow { calc_me_entropy(&patterns, window_size, constant); let num_reads = patterns.len(); let interval = self.start(&Strand::Positive).unwrap() - ..self.end(&Strand::Positive).unwrap().saturating_add(1); + ..self.end(&Strand::Positive).unwrap(); MethylationEntropy::new(me_entropy, num_reads, interval) }) }); @@ -535,7 +549,7 @@ impl GenomeWindow { calc_me_entropy(&patterns, window_size, constant); let num_reads = patterns.len(); let interval = self.start(&Strand::Negative).unwrap() - ..self.end(&Strand::Negative).unwrap().saturating_add(1); + ..self.end(&Strand::Negative).unwrap(); MethylationEntropy::new(me_entropy, num_reads, interval) }) }); @@ -578,17 +592,18 @@ impl GenomeWindows { } fn get_range(&self) -> Range { - // these expects are checked in a few places, make them .unwrap()s let start = self .entropy_windows - .first() - .expect("self.entropy_windows should not be empty") - .leftmost(); + .iter() + .map(GenomeWindow::leftmost) + .min() + .expect("self.entropy_windows should not be empty"); let end = self .entropy_windows - .last() - .expect("self.entropy_windows should not be empty") - .rightmost(); + .iter() + .map(GenomeWindow::exclusive_end) + .max() + .expect("self.entropy_windows should not be empty"); start..end } @@ -657,25 +672,24 @@ impl GenomeWindows { ); // if neg_entropies is empty and there are no fails, we never saw // any negative strand me entropies - let neg_entropy_stats = if neg_entropies.is_empty() - && neg_num_fails == 0 - { - assert!( + let neg_entropy_stats = + if neg_entropies.is_empty() && neg_num_fails == 0 { + assert!( neg_num_reads.is_empty(), "neg num reads and window entropies should both be empty" ); - None - } else { - // this will fail correctly if there are neg_entropies is empty - // but there are fails - Some(DescriptiveStats::new( - &neg_entropies, - &neg_num_reads, - neg_num_fails, - chrom_id, - &interval, - )) - }; + None + } else { + // this will fail correctly if there are neg_entropies is empty + // but there are fails + Some(DescriptiveStats::new( + &neg_entropies, + &neg_num_reads, + neg_num_fails, + chrom_id, + &interval, + )) + }; let region_entropy = RegionEntropy::new( chrom_id, @@ -692,33 +706,682 @@ impl GenomeWindows { } } -#[derive(new)] +#[derive(Debug, new, Clone, Copy, PartialEq, Eq)] struct MotifHit { pos: u64, neg_position: Option, strand: Strand, base: DnaBase, + motif_idx: usize, +} + +#[derive(Debug, Clone)] +struct ReferenceSearchSpace { + record: ReferenceRecord, + sequence: Arc>, + owner: Range, +} + +#[derive(Debug)] +struct ScannedStripe { + hits: Vec, + #[cfg_attr(not(test), allow(dead_code))] + raw_hit_count: usize, +} + +#[derive(Debug)] +struct BufferedStripe { + hits: Vec, + next_hit: usize, +} + +impl BufferedStripe { + fn new(hits: Vec) -> Self { + Self { hits, next_hit: 0 } + } + + fn pop(&mut self) -> Option { + let hit = self.hits.get(self.next_hit).copied()?; + self.next_hit += 1; + Some(hit) + } + + fn is_empty(&self) -> bool { + self.next_hit >= self.hits.len() + } + + #[cfg(test)] + fn remaining(&self) -> usize { + self.hits.len().saturating_sub(self.next_hit) + } +} + +#[derive(Debug, Clone, Copy)] +struct PartnerAnchor { + positive_position: u64, + motif_idx: usize, +} + +#[derive(Debug)] +struct CombinedPartnerValidator { + max_anchor_distance: u64, + partner_to_positive: BTreeMap, + anchors_in_range: VecDeque<(u64, BaseAndPosition)>, +} + +impl CombinedPartnerValidator { + fn new(motif_search_context: usize) -> Self { + let max_anchor_distance = + u64::try_from(motif_search_context.saturating_mul(2)) + .unwrap_or(u64::MAX); + Self { + max_anchor_distance, + partner_to_positive: BTreeMap::new(), + anchors_in_range: VecDeque::new(), + } + } + + fn validate( + &mut self, + hit: &MotifHit, + motifs: &[RegexMotif], + reference_name: &str, + ) -> anyhow::Result<()> { + debug_assert_eq!(hit.strand, Strand::Positive); + let Some(negative_position) = hit.neg_position else { + return Ok(()); + }; + + // For a palindromic motif of length m and focus f, the partner + // displacement is d = m - 1 - 2f, so |d| <= C where C is the maximum + // motif length minus one. If two anchors share a partner, their + // distance is therefore at most 2C. Older entries cannot collide with + // this or any future (monotonically ordered) positive anchor. + while let Some((positive_position, partner)) = + self.anchors_in_range.front().copied() + { + if positive_position.saturating_add(self.max_anchor_distance) + >= hit.pos + { + break; + } + self.anchors_in_range.pop_front(); + if self.partner_to_positive.get(&partner).is_some_and(|anchor| { + anchor.positive_position == positive_position + }) { + self.partner_to_positive.remove(&partner); + } + } + + let negative_partner = (hit.base, negative_position); + if let Some(previous) = + self.partner_to_positive.get(&negative_partner).copied() + { + if previous.positive_position != hit.pos { + bail!( + "conflicting combined-strand motif anchors at \ + {reference_name}: negative {} partner {} maps to positive \ + anchors {} from motif {} and {} from motif {}", + hit.base, + negative_position, + previous.positive_position, + motifs[previous.motif_idx], + hit.pos, + motifs[hit.motif_idx], + ) + } + return Ok(()); + } + + self.partner_to_positive.insert( + negative_partner, + PartnerAnchor { + positive_position: hit.pos, + motif_idx: hit.motif_idx, + }, + ); + self.anchors_in_range.push_back((hit.pos, negative_partner)); + Ok(()) + } +} + +#[derive(Debug)] +struct OrderedMotifScanner { + sequence: Arc>, + owner: Range, + reference_start: u64, + reference_name: String, + motifs: Arc>, + strand_filter: Option, + combine_strands: bool, + motif_search_context: usize, + chunk_size: usize, + stripes_per_batch: usize, + next_stripe_start: usize, + pending_stripes: VecDeque, + partner_validator: Option, + #[cfg(test)] + current_batch_raw_hits: usize, +} + +// Bounded motif-scan pipeline invariant: +// - OrderedMotifScanner emits locally sorted, deduplicated owner stripes in +// reference order and retains at most one fixed-size stripe batch. +// - StrandHitStream retains at most the requested N-hit lookahead; separate +// monotonic strand streams prevent a dense strand from accumulating while +// the other strand searches ahead. +// - CombinedPartnerValidator needs only the preceding 2C anchor span, where C +// is the maximum motif-search context. +// Thus motif-hit storage depends on N and the fixed stripe/context settings, +// never on owner length or requested window width. +impl OrderedMotifScanner { + fn motif_search_context(motifs: &[RegexMotif]) -> usize { + motifs + .iter() + .map(|motif| motif.length().saturating_sub(1)) + .max() + .unwrap_or(0) + } + + fn new( + search_space: &ReferenceSearchSpace, + motifs: Arc>, + strand_filter: Option, + combine_strands: bool, + chunk_size: usize, + stripes_per_batch: usize, + ) -> Self { + assert!(chunk_size > 0, "motif search chunk size must be positive"); + assert!( + stripes_per_batch > 0, + "motif stripe batch size must be positive" + ); + assert!(search_space.owner.start <= search_space.owner.end); + assert!(search_space.owner.end <= search_space.sequence.len()); + let motif_search_context = + Self::motif_search_context(motifs.as_slice()); + let partner_validator = (combine_strands && motifs.len() > 1) + .then(|| CombinedPartnerValidator::new(motif_search_context)); + Self { + sequence: search_space.sequence.clone(), + owner: search_space.owner.clone(), + reference_start: search_space.record.start as u64, + reference_name: search_space.record.name.clone(), + motifs, + strand_filter, + combine_strands, + motif_search_context, + chunk_size, + stripes_per_batch, + next_stripe_start: search_space.owner.start, + pending_stripes: VecDeque::new(), + partner_validator, + #[cfg(test)] + current_batch_raw_hits: 0, + } + } + + fn find_motif_hits_in_owner( + seq: &[char], + motifs: &[RegexMotif], + owner: Range, + reference_start: u64, + motif_search_context: usize, + strand_filter: Option, + ) -> (Vec, usize) { + assert!(owner.start <= owner.end); + assert!(owner.end <= seq.len()); + if owner.is_empty() { + return (Vec::new(), 0); + } + + let search_start = owner.start.saturating_sub(motif_search_context); + let search_end = + owner.end.saturating_add(motif_search_context).min(seq.len()); + let subseq = seq[search_start..search_end].iter().collect::(); + let mut motif_hits = Vec::new(); + let mut raw_hit_count = 0usize; + for (motif_idx, motif) in motifs.iter().enumerate() { + let raw_motif_hits = motif.find_hits(&subseq); + raw_hit_count = raw_hit_count.saturating_add(raw_motif_hits.len()); + for (search_position, strand) in raw_motif_hits { + if strand_filter.is_some() && strand_filter != Some(strand) { + continue; + } + let Some(local_position) = + search_start.checked_add(search_position) + else { + continue; + }; + if !owner.contains(&local_position) { + continue; + } + let position = reference_start + .checked_add(local_position as u64) + .expect("reference position overflow"); + let dna_base = DnaBase::parse(seq[local_position]) + .expect("motif anchor must be a DNA base"); + let base = if strand == Strand::Negative { + dna_base.complement() + } else { + dna_base + }; + let neg_position = motif + .motif_info + .negative_strand_position(position as u32) + .map(|position| position as u64); + motif_hits.push(MotifHit::new( + position, + neg_position, + strand, + base, + motif_idx, + )); + } + } + (motif_hits, raw_hit_count) + } + + fn sort_and_dedup_stripe( + mut motif_hits: Vec, + raw_hit_count: usize, + motifs: &[RegexMotif], + combine_strands: bool, + reference_name: &str, + ) -> anyhow::Result { + motif_hits + .sort_by_key(|hit| (hit.pos, hit.strand, hit.base, hit.motif_idx)); + let mut write_idx = 0usize; + for read_idx in 0..motif_hits.len() { + let hit = motif_hits[read_idx]; + if write_idx > 0 { + let previous = motif_hits[write_idx - 1]; + let same_anchor = previous.pos == hit.pos + && previous.strand == hit.strand + && previous.base == hit.base; + if same_anchor { + if combine_strands + && hit.strand == Strand::Positive + && previous.neg_position != hit.neg_position + { + bail!( + "conflicting combined-strand motif partners at \ + {reference_name}:{} for {} anchor: {:?} from \ + motif {} versus {:?} from motif {}", + hit.pos, + hit.base, + previous.neg_position, + motifs[previous.motif_idx], + hit.neg_position, + motifs[hit.motif_idx], + ) + } + continue; + } + } + if write_idx != read_idx { + motif_hits.swap(write_idx, read_idx); + } + write_idx += 1; + } + motif_hits.truncate(write_idx); + Ok(ScannedStripe { hits: motif_hits, raw_hit_count }) + } + + fn fill_stripe_batch(&mut self) -> anyhow::Result { + if self.next_stripe_start >= self.owner.end { + return Ok(false); + } + debug_assert!(self.pending_stripes.is_empty()); + let mut owner_chunks = Vec::with_capacity(self.stripes_per_batch); + for _ in 0..self.stripes_per_batch { + if self.next_stripe_start >= self.owner.end { + break; + } + let start = self.next_stripe_start; + let end = start.saturating_add(self.chunk_size).min(self.owner.end); + owner_chunks.push(start..end); + self.next_stripe_start = end; + } + + let scanned_stripes = owner_chunks + .into_par_iter() + .map(|owner_chunk| { + let (hits, raw_hit_count) = Self::find_motif_hits_in_owner( + &self.sequence, + self.motifs.as_slice(), + owner_chunk, + self.reference_start, + self.motif_search_context, + self.strand_filter, + ); + Self::sort_and_dedup_stripe( + hits, + raw_hit_count, + self.motifs.as_slice(), + self.combine_strands, + &self.reference_name, + ) + }) + .collect::>(); + + #[cfg(test)] + { + self.current_batch_raw_hits = scanned_stripes + .iter() + .filter_map(|result| result.as_ref().ok()) + .map(|stripe| stripe.raw_hit_count) + .sum(); + } + for scanned_stripe in scanned_stripes { + let ScannedStripe { hits, raw_hit_count: _ } = scanned_stripe?; + if !hits.is_empty() { + self.pending_stripes.push_back(BufferedStripe::new(hits)); + } + } + Ok(true) + } + + fn next_hit(&mut self) -> anyhow::Result> { + loop { + while self + .pending_stripes + .front() + .is_some_and(|stripe| stripe.is_empty()) + { + self.pending_stripes.pop_front(); + } + if let Some(hit) = + self.pending_stripes.front_mut().and_then(BufferedStripe::pop) + { + if let Some(validator) = self.partner_validator.as_mut() { + validator.validate( + &hit, + self.motifs.as_slice(), + &self.reference_name, + )?; + } + return Ok(Some(hit)); + } + + #[cfg(test)] + { + self.current_batch_raw_hits = 0; + } + if !self.fill_stripe_batch()? { + return Ok(None); + } + } + } + + fn is_exhausted(&self) -> bool { + self.next_stripe_start >= self.owner.end + && self.pending_stripes.iter().all(BufferedStripe::is_empty) + } + + #[cfg(test)] + fn live_hits_for_test(&self) -> usize { + let pending_hits = self + .pending_stripes + .iter() + .map(BufferedStripe::remaining) + .sum::(); + pending_hits.saturating_add(self.current_batch_raw_hits) + } +} + +#[derive(Debug)] +struct StrandHitStream { + scanner: OrderedMotifScanner, + lookahead: VecDeque, + minimum_position: u64, +} + +impl StrandHitStream { + fn new( + search_space: &ReferenceSearchSpace, + motifs: Arc>, + strand: Strand, + combine_strands: bool, + chunk_size: usize, + stripes_per_batch: usize, + ) -> Self { + let minimum_position = (search_space.record.start as u64) + .saturating_add(search_space.owner.start as u64); + Self { + scanner: OrderedMotifScanner::new( + search_space, + motifs, + Some(strand), + combine_strands, + chunk_size, + stripes_per_batch, + ), + lookahead: VecDeque::new(), + minimum_position, + } + } + + fn fill_to_len(&mut self, len: usize) -> anyhow::Result<()> { + while self.lookahead.len() < len { + let Some(hit) = self.scanner.next_hit()? else { + break; + }; + if hit.pos >= self.minimum_position { + self.lookahead.push_back(hit); + } + } + Ok(()) + } + + fn discard_before(&mut self, minimum_position: u64) { + self.minimum_position = minimum_position; + while self + .lookahead + .front() + .is_some_and(|hit| hit.pos < minimum_position) + { + self.lookahead.pop_front(); + } + } + + fn first_after(&mut self, position: u64) -> anyhow::Result> { + let minimum_position = + position.checked_add(1).expect("reference position overflow"); + self.discard_before(minimum_position); + self.fill_to_len(1)?; + Ok(self.lookahead.front().map(|hit| hit.pos)) + } + + fn hits_before(&self, end: u64, limit: usize) -> Vec { + self.lookahead + .iter() + .take(limit) + .take_while(|hit| hit.pos < end) + .copied() + .collect() + } + + fn first_position(&self) -> Option { + self.lookahead.front().map(|hit| hit.pos) + } + + fn is_exhausted(&self) -> bool { + self.lookahead.is_empty() && self.scanner.is_exhausted() + } + + #[cfg(test)] + fn live_hits_for_test(&self) -> usize { + self.lookahead.len() + self.scanner.live_hits_for_test() + } +} + +#[derive(Debug)] +struct ActiveReference { + record: ReferenceRecord, + owner: Range, + positive_hits: StrandHitStream, + negative_hits: Option, +} + +impl ActiveReference { + fn new( + search_space: ReferenceSearchSpace, + motifs: Arc>, + combine_strands: bool, + chunk_size: usize, + stripes_per_batch: usize, + ) -> Self { + let positive_hits = StrandHitStream::new( + &search_space, + motifs.clone(), + Strand::Positive, + combine_strands, + chunk_size, + stripes_per_batch, + ); + let negative_hits = (!combine_strands).then(|| { + StrandHitStream::new( + &search_space, + motifs, + Strand::Negative, + false, + chunk_size, + stripes_per_batch, + ) + }); + Self { + record: search_space.record, + owner: search_space.owner, + positive_hits, + negative_hits, + } + } + + fn fill_lookahead(&mut self, len: usize) -> anyhow::Result<()> { + self.positive_hits.fill_to_len(len)?; + if let Some(negative_hits) = self.negative_hits.as_mut() { + negative_hits.fill_to_len(len)?; + } + Ok(()) + } + + fn first_position(&mut self) -> anyhow::Result> { + self.fill_lookahead(1)?; + Ok(self + .positive_hits + .first_position() + .into_iter() + .chain( + self.negative_hits + .as_ref() + .and_then(StrandHitStream::first_position), + ) + .min()) + } + + fn window_hits( + &self, + end: u64, + limit: usize, + ) -> (Vec, Vec) { + let positive_hits = self.positive_hits.hits_before(end, limit); + let negative_hits = self + .negative_hits + .as_ref() + .map(|hits| hits.hits_before(end, limit)) + .unwrap_or_default(); + (positive_hits, negative_hits) + } + + fn discard_before(&mut self, minimum_position: u64) { + self.positive_hits.discard_before(minimum_position); + if let Some(negative_hits) = self.negative_hits.as_mut() { + negative_hits.discard_before(minimum_position); + } + } + + fn first_after(&mut self, position: u64) -> anyhow::Result> { + let positive = self.positive_hits.first_after(position)?; + let negative = match self.negative_hits.as_mut() { + Some(hits) => hits.first_after(position)?, + None => None, + }; + Ok(positive.into_iter().chain(negative).min()) + } + + fn is_exhausted(&self) -> bool { + self.positive_hits.is_exhausted() + && self + .negative_hits + .as_ref() + .is_none_or(StrandHitStream::is_exhausted) + } + + #[cfg(test)] + fn live_hits_for_test(&self) -> usize { + self.positive_hits.live_hits_for_test() + + self + .negative_hits + .as_ref() + .map(StrandHitStream::live_hits_for_test) + .unwrap_or(0) + } } struct SlidingWindows { - motifs: Vec, - work_queue: VecDeque<(ReferenceRecord, Vec)>, + motifs: Arc>, + work_queue: VecDeque, region_names: VecDeque, window_size: usize, num_positions: usize, batch_size: usize, curr_position: usize, - curr_contig: ReferenceRecord, - curr_seq: Vec, + curr_reference: Option, curr_region_name: Option, combine_strands: bool, - /// the longest motif length, so we find motifs that are in the window, but - /// reach outside the window - motif_search_adj: usize, + scan_chunk_size: usize, + stripes_per_batch: usize, done: bool, + #[cfg(test)] + high_water_retained_hits: usize, } impl SlidingWindows { + const START_SEARCH_CHUNK_SIZE: usize = 10_000; + const STRIPES_PER_BATCH: usize = 32; + + fn motif_search_context(motifs: &[RegexMotif]) -> usize { + OrderedMotifScanner::motif_search_context(motifs) + } + + fn preflight_combined_partner_conflicts( + work_queue: &VecDeque, + motifs: Arc>, + combine_strands: bool, + chunk_size: usize, + stripes_per_batch: usize, + ) -> anyhow::Result<()> { + if !(combine_strands && motifs.len() > 1) { + return Ok(()); + } + + // Consume the ordered positive stream for every owner before creating + // the writer. Partner validation uses only the bounded 2C carry. + for search_space in work_queue { + let mut scanner = OrderedMotifScanner::new( + search_space, + motifs.clone(), + Some(Strand::Positive), + true, + chunk_size, + stripes_per_batch, + ); + while scanner.next_hit()?.is_some() {} + } + Ok(()) + } + fn new_with_regions( reference_sequences_lookup: ReferenceSequencesLookup, regions_bed_fp: &PathBuf, @@ -728,6 +1391,7 @@ impl SlidingWindows { window_size: usize, batch_size: usize, ) -> anyhow::Result { + let motif_search_context = Self::motif_search_context(&motifs); let regions_iter = BufReader::new(File::open(regions_bed_fp).with_context(|| { format!("failed to load regions at {regions_bed_fp:?}") @@ -745,28 +1409,47 @@ impl SlidingWindows { // lines .map(|r| { r.and_then(|bed_region| { - let start = bed_region.interval.start; - let end = bed_region.interval.end; - let interval = start..end; - reference_sequences_lookup - .get_subsequence_by_name( - bed_region.chrom.as_str(), - interval, + let chrom = bed_region.chrom.as_str(); + let sequence_length = reference_sequences_lookup + .sequence_length_by_name(chrom)?; + if bed_region.interval.end > sequence_length { + bail!( + "region {}:{}-{} exceeds reference length {}", + bed_region.chrom, + bed_region.interval.start, + bed_region.interval.end, + sequence_length ) - .map(|seq| (bed_region, seq)) + } + let search_start = bed_region + .interval + .start + .saturating_sub(motif_search_context); + let search_end = bed_region + .interval + .end + .saturating_add(motif_search_context) + .min(sequence_length); + let owner = bed_region.interval.start - search_start + ..bed_region.interval.end - search_start; + let seq = reference_sequences_lookup + .get_subsequence_by_name( + chrom, + search_start..search_end, + )?; + let tid = reference_sequences_lookup + .name_to_chrom_id(chrom) + .ok_or_else(|| { + anyhow!("missing reference ID for {chrom}") + })?; + let reference_record = ReferenceRecord::new( + tid, + search_start as u32, + seq.len() as u32, + bed_region.chrom.clone(), + ); + Ok((reference_record, bed_region.name, seq, owner)) }) - }) - .map_ok(|(bed_region, seq)| { - let tid = reference_sequences_lookup - .name_to_chrom_id(bed_region.chrom.as_str()) - .unwrap(); - let start = bed_region.interval.start as u32; - let length = bed_region.length() as u32; - let chrom_name = bed_region.chrom; - let region_name = bed_region.name; - let reference_record = - ReferenceRecord::new(tid, start, length, chrom_name); - (reference_record, region_name, seq) }); // accumulators for the above iterator, could have done this all in a @@ -782,8 +1465,12 @@ impl SlidingWindows { for res in regions_iter { match res { - Ok((reference_record, region_name, subseq)) => { - work_queue.push_back((reference_record, subseq)); + Ok((reference_record, region_name, subseq, owner)) => { + work_queue.push_back(ReferenceSearchSpace { + record: reference_record, + sequence: Arc::new(subseq), + owner, + }); region_queue.push_back(region_name); } Err(e) => { @@ -806,60 +1493,18 @@ impl SlidingWindows { } assert_eq!(region_queue.len(), work_queue.len()); - let (curr_contig, curr_seq, curr_position, curr_region_name) = loop { - let (ref_record, subseq, region_name) = - match (work_queue.pop_front(), region_queue.pop_front()) { - (Some((rr, subseq)), Some(region_name)) => { - anyhow::Ok((rr, subseq, region_name)) - } - _ => bail!( - "didn't find at least 1 sequence with valid start \ - position" - ), - }?; - if let Some(start_position) = - Self::find_start_position(&subseq, &motifs) - { - info!( - "starting with region {region_name} at 0-based position \ - {} on contig {}", - start_position + ref_record.start as usize, - &ref_record.name - ); - break (ref_record, subseq, start_position, region_name); - } else { - info!("region {region_name} has no valid positions, skipping"); - continue; - } - }; - debug!( - "parsed {} regions, starting with {} on contig {}", - region_queue.len() + 1usize, - &curr_region_name, - curr_contig.name - ); - let motif_search_adj = motifs - .iter() - .map(|motif| motif.length()) - .filter(|l| *l > 1) - .max() - .unwrap_or(0); - - Ok(Self { - motifs, + Self::new_from_work_queue( work_queue, - region_names: region_queue, - window_size, + region_queue, + motifs, + combine_strands, num_positions, + window_size, batch_size, - curr_position, - curr_contig, - curr_seq, - curr_region_name: Some(curr_region_name), - combine_strands, - motif_search_adj, - done: false, - }) + Self::START_SEARCH_CHUNK_SIZE, + Self::STRIPES_PER_BATCH, + true, + ) } fn new( @@ -870,52 +1515,170 @@ impl SlidingWindows { window_size: usize, batch_size: usize, ) -> anyhow::Result { - let mut work_queue = - reference_sequence_lookup.into_reference_sequences(); - - let (curr_contig, curr_seq, curr_position) = loop { - let (curr_record, curr_seq) = - work_queue.pop_front().ok_or_else(|| { - anyhow!( - "didn't find at least 1 sequence with a valid start \ - position" - ) + let work_queue = reference_sequence_lookup + .into_reference_sequences() + .into_iter() + .map(|(record, sequence)| { + let owner = 0..sequence.len(); + ReferenceSearchSpace { + record, + sequence: Arc::new(sequence), + owner, + } + }) + .collect::>(); + Self::new_from_work_queue( + work_queue, + VecDeque::new(), + motifs, + combine_strands, + num_positions, + window_size, + batch_size, + Self::START_SEARCH_CHUNK_SIZE, + Self::STRIPES_PER_BATCH, + false, + ) + } + + #[allow(clippy::too_many_arguments)] + fn new_from_work_queue( + mut work_queue: VecDeque, + mut region_names: VecDeque, + motifs: Vec, + combine_strands: bool, + num_positions: usize, + window_size: usize, + batch_size: usize, + scan_chunk_size: usize, + stripes_per_batch: usize, + regions_mode: bool, + ) -> anyhow::Result { + if regions_mode { + assert_eq!(region_names.len(), work_queue.len()); + } else { + assert!(region_names.is_empty()); + } + let motifs = Arc::new(motifs); + Self::preflight_combined_partner_conflicts( + &work_queue, + motifs.clone(), + combine_strands, + scan_chunk_size, + stripes_per_batch, + )?; + + let (curr_reference, curr_position, curr_region_name) = loop { + let search_space = work_queue.pop_front().ok_or_else(|| { + anyhow!( + "didn't find at least 1 sequence with a valid start \ + position" + ) + })?; + let region_name = if regions_mode { + Some(region_names.pop_front().expect( + "region names must remain aligned with search spaces", + )) + } else { + None + }; + let mut active = ActiveReference::new( + search_space, + motifs.clone(), + combine_strands, + scan_chunk_size, + stripes_per_batch, + ); + let first_position = + active.first_position().with_context(|| { + region_name + .as_ref() + .map(|name| { + format!("failed to scan entropy region {name}") + }) + .unwrap_or_else(|| { + format!( + "failed to scan entropy contig {}", + active.record.name + ) + }) })?; - if let Some(pos) = Self::find_start_position(&curr_seq, &motifs) { - info!( - "starting with contig {} at 0-based position {pos}", - &curr_record.name - ); - break (curr_record, curr_seq, pos); + if let Some(first_position) = first_position { + let local_position = first_position + .checked_sub(active.record.start as u64) + .and_then(|position| usize::try_from(position).ok()) + .expect("motif hit must be local to its reference slice"); + if let Some(name) = region_name.as_ref() { + info!( + "starting with region {name} at 0-based position \ + {first_position} on contig {}", + active.record.name + ); + } else { + info!( + "starting with contig {} at 0-based position \ + {local_position}", + active.record.name + ); + } + break (active, local_position, region_name); + } + + if let Some(name) = region_name { + info!("region {name} has no valid positions, skipping"); } else { info!( "contig {} had no valid motif positions, skipping..", - curr_record.name + active.record.name ); } }; - let motif_search_adj = motifs - .iter() - .map(|motif| motif.length()) - .filter(|l| *l > 1) - .max() - .unwrap_or(0); - Ok(Self { + let mut windows = Self { motifs, work_queue, - region_names: VecDeque::new(), + region_names, window_size, num_positions, batch_size, curr_position, - curr_contig, - curr_seq, - curr_region_name: None, + curr_reference: Some(curr_reference), + curr_region_name, combine_strands, - motif_search_adj, + scan_chunk_size, + stripes_per_batch, done: false, - }) + #[cfg(test)] + high_water_retained_hits: 0, + }; + windows.observe_retained_hits(); + Ok(windows) + } + + #[cfg(test)] + #[allow(clippy::too_many_arguments)] + fn new_for_test( + work_queue: VecDeque, + motifs: Vec, + combine_strands: bool, + num_positions: usize, + window_size: usize, + batch_size: usize, + scan_chunk_size: usize, + stripes_per_batch: usize, + ) -> anyhow::Result { + Self::new_from_work_queue( + work_queue, + VecDeque::new(), + motifs, + combine_strands, + num_positions, + window_size, + batch_size, + scan_chunk_size, + stripes_per_batch, + false, + ) } #[inline] @@ -924,7 +1687,7 @@ impl SlidingWindows { motif_hits: &[MotifHit], ) -> Option> { let positions = motif_hits - .into_iter() + .iter() .take(self.num_positions) .map(|mh| (mh.base, mh.pos)) .sorted_by(|(_, a), (_, b)| a.cmp(b)) @@ -941,11 +1704,15 @@ impl SlidingWindows { &self, pos_hits: &[MotifHit], neg_hits: &[MotifHit], - ) -> Option { + ) -> Option<(GenomeWindow, u64)> { if self.combine_strands { + let next_reference_position = pos_hits + .first()? + .pos + .checked_add(1) + .expect("reference position overflow"); let neg_to_pos = pos_hits - .into_iter() - .filter(|x| x.strand == Strand::Positive) + .iter() .take(self.num_positions) .filter_map(|motif_hit| { assert_eq!( @@ -961,21 +1728,19 @@ impl SlidingWindows { if neg_to_pos.len() < self.num_positions { None } else { - let (start, end) = match neg_to_pos - .keys() - .chain(neg_to_pos.values()) - .map(|(_, x)| x) - .minmax() - { - MinMaxResult::MinMax(s, t) => (*s, *t), - MinMaxResult::OneElement(x) => (*x, *x + 1u64), /* should probably fail here too? */ - _ => unreachable!("there must be more than 1 element"), - }; - let interval = start..end; - Some(GenomeWindow::new_combine_strands( - interval, - self.num_positions, - neg_to_pos, + let interval = half_open_interval( + neg_to_pos + .keys() + .chain(neg_to_pos.values()) + .map(|(_, position)| *position), + ); + Some(( + GenomeWindow::new_combine_strands( + interval, + self.num_positions, + neg_to_pos, + ), + next_reference_position, )) } } else { @@ -1003,19 +1768,29 @@ impl SlidingWindows { if leftmost_positive_ref_pos < leftmost_negative_ref_pos { // debug!("(+) is lefter, using {p:?}"); - Some(GenomeWindow::new_stranded( - Some(p), - None, - self.num_positions, + Some(( + GenomeWindow::new_stranded( + Some(p), + None, + self.num_positions, + ), + leftmost_positive_ref_pos + .checked_add(1) + .expect("reference position overflow"), )) } else if leftmost_negative_ref_pos < leftmost_positive_ref_pos { // debug!("(-) is lefter, using {n:?}"); - Some(GenomeWindow::new_stranded( - None, - Some(n), - self.num_positions, + Some(( + GenomeWindow::new_stranded( + None, + Some(n), + self.num_positions, + ), + leftmost_negative_ref_pos + .checked_add(1) + .expect("reference position overflow"), )) } else { assert_eq!( @@ -1024,27 +1799,46 @@ impl SlidingWindows { ); // debug!("they are the same, using {p:?} and // {n:?}"); - Some(GenomeWindow::new_stranded( - Some(p), - Some(n), - self.num_positions, + Some(( + GenomeWindow::new_stranded( + Some(p), + Some(n), + self.num_positions, + ), + leftmost_positive_ref_pos + .checked_add(1) + .expect("reference position overflow"), )) } } (Some(p), None) => { // debug!("(+) only, using {p:?}"); - Some(GenomeWindow::new_stranded( - Some(p), - None, - self.num_positions, + let next_reference_position = p[0] + .1 + .checked_add(1) + .expect("reference position overflow"); + Some(( + GenomeWindow::new_stranded( + Some(p), + None, + self.num_positions, + ), + next_reference_position, )) } (None, Some(n)) => { // debug!("(-) only, using {n:?}"); - Some(GenomeWindow::new_stranded( - None, - Some(n), - self.num_positions, + let next_reference_position = n[0] + .1 + .checked_add(1) + .expect("reference position overflow"); + Some(( + GenomeWindow::new_stranded( + None, + Some(n), + self.num_positions, + ), + next_reference_position, )) } _ => None, @@ -1055,197 +1849,302 @@ impl SlidingWindows { } } + fn advance_cursor(&mut self, next_reference_position: u64) { + let curr_reference = self + .curr_reference + .as_ref() + .expect("active reference must exist while advancing"); + let next_local_position = next_reference_position + .checked_sub(curr_reference.record.start as u64) + .and_then(|position| usize::try_from(position).ok()) + .expect("next motif cursor must be local to its reference slice"); + assert!( + next_local_position > self.curr_position, + "motif cursor must advance: {} -> {}", + self.curr_position, + next_local_position + ); + assert!( + next_local_position <= curr_reference.owner.end, + "motif cursor must remain inside its owner" + ); + self.curr_position = next_local_position; + self.curr_reference + .as_mut() + .expect("active reference must exist while advancing") + .discard_before(next_reference_position); + self.observe_retained_hits(); + } + fn next_window(&mut self) -> Option { while !self.at_end_of_contig() { - // search forward for hits - let end = std::cmp::min( - self.curr_position.saturating_add(self.window_size), - self.curr_seq.len(), - ); - // todo optimize? - // debug!( - // "genome space position at top {}, {}, {}", - // self.curr_position + self.curr_contig.start as usize, - // self.curr_position, - // self.motif_search_adj - // ); - let subseq_start = - self.curr_position.saturating_sub(self.motif_search_adj); - let offset = self.curr_position.checked_sub(subseq_start).expect( - "curr_position should always be greater than subset_start", - ); - let subseq = self.curr_seq[subseq_start..end] - .iter() - .map(|x| *x) - .collect::(); - // debug!("subseq at the top {subseq}"); - // N.B. the 'position' in these tuples are _genome coordinates_! - // this is because when we fetch reads we need to do it with the - // proper genome coordinates. when we're using normal - // sliding windows, the relative coordinates and the - // genome coordinates _should_ be the same however when - // using regions, we slice the reference genome, so the - // relative (to the sequence) and genome coordinates will _not_ be - // the same - let (pos_hits, neg_hits): (Vec, Vec) = self - .motifs - .iter() - .flat_map(|motif| { - motif - .find_hits(&subseq) - .into_iter() - // this filter removes positions found before - // self.curr-position - .filter_map(|(pos, strand)| { - pos.checked_sub(offset).map(|p| (p, strand)) - }) - .map(|(pos, strand)| { - let adjusted_position = pos - .saturating_add(self.curr_position) - .saturating_add( - self.curr_contig.start as usize, - ); - let dna_base = DnaBase::parse( - self.curr_seq[pos + self.curr_position], - ) - .unwrap(); - let base = if strand == Strand::Negative { - dna_base.complement() - } else { - dna_base - }; - let neg_position = motif - .motif_info - .negative_strand_position( - adjusted_position as u32, - ) - .map(|x| x as u64); - MotifHit::new( - adjusted_position as u64, - neg_position, - strand, - base, - ) - }) - .collect::>() - }) - .sorted_by(|a, b| a.pos.cmp(&b.pos)) - .partition(|x| x.strand == Strand::Positive); - if let Some(entropy_window) = + let curr_reference = self + .curr_reference + .as_ref() + .expect("active reference must exist while scanning windows"); + let current_reference_position = curr_reference + .record + .start + .checked_add(self.curr_position as u32) + .expect("reference position overflow") + as u64; + let owner_end = curr_reference + .record + .start + .checked_add(curr_reference.owner.end as u32) + .expect("reference position overflow") + as u64; + let window_end = current_reference_position + .saturating_add(self.window_size as u64) + .min(owner_end); + self.curr_reference + .as_mut() + .expect("active reference must exist while scanning windows") + .fill_lookahead(self.num_positions) + .expect( + "combined motif partner conflicts must be rejected during \ + construction", + ); + self.observe_retained_hits(); + let (pos_hits, neg_hits) = self + .curr_reference + .as_ref() + .expect("active reference must exist while scanning windows") + .window_hits(window_end, self.num_positions); + + if let Some((entropy_window, next_reference_position)) = self.enough_hits_for_window(&pos_hits, &neg_hits) { - let new_genome_space_position = - (entropy_window.leftmost() as usize).saturating_add(1usize); - // info!("new genome position {new_genome_space_position}"); - // need to re-adjust to relative coordinates instead of genome - // coordinates - self.curr_position = new_genome_space_position - .checked_sub(self.curr_contig.start as usize) - .expect( - "should be able to subtract contig start from position", - ); - + self.advance_cursor(next_reference_position); return Some(entropy_window); - } else { - // not enough on (+) or (-) - let hits = pos_hits - .into_iter() - .chain(neg_hits) - .map(|mh| mh.pos as usize) - .map(|p| { - // need to re-adjust to relative coordinates instead of - // genome coordinates - p.checked_sub(self.curr_contig.start as usize) - .expect("should be able to re-adjust position") - }) - .collect::>(); - if let Some(&first) = hits.first() { - // at least 1 - if self.curr_position == first { - match hits.iter().nth(1) { - Some(&second_hit) => { - self.curr_position = second_hit - } - None => { - // there was only 1 - self.curr_position = end; - } - } - } else { - self.curr_position = first; - } - } else { - // hits was empty, set to end - self.curr_position = end; - } - continue; } + + let next_reference_position = self + .curr_reference + .as_mut() + .expect("active reference must exist while scanning windows") + .first_after(current_reference_position) + .expect( + "combined motif partner conflicts must be rejected during \ + construction", + ) + .unwrap_or(owner_end); + self.observe_retained_hits(); + self.advance_cursor(next_reference_position); } None } + #[cfg(test)] + fn sort_and_dedup_motif_hits( + motif_hits: Vec, + motifs: &[RegexMotif], + combine_strands: bool, + reference_name: &str, + ) -> anyhow::Result> { + let raw_hit_count = motif_hits.len(); + let ScannedStripe { hits, raw_hit_count: _ } = + OrderedMotifScanner::sort_and_dedup_stripe( + motif_hits, + raw_hit_count, + motifs, + combine_strands, + reference_name, + )?; + if combine_strands && motifs.len() > 1 { + let mut validator = CombinedPartnerValidator::new( + Self::motif_search_context(motifs), + ); + for hit in hits.iter().filter(|hit| hit.strand == Strand::Positive) + { + validator.validate(hit, motifs, reference_name)?; + } + } + Ok(hits) + } + + #[cfg(test)] + #[allow(clippy::too_many_arguments)] + fn scan_motif_hits_with_config_for_test( + seq: &[char], + motifs: &[RegexMotif], + owner: Range, + reference_start: u64, + reference_name: &str, + combine_strands: bool, + chunk_size: usize, + stripes_per_batch: usize, + strand_filter: Option, + ) -> anyhow::Result> { + let search_space = ReferenceSearchSpace { + record: ReferenceRecord::new( + 0, + reference_start as u32, + seq.len() as u32, + reference_name.to_string(), + ), + sequence: Arc::new(seq.to_vec()), + owner, + }; + let mut scanner = OrderedMotifScanner::new( + &search_space, + Arc::new(motifs.to_vec()), + strand_filter, + combine_strands, + chunk_size, + stripes_per_batch, + ); + let mut hits = Vec::new(); + while let Some(hit) = scanner.next_hit()? { + hits.push(hit); + } + Ok(hits) + } + + #[cfg(test)] + fn scan_motif_hits_with_chunk_size( + seq: &[char], + motifs: &[RegexMotif], + owner: Range, + reference_start: u64, + reference_name: &str, + combine_strands: bool, + chunk_size: usize, + ) -> anyhow::Result> { + Self::scan_motif_hits_with_config_for_test( + seq, + motifs, + owner, + reference_start, + reference_name, + combine_strands, + chunk_size, + Self::STRIPES_PER_BATCH, + None, + ) + } + + #[cfg(test)] fn find_start_position( seq: &[char], motifs: &[RegexMotif], ) -> Option { - seq.par_chunks(10_000).find_map_first(|c| { - let s = c.iter().collect::(); - let min_pos = motifs - .iter() - .flat_map(|motif| { - motif.find_hits(&s).into_iter().nth(0).map(|(pos, _)| pos) - }) - .min(); - min_pos - }) + Self::scan_motif_hits_with_chunk_size( + seq, + motifs, + 0..seq.len(), + 0, + "reference", + false, + Self::START_SEARCH_CHUNK_SIZE, + ) + .expect("non-combined motif scan cannot have partner conflicts") + .first() + .and_then(|hit| usize::try_from(hit.pos).ok()) } #[inline] fn at_end_of_contig(&self) -> bool { - self.curr_position >= self.curr_contig.length as usize + self.curr_reference.as_ref().is_none_or(|active| { + self.curr_position >= active.owner.end || active.is_exhausted() + }) } fn update_current_contig(&mut self) { + self.observe_retained_hits(); + drop(self.curr_reference.take()); 'search: loop { - if let Some((record, seq)) = self.work_queue.pop_front() { - match Self::find_start_position(&seq, &self.motifs) { - Some(start_pos) => { - self.curr_contig = record; - self.curr_position = start_pos; - self.curr_seq = seq; - let region_name = self.region_names.pop_front(); - self.curr_region_name = region_name; - break 'search; - } - None => { - if let Some(region_name) = self.region_names.pop_front() - { - debug!( - "skipping region {region_name}, no valid \ - positions for motifs {:?}", - &self.motifs - ) - } else { - debug!( - "skipping {}, no valid positions for motifs \ - {:?}", - &record.name, &self.motifs - ) - } - continue; - } - } - } else { + let Some(search_space) = self.work_queue.pop_front() else { assert!(self.region_names.is_empty()); self.done = true; break 'search; + }; + assert!( + self.region_names.is_empty() + || self.region_names.len() == self.work_queue.len() + 1, + "region names must remain aligned with search spaces" + ); + let region_name = if self.region_names.is_empty() { + None + } else { + self.region_names.pop_front() + }; + let mut active = ActiveReference::new( + search_space, + self.motifs.clone(), + self.combine_strands, + self.scan_chunk_size, + self.stripes_per_batch, + ); + let first_position = active.first_position().expect( + "combined motif partner conflicts must be rejected during \ + construction", + ); + if let Some(first_position) = first_position { + let start_position = first_position + .checked_sub(active.record.start as u64) + .and_then(|position| usize::try_from(position).ok()) + .expect("motif hit must be local to its reference slice"); + self.curr_position = start_position; + self.curr_reference = Some(active); + self.curr_region_name = region_name; + self.observe_retained_hits(); + break 'search; + } + + if let Some(region_name) = region_name { + debug!( + "skipping region {region_name}, no valid positions for \ + motifs {:?}", + &self.motifs + ); + } else { + debug!( + "skipping {}, no valid positions for motifs {:?}", + &active.record.name, &self.motifs + ); } } } pub(super) fn total_length(&self) -> usize { - self.work_queue.iter().map(|(_, s)| s.len()).sum::() - + self.curr_seq.len() + self.work_queue + .iter() + .map(|search_space| { + search_space.owner.end - search_space.owner.start + }) + .sum::() + + self + .curr_reference + .as_ref() + .map(|active| active.owner.end - active.owner.start) + .unwrap_or(0) + } + + fn current_chrom_id(&self) -> u32 { + self.curr_reference + .as_ref() + .expect("active reference must exist") + .record + .tid + } + + fn observe_retained_hits(&mut self) { + #[cfg(test)] + { + let retained_hits = self + .curr_reference + .as_ref() + .map(ActiveReference::live_hits_for_test) + .unwrap_or(0); + self.high_water_retained_hits = + self.high_water_retained_hits.max(retained_hits); + } + } + + #[cfg(test)] + fn high_water_retained_hits_for_test(&self) -> usize { + self.high_water_retained_hits } } @@ -1276,7 +2175,7 @@ impl Iterator for SlidingWindows { std::mem::replace(&mut self.curr_region_name, None); if !finished_windows.is_empty() { let entropy_windows = GenomeWindows::new( - self.curr_contig.tid, + self.current_chrom_id(), finished_windows, finished_region, ); @@ -1303,7 +2202,7 @@ impl Iterator for SlidingWindows { std::mem::replace(&mut windows, Vec::new()); if !finished_windows.is_empty() { let entropy_windows = GenomeWindows::new( - self.curr_contig.tid, + self.current_chrom_id(), finished_windows, None, ); @@ -1318,7 +2217,7 @@ impl Iterator for SlidingWindows { "region names should be empty here also!" ); let entropy_windows = - GenomeWindows::new(self.curr_contig.tid, windows, None); + GenomeWindows::new(self.current_chrom_id(), windows, None); batch.push(entropy_windows) } @@ -1386,8 +2285,13 @@ impl DescriptiveStats { "measurements and n_reads should be the same length" ); let mean_entropy = Self::mean(measurements); - let median_entropy = - percentile_linear_interp(measurements, 0.5f32)?; + let median_entropy = if measurements.len() == 1 { + measurements[0] + } else { + let mut sorted_measurements = measurements.to_vec(); + sorted_measurements.sort_unstable_by(f32::total_cmp); + percentile_linear_interp(&sorted_measurements, 0.5f32)? + }; // safe because of above check let (min_entropy, max_entropy) = match measurements.iter().minmax() { @@ -1635,10 +2539,6 @@ struct BedRegion { } impl BedRegion { - fn length(&self) -> usize { - self.interval.end - self.interval.start - } - fn parser(raw: &str) -> IResult<&str, Self> { let n_parts = raw.split('\t').count(); let (rest, chrom) = crate::parsing_utils::consume_string(raw)?; @@ -1671,7 +2571,684 @@ impl BedRegion { #[cfg(test)] mod entropy_mod_tests { - use crate::entropy::BedRegion; + use crate::entropy::methylation_entropy::EntropySymbol; + use crate::entropy::{ + BedRegion, DescriptiveStats, GenomeWindow, GenomeWindows, + ReferenceSearchSpace, SlidingWindows, + }; + use crate::mod_bam::BaseModCall; + use crate::mod_base_code::{DnaBase, ModCodeRepr}; + use crate::motifs::motif_bed::RegexMotif; + use crate::util::{ReferenceRecord, Strand}; + use itertools::Itertools; + use rayon::prelude::*; + use rayon::ThreadPoolBuilder; + use rustc_hash::FxHashMap; + use std::collections::{BTreeSet, HashSet, VecDeque}; + use std::sync::Arc; + + #[test] + fn singleton_entropy_summary_uses_its_value_as_the_median() { + let stats = DescriptiveStats::new(&[0.25], &[7], 2, 0, &(10..11)) + .expect("a singleton is a valid regional entropy summary"); + + assert_eq!(stats.mean_entropy, 0.25); + assert_eq!(stats.median_entropy, 0.25); + assert_eq!(stats.min_entropy, 0.25); + assert_eq!(stats.max_entropy, 0.25); + assert_eq!(stats.mean_num_reads, 7.0); + assert_eq!(stats.min_num_reads, 7); + assert_eq!(stats.max_num_reads, 7); + assert_eq!(stats.successful_count, 1); + assert_eq!(stats.failed_count, 2); + } + + #[test] + fn multi_window_entropy_summary_keeps_existing_exact_statistics() { + let stats = DescriptiveStats::new( + &[0.25, 0.5, 0.75], + &[2, 4, 6], + 1, + 0, + &(10..20), + ) + .unwrap(); + + assert_eq!(stats.mean_entropy, 0.5); + assert_eq!(stats.median_entropy, 0.5); + assert_eq!(stats.min_entropy, 0.25); + assert_eq!(stats.max_entropy, 0.75); + assert_eq!(stats.mean_num_reads, 4.0); + assert_eq!(stats.min_num_reads, 2); + assert_eq!(stats.max_num_reads, 6); + assert_eq!(stats.successful_count, 3); + assert_eq!(stats.failed_count, 1); + } + + #[test] + fn region_median_is_independent_of_window_encounter_order() { + let observed = DescriptiveStats::new( + &[0.9, 0.1, 0.4, 0.2], + &[9, 1, 4, 2], + 0, + 0, + &(10..20), + ) + .unwrap(); + let sorted = DescriptiveStats::new( + &[0.1, 0.2, 0.4, 0.9], + &[1, 2, 4, 9], + 0, + 0, + &(10..20), + ) + .unwrap(); + + assert_eq!(observed.median_entropy, sorted.median_entropy); + assert_eq!(observed.median_entropy, 0.3); + } + + #[test] + fn odd_region_median_is_independent_of_window_encounter_order() { + let observed = DescriptiveStats::new( + &[0.9, 0.1, 0.4], + &[9, 1, 4], + 0, + 0, + &(10..20), + ) + .unwrap(); + let sorted = DescriptiveStats::new( + &[0.1, 0.4, 0.9], + &[1, 4, 9], + 0, + 0, + &(10..20), + ) + .unwrap(); + + assert_eq!(observed.median_entropy, sorted.median_entropy); + assert_eq!(observed.median_entropy, 0.4); + } + + #[test] + fn regional_medians_are_stable_across_all_small_encounter_orders() { + for (values, expected_median) in + [(vec![0.75, 0.25, 0.5], 0.5), (vec![0.75, 0.0, 0.25, 0.5], 0.375)] + { + for measurements in + values.iter().copied().permutations(values.len()) + { + let original_order = measurements.clone(); + let reads = vec![1; measurements.len()]; + let stats = DescriptiveStats::new( + &measurements, + &reads, + 0, + 0, + &(10..20), + ) + .unwrap(); + + assert_eq!(stats.median_entropy, expected_median); + assert_eq!(measurements, original_order); + } + } + } + + fn combined_window_with_code_count(code_count: usize) -> GenomeWindow { + let read_patterns = (0..code_count) + .map(|idx| { + vec![BaseModCall::Modified(1.0, ModCodeRepr::ChEbi(idx as u32))] + }) + .collect(); + GenomeWindow::CombineStrands { + interval: 0..1, + neg_to_pos_positions: FxHashMap::default(), + read_patterns, + position_valid_coverages: vec![code_count as u32], + } + } + + fn mixed_codes() -> Vec { + vec![ + ModCodeRepr::Code('z'), + ModCodeRepr::ChEbi(900), + ModCodeRepr::Code('a'), + ModCodeRepr::ChEbi(1), + ModCodeRepr::Code('m'), + ModCodeRepr::ChEbi(42), + ModCodeRepr::Code('b'), + ModCodeRepr::ChEbi(7), + ModCodeRepr::Code('q'), + ModCodeRepr::ChEbi(3), + ModCodeRepr::Code('x'), + ModCodeRepr::ChEbi(100), + ModCodeRepr::Code('c'), + ModCodeRepr::ChEbi(2), + ModCodeRepr::Code('d'), + ModCodeRepr::ChEbi(10), + ModCodeRepr::Code('e'), + ] + } + + fn stranded_lookup_window( + pos_codes: &[ModCodeRepr], + neg_codes: &[ModCodeRepr], + ) -> GenomeWindow { + let make_patterns = |codes: &[ModCodeRepr]| { + codes + .iter() + .map(|code| vec![BaseModCall::Modified(1.0, *code)]) + .collect::>() + }; + GenomeWindow::Stranded { + pos_interval: None, + neg_interval: None, + pos_positions: None, + neg_positions: None, + pos_read_patterns: make_patterns(pos_codes), + neg_read_patterns: make_patterns(neg_codes), + pos_position_valid_coverages: Vec::new(), + neg_position_valid_coverages: Vec::new(), + } + } + + fn p_q_w_calls(codes: &[ModCodeRepr]) -> Vec> { + let p = codes + .iter() + .map(|code| BaseModCall::Modified(1.0, *code)) + .collect::>(); + let mut q = p.clone(); + q[0] = BaseModCall::Canonical(1.0); + let mut wildcard = p.clone(); + wildcard[0] = BaseModCall::Filtered; + vec![p, q, wildcard] + } + + fn oracle_window( + codes: &[ModCodeRepr], + pos_order: &[usize; 3], + neg_order: &[usize; 3], + ) -> GenomeWindow { + let patterns = p_q_w_calls(codes); + let reorder = |order: &[usize; 3]| { + order.iter().map(|idx| patterns[*idx].clone()).collect::>() + }; + let mut coverages = vec![3; codes.len()]; + coverages[0] = 2; + let interval_end = codes.len().saturating_sub(1) as u64; + GenomeWindow::Stranded { + pos_interval: Some(0..interval_end), + neg_interval: Some(100..100 + interval_end), + pos_positions: None, + neg_positions: None, + pos_read_patterns: reorder(pos_order), + neg_read_patterns: reorder(neg_order), + pos_position_valid_coverages: coverages.clone(), + neg_position_valid_coverages: coverages, + } + } + + fn entropy_snapshot(window: &GenomeWindow) -> [(f32, usize); 2] { + let entropy = window.into_entropy(7, 2); + let pos = entropy.pos_me_entropy.unwrap().unwrap(); + let neg = entropy.neg_me_entropy.unwrap().unwrap(); + [(pos.me_entropy, pos.num_reads), (neg.me_entropy, neg.num_reads)] + } + + fn motif_sequence(length: usize, start: usize, motif: &str) -> Vec { + let mut sequence = vec!['N'; length]; + for (idx, base) in motif.chars().enumerate() { + sequence[start + idx] = base; + } + sequence + } + + fn sliding_windows_for_test( + sequence: &str, + motifs: Vec, + combine_strands: bool, + num_positions: usize, + ) -> SlidingWindows { + let sequence = sequence.chars().collect::>(); + let sequence_length = sequence.len(); + SlidingWindows::new_for_test( + VecDeque::from([ReferenceSearchSpace { + record: ReferenceRecord::new( + 0, + 0, + sequence_length as u32, + "chr1".to_string(), + ), + sequence: Arc::new(sequence), + owner: 0..sequence_length, + }]), + motifs, + combine_strands, + num_positions, + sequence_length, + 1, + SlidingWindows::START_SEARCH_CHUNK_SIZE, + SlidingWindows::STRIPES_PER_BATCH, + ) + .unwrap() + } + + #[test] + fn one_position_window_is_one_base_half_open() { + let mut window = + GenomeWindow::new_stranded(Some(vec![(DnaBase::C, 9)]), None, 1); + window.inc_coverage(0, &Strand::Positive); + window + .add_pattern(&Strand::Positive, vec![BaseModCall::Canonical(1.0)]); + + let entropy = window.into_entropy(0, 1); + assert_eq!(entropy.pos_me_entropy.unwrap().unwrap().interval, 9..10); + } + + #[test] + fn reverse_anchor_at_final_base_is_one_base_half_open() { + let mut window = + GenomeWindow::new_stranded(None, Some(vec![(DnaBase::C, 3)]), 1); + window.inc_coverage(0, &Strand::Negative); + window + .add_pattern(&Strand::Negative, vec![BaseModCall::Canonical(1.0)]); + + let entropy = window.into_entropy(0, 1); + assert_eq!(entropy.neg_me_entropy.unwrap().unwrap().interval, 3..4); + } + + #[test] + fn entropy_start_position_is_global_after_first_search_chunk() { + let sequence = motif_sequence(20_010, 15_000, "GATC"); + let motif = RegexMotif::parse_string("GATC", 1).unwrap(); + + assert_eq!( + super::SlidingWindows::find_start_position(&sequence, &[motif]), + Some(15_001) + ); + } + + #[test] + fn entropy_start_position_finds_motif_across_search_chunk_seam() { + let sequence = motif_sequence(10_010, 9_998, "GATC"); + let motif = RegexMotif::parse_string("GATC", 1).unwrap(); + + assert_eq!( + super::SlidingWindows::find_start_position(&sequence, &[motif]), + Some(9_999) + ); + } + + #[test] + fn fetch_range_uses_global_min_start_and_max_exclusive_end() { + let first = + GenomeWindow::new_combine_strands(10..21, 0, FxHashMap::default()); + let second = + GenomeWindow::new_combine_strands(15..17, 0, FxHashMap::default()); + let windows = GenomeWindows::new(0, vec![first, second], None); + + assert_eq!(windows.get_range(), 10..21); + } + + #[test] + fn combined_window_advances_from_owned_anchor_not_left_partner() { + let motifs = vec![RegexMotif::parse_string("GATC", 3).unwrap()]; + let mut windows = sliding_windows_for_test("GATC", motifs, true, 1); + + assert!(windows.next_window().is_some()); + assert_eq!(windows.curr_position, 4); + } + + #[test] + fn overlapping_motifs_do_not_duplicate_non_combined_anchor() { + let motifs = vec![ + RegexMotif::parse_string("CG", 0).unwrap(), + RegexMotif::parse_string("CGN", 0).unwrap(), + ]; + let mut windows = sliding_windows_for_test("CGA", motifs, false, 2); + + assert!(windows.next_window().is_none()); + } + + #[test] + fn combined_motifs_reject_two_anchors_for_one_negative_partner() { + let motifs = vec![ + RegexMotif::parse_string("GATC", 1).unwrap(), + RegexMotif::parse_string("GATC", 2).unwrap(), + ]; + let hits = vec![ + super::MotifHit::new(3, Some(10), Strand::Positive, DnaBase::C, 0), + super::MotifHit::new(5, Some(10), Strand::Positive, DnaBase::C, 1), + ]; + + let error = SlidingWindows::sort_and_dedup_motif_hits( + hits, &motifs, true, "chr1", + ) + .unwrap_err() + .to_string(); + assert!(error.contains("negative C partner 10"), "{error}"); + assert!(error.contains("positive anchors 3"), "{error}"); + assert!(error.contains("and 5"), "{error}"); + assert!(error.contains("motif GATC,1"), "{error}"); + assert!(error.contains("motif GATC,2"), "{error}"); + } + + #[test] + fn combined_motifs_deduplicate_identical_anchor_partner_pairs() { + let motifs = vec![ + RegexMotif::parse_string("CG", 0).unwrap(), + RegexMotif::parse_string("CGN", 0).unwrap(), + ]; + let hits = vec![ + super::MotifHit::new(4, Some(5), Strand::Positive, DnaBase::C, 0), + super::MotifHit::new(4, Some(5), Strand::Positive, DnaBase::C, 1), + ]; + + let deduped = SlidingWindows::sort_and_dedup_motif_hits( + hits, &motifs, true, "chr1", + ) + .unwrap(); + assert_eq!(deduped.len(), 1); + } + + #[test] + fn combined_partner_conflict_is_detected_across_scan_stripes() { + let sequence = "CGC".chars().collect::>(); + let motifs = vec![ + RegexMotif::parse_string("CG", 0).unwrap(), + RegexMotif::parse_string("GC", 1).unwrap(), + ]; + + let error = SlidingWindows::scan_motif_hits_with_config_for_test( + &sequence, + &motifs, + 0..sequence.len(), + 0, + "chr1", + true, + 2, + 1, + Some(Strand::Positive), + ) + .unwrap_err() + .to_string(); + + assert!(error.contains("negative C partner 1"), "{error}"); + assert!(error.contains("positive anchors 0"), "{error}"); + assert!(error.contains("and 2"), "{error}"); + } + + #[test] + fn dense_base_scan_is_identical_across_threads_and_chunk_seams() { + let sequence = "ACGT".repeat(257).chars().collect::>(); + let motifs = ["A", "C", "G", "T"] + .into_iter() + .map(|base| RegexMotif::parse_string(base, 0).unwrap()) + .collect::>(); + let mut baseline = None; + + for (threads, chunk_size) in [(1, 7), (4, 7), (4, 13)] { + let pool = + ThreadPoolBuilder::new().num_threads(threads).build().unwrap(); + let hits = pool + .install(|| { + SlidingWindows::scan_motif_hits_with_config_for_test( + &sequence, + &motifs, + 0..sequence.len(), + 0, + "chr1", + false, + chunk_size, + 3, + None, + ) + }) + .unwrap(); + let snapshot = hits + .iter() + .map(|hit| { + ( + hit.pos, + hit.strand, + hit.base, + hit.neg_position, + hit.motif_idx, + ) + }) + .collect::>(); + + assert_eq!(snapshot.len(), sequence.len() * 2); + assert!(snapshot.windows(2).all(|pair| pair[0] <= pair[1])); + if let Some(expected) = baseline.as_ref() { + assert_eq!(&snapshot, expected); + } else { + baseline = Some(snapshot); + } + } + } + + #[test] + fn duplicate_anchor_groups_remain_unique_across_scan_stripes() { + let sequence = "CGACGA".chars().collect::>(); + let motifs = vec![ + RegexMotif::parse_string("CG", 0).unwrap(), + RegexMotif::parse_string("CGN", 0).unwrap(), + ]; + + let hits = SlidingWindows::scan_motif_hits_with_config_for_test( + &sequence, + &motifs, + 0..sequence.len(), + 0, + "chr1", + false, + 1, + 2, + Some(Strand::Positive), + ) + .unwrap(); + + assert_eq!( + hits.iter().map(|hit| hit.pos).collect::>(), + vec![0, 3] + ); + } + + #[test] + fn out_of_window_lookahead_is_not_emitted_and_scan_terminates() { + let make_search_space = |tid, name: &str| { + let sequence = "AAAA".chars().collect::>(); + ReferenceSearchSpace { + record: ReferenceRecord::new( + tid, + 0, + sequence.len() as u32, + name.to_string(), + ), + owner: 0..sequence.len(), + sequence: Arc::new(sequence), + } + }; + let mut windows = SlidingWindows::new_for_test( + VecDeque::from([ + make_search_space(0, "chr1"), + make_search_space(1, "chr2"), + ]), + vec![RegexMotif::parse_string("A", 0).unwrap()], + false, + 2, + 1, + 1, + 2, + 2, + ) + .unwrap(); + + // Every lookahead contains the next anchor at or beyond the exclusive + // one-base window end, so no two-position window is eligible. + assert!(windows.next().is_none()); + assert!(windows.done); + assert!(windows.next().is_none()); + } + + #[test] + fn dense_two_owner_scan_has_fixed_live_hit_high_water() { + const CHUNK_SIZE: usize = 31; + const STRIPES_PER_BATCH: usize = 3; + const NUM_POSITIONS: usize = 4; + let make_search_space = |tid, name: &str, repeats| { + let sequence = "ACGT".repeat(repeats).chars().collect::>(); + let length = sequence.len(); + ReferenceSearchSpace { + record: ReferenceRecord::new( + tid, + 0, + length as u32, + name.to_string(), + ), + sequence: Arc::new(sequence), + owner: 0..length, + } + }; + let work_queue = VecDeque::from([ + make_search_space(0, "chr1", 513), + make_search_space(1, "chr2", 777), + ]); + let motifs = ["A", "C", "G", "T"] + .into_iter() + .map(|base| RegexMotif::parse_string(base, 0).unwrap()) + .collect::>(); + let mut windows = SlidingWindows::new_for_test( + work_queue, + motifs, + false, + NUM_POSITIONS, + 17, + 1, + CHUNK_SIZE, + STRIPES_PER_BATCH, + ) + .unwrap(); + + while windows.next().is_some() {} + + // Four single-base motifs produce two raw strand hits per base before + // filtering. The conservative metric counts raw and deduplicated + // stripe storage separately for both strand streams, plus both N-hit + // lookaheads. + let fixed_bound = + 6 * CHUNK_SIZE * STRIPES_PER_BATCH + 2 * NUM_POSITIONS; + let raw_or_dedup_max = + 4 * CHUNK_SIZE * STRIPES_PER_BATCH + 2 * NUM_POSITIONS; + let high_water = windows.high_water_retained_hits_for_test(); + assert!(high_water > raw_or_dedup_max); + assert!(high_water <= fixed_bound); + assert!(high_water < 513 * 4); + } + + #[test] + fn nine_distinct_modification_codes_have_atomic_states() { + let window = combined_window_with_code_count(9); + let lookup = window.get_mod_code_lookup(); + assert_eq!(lookup.len(), 9); + for id in 1..=9 { + assert_eq!( + lookup.get(&ModCodeRepr::ChEbi(id - 1)), + Some(&EntropySymbol::called(id)) + ); + } + } + + #[test] + fn ten_distinct_modification_codes_have_atomic_states() { + let window = combined_window_with_code_count(10); + assert_eq!(window.get_mod_code_lookup().len(), 10); + } + + #[test] + fn mixed_code_strand_union_preserves_current_btree_ranking() { + let codes = mixed_codes(); + let first = stranded_lookup_window(&codes[..8], &codes[8..]); + let mut reversed_pos = codes[8..].to_vec(); + reversed_pos.reverse(); + let mut reversed_neg = codes[..8].to_vec(); + reversed_neg.reverse(); + let permuted = stranded_lookup_window(&reversed_pos, &reversed_neg); + let first_lookup = first.get_mod_code_lookup(); + let permuted_lookup = permuted.get_mod_code_lookup(); + assert_eq!(first_lookup.len(), 17); + assert_eq!(first_lookup, permuted_lookup); + assert_eq!( + first_lookup.keys().copied().collect::>(), + codes.iter().copied().collect::>() + ); + let distinct_symbols = + first_lookup.values().copied().collect::>(); + let expected_symbols = + (1..=17).map(EntropySymbol::called).collect::>(); + assert_eq!(distinct_symbols, expected_symbols); + assert!(!distinct_symbols.contains(&EntropySymbol::CANONICAL)); + assert!(!distinct_symbols.contains(&EntropySymbol::FILTERED)); + + let mut code_symbols = codes + .iter() + .filter_map(|code| match code { + ModCodeRepr::Code(c) => Some((*c, first_lookup[code])), + ModCodeRepr::ChEbi(_) => None, + }) + .collect::>(); + code_symbols.sort_by_key(|(code, _)| *code); + assert!(code_symbols.windows(2).all(|pair| pair[0].1 < pair[1].1)); + let mut chebi_symbols = codes + .iter() + .filter_map(|code| match code { + ModCodeRepr::ChEbi(id) => Some((*id, first_lookup[code])), + ModCodeRepr::Code(_) => None, + }) + .collect::>(); + chebi_symbols.sort_by_key(|(code, _)| *code); + assert!(chebi_symbols.windows(2).all(|pair| pair[0].1 < pair[1].1)); + + let expected_entropy = [(1.0 / 17.0, 3), (1.0 / 17.0, 3)]; + assert_eq!( + entropy_snapshot(&oracle_window(&codes, &[0, 1, 2], &[2, 0, 1])), + expected_entropy + ); + let mut reversed_codes = codes.clone(); + reversed_codes.reverse(); + assert_eq!( + entropy_snapshot(&oracle_window( + &reversed_codes, + &[1, 2, 0], + &[0, 2, 1], + )), + expected_entropy + ); + } + + #[test] + fn seventeen_codes_preserve_oracle_across_threads_and_permutations() { + let codes = (1..=17).map(ModCodeRepr::ChEbi).collect::>(); + let mut reversed_codes = codes.clone(); + reversed_codes.reverse(); + let expected = [(1.0 / 17.0, 3), (1.0 / 17.0, 3)]; + + for threads in [1, 2, 4] { + let windows = vec![ + oracle_window(&codes, &[0, 1, 2], &[2, 0, 1]), + oracle_window(&reversed_codes, &[1, 2, 0], &[0, 2, 1]), + ]; + let pool = + ThreadPoolBuilder::new().num_threads(threads).build().unwrap(); + let observed = pool.install(|| { + windows.par_iter().map(entropy_snapshot).collect::>() + }); + assert_eq!(observed, vec![expected, expected]); + } + } #[test] fn test_bed_region_parsing() { diff --git a/modkit-core/src/entropy/subcommand.rs b/modkit-core/src/entropy/subcommand.rs index cbfe1257..c23abe16 100644 --- a/modkit-core/src/entropy/subcommand.rs +++ b/modkit-core/src/entropy/subcommand.rs @@ -203,33 +203,6 @@ impl MethylationEntropy { })?; } - let mut writer: Box = - match (self.out_bed.as_ref(), self.regions_fp.is_some()) { - (Some(out_fp), false) => Box::new( - WindowsWriter::new_file(out_fp, self.header, self.verbose) - .context("failed to make writer to file")?, - ), - (Some(out_dir), true) => Box::new( - RegionsWriter::new( - out_dir, - self.prefix.as_ref(), - self.header, - self.verbose, - ) - .context( - "failed to make regions writer, output must be a \ - directory", - )?, - ), - (None, false) => Box::new( - WindowsWriter::new_stdout(self.header, self.verbose) - .context("failed to make writer to stdout")?, - ), - (None, true) => { - bail!("must provide output directory with regions") - } - }; - let pool = rayon::ThreadPoolBuilder::new() .num_threads(self.threads) .build()?; @@ -335,6 +308,36 @@ impl MethylationEntropy { } })?; + // Motif ownership and combined-strand partner validation happen while + // constructing the sliding windows. Do not create or truncate output + // until that scientific preflight has succeeded. + let mut writer: Box = + match (self.out_bed.as_ref(), self.regions_fp.is_some()) { + (Some(out_fp), false) => Box::new( + WindowsWriter::new_file(out_fp, self.header, self.verbose) + .context("failed to make writer to file")?, + ), + (Some(out_dir), true) => Box::new( + RegionsWriter::new( + out_dir, + self.prefix.as_ref(), + self.header, + self.verbose, + ) + .context( + "failed to make regions writer, output must be a \ + directory", + )?, + ), + (None, false) => Box::new( + WindowsWriter::new_stdout(self.header, self.verbose) + .context("failed to make writer to stdout")?, + ), + (None, true) => { + bail!("must provide output directory with regions") + } + }; + let threshold_caller = self.get_threshold_caller(&pool).map(|c| Arc::new(c))?; diff --git a/modkit-core/src/reads_sampler/sampling_schedule.rs b/modkit-core/src/reads_sampler/sampling_schedule.rs index 21a99d29..f9d856f4 100644 --- a/modkit-core/src/reads_sampler/sampling_schedule.rs +++ b/modkit-core/src/reads_sampler/sampling_schedule.rs @@ -925,6 +925,17 @@ impl ReferenceSequencesLookup { .map(|id| *self.id_to_tid.get(&id).unwrap()) } + pub(crate) fn sequence_length_by_name( + &self, + name: &str, + ) -> anyhow::Result { + let id = self + .reference_sequence_names + .get_index_of(name) + .ok_or_else(|| anyhow!("seq {name} not in used references"))?; + Ok(self.reference_sequences.get(&id).unwrap().len()) + } + pub(crate) fn get_subsequence_by_name( &self, name: &str, diff --git a/modkit/tests/test_entropy_geometry.rs b/modkit/tests/test_entropy_geometry.rs new file mode 100644 index 00000000..f5ec6faa --- /dev/null +++ b/modkit/tests/test_entropy_geometry.rs @@ -0,0 +1,309 @@ +use rust_htslib::bam::{ + self, + header::HeaderRecord, + record::{Aux, Cigar, CigarString}, +}; +use std::fs; +use std::path::{Path, PathBuf}; +use std::process::{Command, Output}; + +fn write_reference(root: &Path, sequence: &str) -> PathBuf { + let reference = root.join("reference.fa"); + fs::write(&reference, format!(">chr1\n{sequence}\n")).unwrap(); + fs::write( + root.join("reference.fa.fai"), + format!( + "chr1\t{}\t6\t{}\t{}\n", + sequence.len(), + sequence.len(), + sequence.len() + 1 + ), + ) + .unwrap(); + reference +} + +fn write_bam( + root: &Path, + sequence: &str, + mm_tag: &str, + ml_count: usize, + strands: &[bool], +) -> PathBuf { + let bam_path = root.join("reads.bam"); + let mut header = bam::Header::new(); + let mut sq = HeaderRecord::new(b"SQ"); + sq.push_tag(b"SN", "chr1").push_tag(b"LN", sequence.len()); + header.push_record(&sq); + + let cigar = CigarString(vec![Cigar::Match(sequence.len() as u32)]); + let mut writer = + bam::Writer::from_path(&bam_path, &header, bam::Format::Bam).unwrap(); + for (idx, reverse) in strands.iter().copied().enumerate() { + let mut record = bam::Record::new(); + record.set( + format!("read-{idx}").as_bytes(), + Some(&cigar), + sequence.as_bytes(), + &vec![30; sequence.len()], + ); + record.set_tid(0); + record.set_pos(0); + record.set_mapq(60); + if reverse { + record.set_flags(16); + } + let ml = vec![255u8; ml_count]; + record.push_aux(b"MM", Aux::String(mm_tag)).unwrap(); + record.push_aux(b"ML", Aux::ArrayU8((&ml[..]).into())).unwrap(); + record.push_aux(b"MN", Aux::U32(sequence.len() as u32)).unwrap(); + record.push_aux(b"NM", Aux::U32(0)).unwrap(); + writer.write(&record).unwrap(); + } + drop(writer); + bam::index::build(&bam_path, None, bam::index::Type::Bai, 1).unwrap(); + bam_path +} + +fn write_two_contig_conflict_fixture(root: &Path) -> (PathBuf, PathBuf) { + let reference = root.join("two-contigs.fa"); + fs::write(&reference, ">chr1\nCG\n>chr2\nCGCG\n").unwrap(); + fs::write( + root.join("two-contigs.fa.fai"), + "chr1\t2\t6\t2\t3\nchr2\t4\t15\t4\t5\n", + ) + .unwrap(); + + let bam_path = root.join("two-contigs.bam"); + let mut header = bam::Header::new(); + for (name, length) in [("chr1", 2), ("chr2", 4)] { + let mut sq = HeaderRecord::new(b"SQ"); + sq.push_tag(b"SN", name).push_tag(b"LN", length); + header.push_record(&sq); + } + let mut writer = + bam::Writer::from_path(&bam_path, &header, bam::Format::Bam).unwrap(); + for (tid, sequence, mm_tag, ml_count) in + [(0, "CG", "C+m?,0;", 1), (1, "CGCG", "C+m?,0,0;", 2)] + { + let cigar = CigarString(vec![Cigar::Match(sequence.len() as u32)]); + let mut record = bam::Record::new(); + record.set( + format!("read-{tid}").as_bytes(), + Some(&cigar), + sequence.as_bytes(), + &vec![30; sequence.len()], + ); + record.set_tid(tid); + record.set_pos(0); + record.set_mapq(60); + let ml = vec![255u8; ml_count]; + record.push_aux(b"MM", Aux::String(mm_tag)).unwrap(); + record.push_aux(b"ML", Aux::ArrayU8((&ml[..]).into())).unwrap(); + record.push_aux(b"MN", Aux::U32(sequence.len() as u32)).unwrap(); + record.push_aux(b"NM", Aux::U32(0)).unwrap(); + writer.write(&record).unwrap(); + } + drop(writer); + bam::index::build(&bam_path, None, bam::index::Type::Bai, 1).unwrap(); + (reference, bam_path) +} + +fn run_entropy( + bam: &Path, + reference: &Path, + output: &Path, + window_size: usize, + regions: Option<&Path>, +) -> Output { + let mut command = Command::new(env!("CARGO_BIN_EXE_modkit")); + command.args([ + "entropy", + "--in-bam", + bam.to_str().unwrap(), + "--out-bed", + output.to_str().unwrap(), + "--ref", + reference.to_str().unwrap(), + "--motif", + "CG", + "0", + "--num-positions", + "1", + "--window-size", + &window_size.to_string(), + "--min-coverage", + "1", + "--max-filtered-positions", + "0", + "--filter-threshold", + "0", + "--threads", + "1", + "--io-threads", + "1", + "--suppress-progress", + "--force", + ]); + if let Some(regions) = regions { + command.args([ + "--regions", + regions.to_str().unwrap(), + "--prefix", + "anchor", + ]); + } + command.output().unwrap() +} + +#[test] +fn cg_rows_are_invariant_to_search_window_size() { + let temp_dir = tempfile::tempdir().unwrap(); + let reference = write_reference(temp_dir.path(), "CGTACG"); + let bam = + write_bam(temp_dir.path(), "CGTACG", "C+m?,0,0;", 2, &[false, true]); + let expected = concat!( + "chr1\t0\t1\t0\t+\t1\n", + "chr1\t1\t2\t0\t-\t1\n", + "chr1\t4\t5\t0\t+\t1\n", + "chr1\t5\t6\t0\t-\t1\n", + ); + + for window_size in [1, 2, 3, 6, 100] { + let output = temp_dir.path().join(format!("entropy-{window_size}.bed")); + let result = run_entropy(&bam, &reference, &output, window_size, None); + assert!( + result.status.success(), + "window_size={window_size}:\n{}", + String::from_utf8_lossy(&result.stderr) + ); + assert_eq!( + fs::read(output).unwrap(), + expected.as_bytes(), + "window_size={window_size}" + ); + } +} + +#[test] +fn region_owning_only_the_anchor_uses_reference_context() { + let temp_dir = tempfile::tempdir().unwrap(); + let reference = write_reference(temp_dir.path(), "AACGAA"); + let bam = write_bam(temp_dir.path(), "AACGAA", "C+m?,0;", 1, &[false]); + let regions = temp_dir.path().join("regions.bed"); + fs::write(®ions, "chr1\t2\t3\tanchor-only\n").unwrap(); + let output_dir = temp_dir.path().join("entropy-regions"); + + let result = run_entropy(&bam, &reference, &output_dir, 1, Some(®ions)); + assert!( + result.status.success(), + "{}", + String::from_utf8_lossy(&result.stderr) + ); + assert_eq!( + fs::read(output_dir.join("anchor_windows.bedgraph")).unwrap(), + b"chr1\t2\t3\t0\t+\t1\n" + ); + // Regional singleton statistics are covered independently by issue #682; + // this fixture isolates motif context and proves that the exact region + // owner emits its biological anchor. +} + +#[test] +fn conflicting_combined_motif_partners_fail_before_output_creation() { + let temp_dir = tempfile::tempdir().unwrap(); + let reference = write_reference(temp_dir.path(), "CGCG"); + let bam = write_bam(temp_dir.path(), "CGCG", "C+m?,0,0;", 2, &[false]); + let output = temp_dir.path().join("conflict.bed"); + let result = Command::new(env!("CARGO_BIN_EXE_modkit")) + .args([ + "entropy", + "--in-bam", + bam.to_str().unwrap(), + "--out-bed", + output.to_str().unwrap(), + "--ref", + reference.to_str().unwrap(), + "--motif", + "CG", + "0", + "--motif", + "CGCG", + "0", + "--combine-strands", + "--num-positions", + "1", + "--window-size", + "4", + "--min-coverage", + "1", + "--filter-threshold", + "0", + "--threads", + "1", + "--io-threads", + "1", + "--suppress-progress", + ]) + .output() + .unwrap(); + + assert!(!result.status.success()); + let stderr = String::from_utf8_lossy(&result.stderr); + assert!( + stderr.contains("conflicting combined-strand motif partners"), + "{stderr}" + ); + assert!(stderr.contains("chr1:0"), "{stderr}"); + assert!(!output.exists()); +} + +#[test] +fn conflict_on_later_contig_fails_before_output_creation() { + let temp_dir = tempfile::tempdir().unwrap(); + let (reference, bam) = write_two_contig_conflict_fixture(temp_dir.path()); + let output = temp_dir.path().join("later-conflict.bed"); + fs::write(&output, b"sentinel\n").unwrap(); + let result = Command::new(env!("CARGO_BIN_EXE_modkit")) + .args([ + "entropy", + "--in-bam", + bam.to_str().unwrap(), + "--out-bed", + output.to_str().unwrap(), + "--ref", + reference.to_str().unwrap(), + "--motif", + "CG", + "0", + "--motif", + "CGCG", + "0", + "--combine-strands", + "--num-positions", + "1", + "--window-size", + "4", + "--min-coverage", + "1", + "--filter-threshold", + "0", + "--threads", + "1", + "--io-threads", + "1", + "--suppress-progress", + "--force", + ]) + .output() + .unwrap(); + + assert!(!result.status.success()); + let stderr = String::from_utf8_lossy(&result.stderr); + assert!( + stderr.contains("conflicting combined-strand motif partners"), + "{stderr}" + ); + assert!(stderr.contains("chr2:0"), "{stderr}"); + assert_eq!(fs::read(output).unwrap(), b"sentinel\n"); +} diff --git a/modkit/tests/test_entropy_region_statistics.rs b/modkit/tests/test_entropy_region_statistics.rs new file mode 100644 index 00000000..e0675076 --- /dev/null +++ b/modkit/tests/test_entropy_region_statistics.rs @@ -0,0 +1,102 @@ +use rust_htslib::bam::{ + self, + header::HeaderRecord, + record::{Aux, Cigar, CigarString}, +}; +use std::fs; +use std::path::{Path, PathBuf}; +use std::process::Command; + +fn write_reference(root: &Path) -> PathBuf { + let reference = root.join("reference.fa"); + fs::write(&reference, ">chr1\nAACGAA\n").unwrap(); + fs::write(root.join("reference.fa.fai"), "chr1\t6\t6\t6\t7\n").unwrap(); + reference +} + +fn write_bam(root: &Path) -> PathBuf { + let bam_path = root.join("reads.bam"); + let mut header = bam::Header::new(); + let mut sq = HeaderRecord::new(b"SQ"); + sq.push_tag(b"SN", "chr1").push_tag(b"LN", 6); + header.push_record(&sq); + + let cigar = CigarString(vec![Cigar::Match(6)]); + let mut record = bam::Record::new(); + record.set(b"read-0", Some(&cigar), b"AACGAA", &[30; 6]); + record.set_tid(0); + record.set_pos(0); + record.set_mapq(60); + record.push_aux(b"MM", Aux::String("C+m?,0;")).unwrap(); + record.push_aux(b"ML", Aux::ArrayU8((&[255][..]).into())).unwrap(); + record.push_aux(b"MN", Aux::U32(6)).unwrap(); + record.push_aux(b"NM", Aux::U32(0)).unwrap(); + + let mut writer = + bam::Writer::from_path(&bam_path, &header, bam::Format::Bam).unwrap(); + writer.write(&record).unwrap(); + drop(writer); + bam::index::build(&bam_path, None, bam::index::Type::Bai, 1).unwrap(); + bam_path +} + +#[test] +fn singleton_region_emits_exact_summary_statistics() { + let temp_dir = tempfile::tempdir().unwrap(); + let reference = write_reference(temp_dir.path()); + let bam = write_bam(temp_dir.path()); + let regions = temp_dir.path().join("regions.bed"); + fs::write(®ions, "chr1\t2\t3\tsingleton\n").unwrap(); + let output_dir = temp_dir.path().join("output"); + + let result = Command::new(env!("CARGO_BIN_EXE_modkit")) + .args([ + "entropy", + "--in-bam", + bam.to_str().unwrap(), + "--out-bed", + output_dir.to_str().unwrap(), + "--ref", + reference.to_str().unwrap(), + "--base", + "C", + "--num-positions", + "1", + "--window-size", + "1", + "--min-coverage", + "1", + "--max-filtered-positions", + "0", + "--filter-threshold", + "0", + "--threads", + "1", + "--io-threads", + "1", + "--regions", + regions.to_str().unwrap(), + "--prefix", + "singleton", + "--suppress-progress", + ]) + .output() + .unwrap(); + + assert!( + result.status.success(), + "{}", + String::from_utf8_lossy(&result.stderr) + ); + let windows = + fs::read_to_string(output_dir.join("singleton_windows.bedgraph")) + .unwrap(); + let window_fields = windows.trim_end().split('\t').collect::>(); + assert_eq!(window_fields[0], "chr1"); + assert_eq!(window_fields[1], "2"); + assert_eq!(&window_fields[3..], &["0", "+", "1"]); + assert_eq!( + fs::read(output_dir.join("singleton_regions.bed")).unwrap(), + b"chr1\t2\t3\tsingleton\t0\t+\t0\t0\t0\t1\t1\t1\t1\t0\n" + ); +} diff --git a/modkit/tests/test_entropy_state_cardinality.rs b/modkit/tests/test_entropy_state_cardinality.rs new file mode 100644 index 00000000..23fede5c --- /dev/null +++ b/modkit/tests/test_entropy_state_cardinality.rs @@ -0,0 +1,131 @@ +use rust_htslib::bam::{ + self, + header::HeaderRecord, + record::{Aux, Cigar, CigarString}, +}; +use std::fs; +use std::path::{Path, PathBuf}; +use std::process::Command; + +fn write_reference(temp_dir: &Path) -> PathBuf { + let reference = temp_dir.join("reference.fa"); + fs::write(&reference, ">chr1\nC\n").unwrap(); + let fai = PathBuf::from(format!("{}.fai", reference.display())); + fs::write(fai, "chr1\t1\t6\t1\t2\n").unwrap(); + reference +} + +fn write_code_bam(temp_dir: &Path, name: &str, codes: &[char]) -> PathBuf { + let bam_path = temp_dir.join(format!("{name}.bam")); + let mut header = bam::Header::new(); + let mut sq = HeaderRecord::new(b"SQ"); + sq.push_tag(b"SN", "chr1").push_tag(b"LN", 1); + header.push_record(&sq); + + let cigar = CigarString(vec![Cigar::Match(1)]); + let mut writer = + bam::Writer::from_path(&bam_path, &header, bam::Format::Bam).unwrap(); + for (idx, code) in codes.iter().enumerate() { + let mut record = bam::Record::new(); + record.set( + format!("read-{idx:02}").as_bytes(), + Some(&cigar), + b"C", + &[30], + ); + record.set_tid(0); + record.set_pos(0); + record.set_mapq(60); + record.set_flags(0); + let mm = format!("C+{code}?,0;"); + record.push_aux(b"MM", Aux::String(&mm)).unwrap(); + record.push_aux(b"ML", Aux::ArrayU8((&[255][..]).into())).unwrap(); + record.push_aux(b"MN", Aux::U32(1)).unwrap(); + writer.write(&record).unwrap(); + } + drop(writer); + bam::index::build(&bam_path, None, bam::index::Type::Bai, 1).unwrap(); + bam_path +} + +fn run_entropy(input: &Path, reference: &Path, output: &Path, threads: usize) { + let result = Command::new(env!("CARGO_BIN_EXE_modkit")) + .args([ + "entropy", + "--in-bam", + input.to_str().unwrap(), + "--out-bed", + output.to_str().unwrap(), + "--ref", + reference.to_str().unwrap(), + "--base", + "C", + "--num-positions", + "1", + "--window-size", + "1", + "--min-coverage", + "1", + "--filter-threshold", + "0", + "--threads", + &threads.to_string(), + "--io-threads", + "1", + "--suppress-progress", + ]) + .output() + .unwrap(); + assert!( + result.status.success(), + "entropy failed:\n{}", + String::from_utf8_lossy(&result.stderr) + ); +} + +#[test] +fn code_cardinality_cli_is_stable_across_encounter_order_and_threads() { + let temp_dir = tempfile::tempdir().unwrap(); + let reference = write_reference(temp_dir.path()); + + for code_count in [9, 10, 12, 17] { + let codes = ('a'..='z').take(code_count).collect::>(); + let forward = write_code_bam( + temp_dir.path(), + &format!("forward-{code_count}"), + &codes, + ); + let mut reversed_codes = codes; + reversed_codes.reverse(); + let reversed = write_code_bam( + temp_dir.path(), + &format!("reversed-{code_count}"), + &reversed_codes, + ); + + let mut expected = None; + for (order, input) in [("forward", &forward), ("reversed", &reversed)] { + for threads in [1, 4] { + let output = temp_dir + .path() + .join(format!("{code_count}-{order}-{threads}.bed")); + run_entropy(input, &reference, &output, threads); + let observed = fs::read(&output).unwrap(); + if let Some(expected) = expected.as_ref() { + assert_eq!(&observed, expected); + } else { + expected = Some(observed); + } + } + } + + let output = String::from_utf8(expected.unwrap()).unwrap(); + let rows = output.lines().collect::>(); + assert_eq!(rows.len(), 1); + let fields = rows[0].split('\t').collect::>(); + assert_eq!(fields.len(), 6); + assert_eq!(fields[5], code_count.to_string()); + let entropy = fields[3].parse::().unwrap(); + assert!((entropy - (code_count as f32).log2()).abs() < 0.000_01); + } +}