Skip to main content

fibertools_rs/subcommands/
pileup.rs

1/// This module is used to extract the fire calls as well as nucs and msps from a bam file
2/// for every position in the bam file and output the results to a bed file.
3/// all calculations are done in total as well as for haplotype 1 and haplotype 2.
4use crate::cli::PileupOptions;
5use crate::fiber::FiberseqData;
6use crate::utils::bamannotations;
7use crate::utils::bio_io;
8use crate::*;
9use anyhow::{anyhow, Ok};
10use ordered_float::NotNan;
11use rust_htslib::bam::ext::BamRecordExtensions;
12use rust_htslib::bam::{FetchDefinition, IndexedReader};
13use std::collections::HashMap;
14use std::io::BufRead;
15
16const MIN_FIRE_COVERAGE: i32 = 4;
17const MIN_FIRE_QUAL: u8 = 229; // floor(255*0.9)
18static WINDOW_SIZE: usize = 1_000_000;
19
20/// Options for FireTrack that don't require the full PileupOptions
21/// This allows FireTrack to be used independently
22#[derive(Debug, Clone, Default)]
23pub struct FireTrackOptions {
24    pub no_nuc: bool,
25    pub no_msp: bool,
26    pub m6a: bool,
27    pub cpg: bool,
28    pub fiber_coverage: bool,
29    pub shuffle: bool,             // Track if shuffling is enabled
30    pub random_shuffle: bool, // If true, generate random positions instead of using ShuffledFibers
31    pub shuffle_seed: Option<u64>, // Optional seed for reproducible random shuffling
32    pub rolling_max: Option<usize>,
33    pub track_fire_elements: bool, // If true, store individual FIRE element positions per base
34}
35
36impl From<&PileupOptions> for FireTrackOptions {
37    fn from(opts: &PileupOptions) -> Self {
38        Self {
39            no_nuc: opts.no_nuc,
40            no_msp: opts.no_msp,
41            m6a: opts.m6a,
42            cpg: opts.cpg,
43            fiber_coverage: opts.effective_fiber_coverage(),
44            shuffle: opts.shuffle.is_some(),
45            random_shuffle: false, // PileupOptions doesn't have this yet
46            shuffle_seed: None,
47            rolling_max: opts.rolling_max,
48            track_fire_elements: false, // Default to false for pileup command
49        }
50    }
51}
52
53#[derive(Debug)]
54pub struct FireRow<'a> {
55    pub coverage: &'a i32,
56    pub fire_coverage: &'a i32,
57    pub score: &'a f32,
58    pub nuc_coverage: &'a i32,
59    pub msp_coverage: &'a i32,
60    pub cpg_coverage: &'a i32,
61    pub m6a_coverage: &'a i32,
62    fire_track_opts: &'a FireTrackOptions,
63}
64
65impl PartialEq for FireRow<'_> {
66    fn eq(&self, other: &Self) -> bool {
67        let m6a = if self.fire_track_opts.m6a {
68            self.m6a_coverage == other.m6a_coverage
69        } else {
70            true
71        };
72        let cpg = if self.fire_track_opts.cpg {
73            self.cpg_coverage == other.cpg_coverage
74        } else {
75            true
76        };
77
78        self.coverage == other.coverage
79            && self.fire_coverage == other.fire_coverage
80            && self.score == other.score
81            && self.nuc_coverage == other.nuc_coverage
82            && self.msp_coverage == other.msp_coverage
83            && cpg
84            && m6a
85    }
86}
87
88impl std::fmt::Display for FireRow<'_> {
89    fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
90        let mut rtn = format!(
91            "\t{}\t{}\t{}",
92            self.coverage, self.fire_coverage, self.score
93        );
94        if !self.fire_track_opts.no_nuc {
95            rtn += &format!("\t{}", self.nuc_coverage);
96        }
97        if !self.fire_track_opts.no_msp {
98            rtn += &format!("\t{}", self.msp_coverage);
99        }
100        if self.fire_track_opts.m6a {
101            rtn += &format!("\t{}", self.m6a_coverage);
102        }
103        if self.fire_track_opts.cpg {
104            rtn += &format!("\t{}", self.cpg_coverage);
105        }
106        write!(f, "{rtn}")
107    }
108}
109
110#[derive(Debug)]
111pub struct ShuffledFibers {
112    pub shuffled_fiber_starts: HashMap<(String, String, i64), i64>,
113}
114
115impl ShuffledFibers {
116    pub fn new(file_path: &str) -> Result<Self> {
117        let buffer = bio_io::buffer_from(file_path)?;
118        let mut shuffled_fiber_starts = HashMap::new();
119        for line in buffer.lines() {
120            let line = line?;
121            if line.starts_with('#') {
122                continue;
123            }
124            let mut parts = line.split('\t');
125            // error if there are not at least 4 parts
126            let chrom = parts.next().ok_or(anyhow!("missing chrom"))?;
127            let start = parts
128                .next()
129                .ok_or(anyhow!("missing fiber start"))?
130                .parse::<i64>()?;
131            let _end = parts
132                .next()
133                .ok_or(anyhow!("missing fiber end"))?
134                .parse::<i64>()?;
135            let fiber_name = parts.next().ok_or(anyhow!("missing fiber name"))?;
136            let original_start = parts
137                .next()
138                .ok_or(anyhow!("missing original start"))?
139                .parse::<i64>()?;
140            shuffled_fiber_starts.insert(
141                (chrom.to_string(), fiber_name.to_string(), original_start),
142                start,
143            );
144        }
145        log::info!("Read {} shuffled fibers", shuffled_fiber_starts.len());
146        Ok(Self {
147            shuffled_fiber_starts,
148        })
149    }
150
151    pub fn get_shuffled_start(&self, fiber: &FiberseqData) -> Option<i64> {
152        let target_name = fiber.target_name.clone();
153        let fiber_name = fiber.get_qname();
154        self.shuffled_fiber_starts
155            .get(&(target_name, fiber_name, fiber.record.reference_start()))
156            .copied()
157    }
158
159    pub fn has_fiber(&self, fiber: &FiberseqData) -> bool {
160        let target_name = fiber.target_name.clone();
161        let fiber_name = fiber.get_qname();
162        self.shuffled_fiber_starts.contains_key(&(
163            target_name,
164            fiber_name,
165            fiber.record.reference_start(),
166        ))
167    }
168
169    /// to get the shuffle offset which we ADD(+) to features of the fiber
170    /// to create a shuffled version of the fiber
171    pub fn get_shuffle_offset(&self, fiber: &FiberseqData) -> Option<i64> {
172        let start = fiber.record.reference_start();
173        let shuffled_start = self.get_shuffled_start(fiber)?;
174        Some(shuffled_start - start)
175    }
176}
177
178/// Generate a random shuffle offset for a fiber using uniform distribution
179/// Uses deterministic PRNG seeded from fiber name + seed for reproducibility
180fn generate_random_shuffle_offset(
181    fiber: &FiberseqData,
182    chrom_len: usize,
183    seed: Option<u64>,
184) -> Option<i64> {
185    use rand::rngs::StdRng;
186    use rand::{Rng, SeedableRng};
187    use std::collections::hash_map::DefaultHasher;
188    use std::hash::{Hash, Hasher};
189
190    let fiber_len = fiber.record.reference_end() - fiber.record.reference_start();
191    let original_start = fiber.record.reference_start();
192
193    // Check if fiber can fit in chromosome
194    let max_start = (chrom_len as i64 - fiber_len).max(0);
195    if max_start <= 0 {
196        return Some(0); // Fiber too long, keep at position 0
197    }
198
199    // Create deterministic seed from fiber name + optional seed
200    let mut hasher = DefaultHasher::new();
201    fiber.get_qname().hash(&mut hasher);
202    if let Some(s) = seed {
203        s.hash(&mut hasher);
204    }
205    let fiber_seed = hasher.finish();
206
207    // Use StdRng for uniform distribution
208    let mut rng = StdRng::seed_from_u64(fiber_seed);
209    let shuffled_start = rng.gen_range(0..=max_start);
210
211    // Return offset (shuffled_start - original_start)
212    Some(shuffled_start - original_start)
213}
214
215/// Represents a single FIRE element (MSP) with its genomic coordinates and unique ID
216#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
217pub struct FireElement {
218    pub start: i64,
219    pub end: i64,
220    pub id: usize, // Unique ID for this FIRE element (for tracking in merging)
221}
222
223#[derive(Debug)]
224pub struct FireTrack<'a> {
225    pub chrom: String,
226    pub chrom_start: usize,
227    pub chrom_end: usize,
228    pub track_len: usize,
229    pub raw_scores: Vec<f32>,
230    pub scores: Vec<f32>,
231    pub coverage: Vec<i32>,
232    pub fire_coverage: Vec<i32>,
233    pub msp_coverage: Vec<i32>,
234    pub nuc_coverage: Vec<i32>,
235    pub cpg_coverage: Vec<i32>,
236    pub m6a_coverage: Vec<i32>,
237    fire_track_opts: FireTrackOptions, // Now owned, not borrowed
238    shuffled_fibers: &'a Option<ShuffledFibers>,
239    cur_offset: i64,
240    // Store fiber information for later shuffle generation
241    // Key: (fiber_name, original_start), Value: fiber_length
242    fibers_seen: HashMap<(String, i64), i64>,
243    // Optional: Store individual FIRE elements per position
244    // fire_elements[position] = Vec of FireElements overlapping that position
245    pub fire_elements: Option<Vec<Vec<FireElement>>>,
246    // Counter for assigning unique IDs to FIRE elements
247    next_fire_id: usize,
248}
249
250impl<'a> FireTrack<'a> {
251    pub fn new(
252        chrom: String,
253        chrom_start: usize,
254        chrom_end: usize,
255        fire_track_opts: FireTrackOptions, // Take ownership
256        shuffled_fibers: &'a Option<ShuffledFibers>,
257    ) -> Self {
258        let track_len = chrom_end - chrom_start + 1;
259        let raw_scores = vec![-1.0; track_len];
260        let scores = vec![-1.0; track_len];
261
262        // Initialize fire_elements only if tracking is enabled
263        let fire_elements = if fire_track_opts.track_fire_elements {
264            Some(vec![Vec::new(); track_len])
265        } else {
266            None
267        };
268
269        Self {
270            chrom,
271            chrom_start,
272            chrom_end,
273            track_len,
274            raw_scores,
275            scores,
276            coverage: vec![0; track_len],
277            fire_coverage: vec![0; track_len],
278            msp_coverage: vec![0; track_len],
279            nuc_coverage: vec![0; track_len],
280            cpg_coverage: vec![0; track_len],
281            m6a_coverage: vec![0; track_len],
282            fire_track_opts,
283            shuffled_fibers,
284            cur_offset: 0,
285            fibers_seen: HashMap::new(),
286            fire_elements,
287            next_fire_id: 0,
288        }
289    }
290
291    //#[inline]
292    fn add_range_set(
293        array: &mut [i32],
294        view: &bamannotations::AnnotationTypeView<'_>,
295        cur_offset: i64,
296        chrom_start: usize,
297    ) {
298        for info in view.infos() {
299            match (info.ref_start, info.ref_end) {
300                (Some(rs), Some(re)) => {
301                    let rs = rs as i64;
302                    let re = re as i64;
303                    let re = if rs == re { re + 1 } else { re };
304                    for i in rs..re {
305                        let pos = i + cur_offset - chrom_start as i64;
306                        // skip if pos is before the start of the track
307                        if pos < 0 || pos >= array.len() as i64 {
308                            continue;
309                        }
310                        array[pos as usize] += 1;
311                    }
312                }
313                _ => continue,
314            }
315        }
316    }
317
318    fn fiber_start_and_end(&self, fiber: &FiberseqData) -> (i64, i64) {
319        if !self.fire_track_opts.fiber_coverage {
320            return (
321                fiber.record.reference_start() + self.cur_offset,
322                fiber.record.reference_end() + self.cur_offset,
323            );
324        }
325        let mut start = i64::MAX;
326        let mut end = i64::MIN;
327        for info in fiber.msp().infos() {
328            if let (Some(rs), Some(re)) = (info.ref_start, info.ref_end) {
329                start = std::cmp::min(start, rs as i64);
330                end = std::cmp::max(end, re as i64);
331            }
332        }
333        for info in fiber.nuc().infos() {
334            if let (Some(rs), Some(re)) = (info.ref_start, info.ref_end) {
335                start = std::cmp::min(start, rs as i64);
336                end = std::cmp::max(end, re as i64);
337            }
338        }
339        if start == i64::MAX {
340            start = fiber.record.reference_start();
341        }
342        if end == i64::MIN {
343            end = fiber.record.reference_end();
344        }
345        (start + self.cur_offset, end + self.cur_offset)
346    }
347
348    pub fn update_with_fiber(&mut self, fiber: &FiberseqData) {
349        // skip this fiber if it has no MSP/NUC information
350        // and we are looking at fiber_coverage
351        if self.fire_track_opts.fiber_coverage && fiber.msp().is_empty() && fiber.nuc().is_empty() {
352            return;
353        }
354
355        // Store fiber information for later shuffle generation (only for real, not shuffled)
356        if self.cur_offset == 0 && !self.fire_track_opts.shuffle {
357            let fiber_name = fiber.get_qname();
358            let original_start = fiber.record.reference_start();
359            let fiber_len = fiber.record.reference_end() - original_start;
360            self.fibers_seen
361                .insert((fiber_name, original_start), fiber_len);
362        }
363
364        // find the offset if we are shuffling data
365        // Priority: 1) shuffled_fibers from file, 2) random shuffle, 3) no shuffle
366        self.cur_offset = match self.shuffled_fibers {
367            Some(shuffled_fibers) => {
368                // Use pre-computed shuffle from file
369                match shuffled_fibers.get_shuffle_offset(fiber) {
370                    Some(offset) => offset,
371                    None => return, // skip missing fiber if it is not in the shuffle
372                }
373            }
374            None if self.fire_track_opts.random_shuffle => {
375                // Generate random shuffle offset
376                generate_random_shuffle_offset(
377                    fiber,
378                    self.chrom_end,
379                    self.fire_track_opts.shuffle_seed,
380                )
381                .unwrap_or(0)
382            }
383            None => 0, // No shuffling
384        };
385
386        if self.cur_offset != 0 && self.chrom_start != 0 {
387            panic!("Cannot apply shuffling unless the entire chromosome is being read at once.");
388        }
389
390        let (start, end) = self.fiber_start_and_end(fiber);
391        // calculate the coverage
392        for i in start..end {
393            let pos = i - self.chrom_start as i64;
394            if pos < 0 || pos >= self.track_len as i64 {
395                continue;
396            }
397            self.coverage[pos as usize] += 1;
398        }
399
400        // calculate the fire coverage and fire score. FIRE is its own
401        // annotation type — already filtered to non-zero precision in
402        // add_fire_to_rec. MIN_FIRE_QUAL is the stricter pileup-level gate.
403        for info in fiber.fire().infos() {
404            let (rs, re) = match (info.ref_start, info.ref_end) {
405                (Some(rs), Some(re)) => (rs as i64, re as i64),
406                _ => continue,
407            };
408            let qual = info.qualities.first().copied().unwrap_or(0);
409            if qual < MIN_FIRE_QUAL {
410                continue;
411            }
412            // Cap quality at 253 to avoid log10(0) issues, and cap score at 100
413            let capped_qual = qual.min(253) as f32;
414            let score_update = ((1.0 - capped_qual / 255.0).log10() * -50.0).min(100.0);
415
416            // If tracking FIRE elements, create a FireElement for this MSP
417            let fire_element = if self.fire_track_opts.track_fire_elements {
418                let elem_start = rs + self.cur_offset;
419                let elem_end = re + self.cur_offset;
420                let fire_id = self.next_fire_id;
421                self.next_fire_id += 1;
422                Some(FireElement {
423                    start: elem_start,
424                    end: elem_end,
425                    id: fire_id,
426                })
427            } else {
428                None
429            };
430
431            for i in rs..re {
432                let pos = i + self.cur_offset - self.chrom_start as i64;
433                if pos < 0 || pos >= self.track_len as i64 {
434                    continue;
435                }
436                self.fire_coverage[pos as usize] += 1;
437                self.raw_scores[pos as usize] += score_update;
438
439                // Store the FIRE element at this position if tracking is enabled
440                if let (Some(fire_elements), Some(elem)) = (&mut self.fire_elements, fire_element) {
441                    fire_elements[pos as usize].push(elem);
442                }
443            }
444        }
445
446        // add other sets of data to the FireTrack depending on CLI opts
447        let mut pairs = vec![];
448        if !self.fire_track_opts.no_nuc {
449            pairs.push((&mut self.nuc_coverage, fiber.nuc()));
450        }
451        if !self.fire_track_opts.no_msp {
452            pairs.push((&mut self.msp_coverage, fiber.msp()));
453        }
454        if self.fire_track_opts.m6a {
455            pairs.push((&mut self.m6a_coverage, fiber.m6a()));
456        }
457        if self.fire_track_opts.cpg {
458            pairs.push((&mut self.cpg_coverage, fiber.cpg()));
459        }
460
461        for (array, view) in pairs {
462            Self::add_range_set(array, &view, self.cur_offset, self.chrom_start);
463        }
464    }
465
466    pub fn calculate_scores(&mut self, min_fire_coverage: Option<i32>) {
467        let min_fire_coverage = min_fire_coverage.unwrap_or(MIN_FIRE_COVERAGE);
468        for i in 0..self.track_len {
469            if self.fire_coverage[i] <= 0 {
470                self.scores[i] = -1.0;
471            } else if self.fire_coverage[i] < min_fire_coverage && !self.fire_track_opts.shuffle {
472                // there is no minimum fire coverage if we are shuffling
473                self.scores[i] = -1.0;
474            } else {
475                self.scores[i] = self.raw_scores[i] / self.coverage[i] as f32;
476            }
477        }
478    }
479
480    pub fn calculate_rolling_max_score(&mut self) -> Vec<f32> {
481        let mut rolling_max = vec![-1.0; self.track_len];
482        let window_size = self.fire_track_opts.rolling_max.unwrap();
483        let look_back = window_size / 2;
484        for (i, cur_roll_max) in rolling_max.iter_mut().enumerate().take(self.track_len) {
485            let start = i.saturating_sub(look_back);
486            let mut end = i + look_back;
487            if end > self.track_len {
488                end = self.track_len;
489            }
490            let relevant_scores = &self.scores[start..end];
491            *cur_roll_max = *relevant_scores
492                .iter()
493                .max_by_key(|x| NotNan::new(**x).unwrap())
494                .unwrap_or(&-1.0);
495        }
496        rolling_max
497    }
498
499    pub fn row(&self, i: usize) -> FireRow<'_> {
500        FireRow {
501            score: &self.scores[i],
502            coverage: &self.coverage[i],
503            fire_coverage: &self.fire_coverage[i],
504            msp_coverage: &self.msp_coverage[i],
505            nuc_coverage: &self.nuc_coverage[i],
506            cpg_coverage: &self.cpg_coverage[i],
507            m6a_coverage: &self.m6a_coverage[i],
508            fire_track_opts: &self.fire_track_opts,
509        }
510    }
511
512    /// Calculate median coverage across the track
513    /// Returns (median_coverage, positions_with_coverage, positions_with_fire)
514    pub fn median_coverage(&self) -> (f64, usize) {
515        let mut coverages: Vec<i32> = self.coverage.iter().filter(|&&c| c > 0).copied().collect();
516
517        let positions_with_coverage = coverages.len();
518
519        if coverages.is_empty() {
520            return (0.0, 0);
521        }
522
523        coverages.sort_unstable();
524
525        let median = if coverages.len().is_multiple_of(2) {
526            let mid = coverages.len() / 2;
527            (coverages[mid - 1] as f64 + coverages[mid] as f64) / 2.0
528        } else {
529            coverages[coverages.len() / 2] as f64
530        };
531
532        (median, positions_with_coverage)
533    }
534
535    /// Calculate median and estimated standard deviation (sqrt of median for Poisson)
536    /// Returns (median, std_dev, positions_with_coverage)
537    pub fn median_and_std_coverage(&self) -> (f64, f64, usize) {
538        let (median, positions_with_coverage) = self.median_coverage();
539
540        // For sequencing data, we assume Poisson distribution where std_dev ≈ sqrt(median)
541        let std_dev = median.sqrt();
542
543        (median, std_dev, positions_with_coverage)
544    }
545
546    /// Generate a ShuffledFibers HashMap from a list of fibers
547    /// This creates random shuffled positions for each fiber within the chromosome.
548    /// Uses the FireTrack's coverage to avoid placing shuffled fibers in regions with:
549    /// - Zero coverage (always avoided)
550    /// - Coverage below min_cov (if specified)
551    /// - Coverage above max_cov (if specified)
552    ///
553    /// Will retry up to 1000 times to find a valid position.
554    pub fn generate_shuffled_positions(
555        &self,
556        seed: Option<u64>,
557        min_cov: Option<i32>,
558        max_cov: Option<i32>,
559    ) -> ShuffledFibers {
560        use rand::rngs::StdRng;
561        use rand::{Rng, SeedableRng};
562        use std::collections::hash_map::DefaultHasher;
563        use std::hash::{Hash, Hasher};
564
565        let mut shuffled_fiber_starts = HashMap::new();
566        let mut regenerated_count = 0;
567
568        for ((fiber, original_start), fiber_len) in &self.fibers_seen {
569            let max_start = (self.track_len as i64 - fiber_len).max(0);
570            let coverage = self.coverage[*original_start as usize];
571
572            // filter by coverage constraints before attempting to shuffle
573            let has_valid_coverage = coverage > 0
574                && min_cov.is_none_or(|min| coverage >= min)
575                && max_cov.is_none_or(|max| coverage <= max);
576            if !has_valid_coverage {
577                continue;
578            }
579
580            // Create deterministic seed from fiber name
581            let mut hasher = DefaultHasher::new();
582            fiber.hash(&mut hasher);
583            if let Some(s) = seed {
584                s.hash(&mut hasher);
585            }
586            let fiber_seed = hasher.finish();
587
588            // Generate random position, retrying up to 1000 times if coverage is invalid
589            let mut rng = StdRng::seed_from_u64(fiber_seed);
590            let mut shuffled_start = rng.gen_range(0..=max_start);
591
592            // Try up to 1000 times to find a position with valid coverage
593            let mut attempts = 0;
594            while attempts < 1000 {
595                // Check if this position has valid coverage
596                // No bounds check needed: shuffled_start is guaranteed to be in [0, max_start]
597                // where max_start = track_len - fiber_len, so it's always valid
598                let cov = self.coverage[shuffled_start as usize];
599
600                // Check coverage constraints
601                let has_valid_coverage = cov > 0
602                    && min_cov.is_none_or(|min| cov >= min)
603                    && max_cov.is_none_or(|max| cov <= max);
604
605                if has_valid_coverage {
606                    if attempts > 0 {
607                        regenerated_count += 1;
608                    }
609                    break;
610                }
611
612                // Regenerate position
613                shuffled_start = rng.gen_range(0..=max_start);
614                attempts += 1;
615            }
616
617            // Store as (chrom, fiber_name, original_start) -> shuffled_start
618            let key = (self.chrom.clone(), fiber.clone(), *original_start);
619            shuffled_fiber_starts.insert(key, shuffled_start);
620        }
621
622        log::debug!(
623            "Generated shuffle positions for {} fibers ({} regenerated for valid coverage [min={:?}, max={:?}])",
624            shuffled_fiber_starts.len(),
625            regenerated_count,
626            min_cov,
627            max_cov
628        );
629
630        ShuffledFibers {
631            shuffled_fiber_starts,
632        }
633    }
634}
635
636/// Options needed for FiberseqPileup
637/// This is a lightweight struct that only contains the options actually used by FiberseqPileup
638#[derive(Debug, Clone)]
639pub struct FiberseqPileupOptions {
640    /// Track options for FireTrack (coverage, marks, etc.)
641    pub fire_track_opts: FireTrackOptions,
642    /// Output rolling max of the score column over X bases
643    pub rolling_max: Option<usize>,
644    /// Include haplotype-specific tracks
645    pub haps: bool,
646    /// Write output one base at a time even if values don't change
647    pub per_base: bool,
648    /// Keep zero coverage regions
649    pub keep_zeros: bool,
650    /// Minimum FIRE coverage required to calculate a score (default: 4)
651    pub min_fire_coverage: Option<i32>,
652}
653
654impl From<&PileupOptions> for FiberseqPileupOptions {
655    fn from(opts: &PileupOptions) -> Self {
656        Self {
657            fire_track_opts: FireTrackOptions::from(opts),
658            rolling_max: opts.rolling_max,
659            haps: opts.haps,
660            per_base: opts.per_base,
661            keep_zeros: opts.keep_zeros,
662            min_fire_coverage: None, // Use default
663        }
664    }
665}
666
667#[derive(Debug)]
668pub struct FiberseqPileup<'a> {
669    pub all_data: FireTrack<'a>,
670    pub hap1_data: Option<FireTrack<'a>>,
671    pub hap2_data: Option<FireTrack<'a>>,
672    pub shuffled_data: Option<FireTrack<'a>>,
673    pub chrom: String,
674    pub chrom_start: usize,
675    pub chrom_end: usize,
676    pub track_len: usize,
677    has_data: bool,
678    pileup_opts: FiberseqPileupOptions,
679    shuffled_fibers: &'a Option<ShuffledFibers>,
680    pub rolling_max: Option<Vec<f32>>,
681    pub fdr_scores: Option<Vec<f64>>,
682    pub region_name: Option<String>,
683}
684
685impl<'a> FiberseqPileup<'a> {
686    pub fn new(
687        chrom: &str,
688        chrom_start: usize,
689        chrom_end: usize,
690        pileup_opts: FiberseqPileupOptions,
691        shuffled_fibers: &'a Option<ShuffledFibers>,
692        region_name: Option<String>,
693    ) -> Self {
694        let track_len = chrom_end - chrom_start + 1;
695        let fire_track_opts = pileup_opts.fire_track_opts.clone();
696        let all_data = FireTrack::new(
697            chrom.to_string(),
698            chrom_start,
699            chrom_end,
700            fire_track_opts.clone(),
701            &None,
702        );
703        let (hap1_data, hap2_data) = if pileup_opts.haps {
704            (
705                Some(FireTrack::new(
706                    chrom.to_string(),
707                    chrom_start,
708                    chrom_end,
709                    fire_track_opts.clone(),
710                    &None,
711                )),
712                Some(FireTrack::new(
713                    chrom.to_string(),
714                    chrom_start,
715                    chrom_end,
716                    fire_track_opts.clone(),
717                    &None,
718                )),
719            )
720        } else {
721            (None, None)
722        };
723
724        let shuffled_data = if shuffled_fibers.is_some() {
725            let mut shuffled_opts = fire_track_opts.clone();
726            shuffled_opts.shuffle = true;
727            Some(FireTrack::new(
728                chrom.to_string(),
729                chrom_start,
730                chrom_end,
731                shuffled_opts,
732                shuffled_fibers,
733            ))
734        } else {
735            None
736        };
737
738        Self {
739            all_data,
740            hap1_data,
741            hap2_data,
742            shuffled_data,
743            chrom: chrom.to_string(),
744            chrom_start,
745            chrom_end,
746            track_len,
747            has_data: false,
748            pileup_opts,
749            shuffled_fibers,
750            rolling_max: None,
751            fdr_scores: None,
752            region_name,
753        }
754    }
755
756    pub fn has_data(&self) -> bool {
757        self.has_data
758    }
759
760    /// Calculate and store FDR scores for each position based on the FDR table
761    pub fn calculate_fdr_scores(&mut self, fdr_table: &[crate::subcommands::call_peaks::FdrEntry]) {
762        let mut fdr_scores = vec![1.0; self.track_len];
763
764        if fdr_table.is_empty() {
765            self.fdr_scores = Some(fdr_scores);
766            return;
767        }
768
769        for (i, &score) in self.all_data.scores.iter().enumerate() {
770            if score < 0.0 {
771                // No coverage, FDR = 1.0 (already set)
772                continue;
773            }
774
775            // Binary search to find the FDR for this score
776            let search_result = fdr_table.binary_search_by(|entry| {
777                entry
778                    .threshold
779                    .partial_cmp(&(score as f64))
780                    .unwrap_or(std::cmp::Ordering::Equal)
781            });
782
783            fdr_scores[i] = match search_result {
784                Result::Ok(idx) => fdr_table[idx].fdr,
785                Result::Err(idx) => {
786                    if idx == 0 {
787                        fdr_table[0].fdr
788                    } else {
789                        fdr_table[idx - 1].fdr
790                    }
791                }
792            };
793        }
794
795        self.fdr_scores = Some(fdr_scores);
796    }
797
798    /// Add fibers from an iterator
799    /// This is more efficient than add_records as it works directly with FiberseqData
800    pub fn add_fibers(&mut self, fibers: impl Iterator<Item = FiberseqData>) {
801        for fiber in fibers {
802            self.has_data = true;
803
804            // skip if the fiber was unable to be shuffled
805            if self.shuffled_fibers.is_some()
806                && !self.shuffled_fibers.as_ref().unwrap().has_fiber(&fiber)
807            {
808                continue;
809            }
810
811            self.all_data.update_with_fiber(&fiber);
812            // add hap1 data
813            if let Some(hap1_data) = &mut self.hap1_data {
814                if fiber.get_hp() == "H1" {
815                    hap1_data.update_with_fiber(&fiber);
816                }
817            }
818            // add hap2 data
819            if let Some(hap2_data) = &mut self.hap2_data {
820                if fiber.get_hp() == "H2" {
821                    hap2_data.update_with_fiber(&fiber);
822                }
823            }
824            // add shuffled data
825            if let Some(shuffled_data) = &mut self.shuffled_data {
826                shuffled_data.update_with_fiber(&fiber);
827            }
828        }
829
830        self.calculate_scores();
831    }
832
833    pub fn header(pileup_opts: &PileupOptions, include_name: bool) -> String {
834        let mut header = format!("{}\t{}\t{}", "#chrom", "start", "end");
835
836        let mut suffixes = vec![""];
837        if pileup_opts.haps {
838            suffixes.push("_H1");
839            suffixes.push("_H2");
840        }
841        if pileup_opts.shuffle.is_some() {
842            suffixes.push("_shuffled");
843        }
844
845        for suffix in suffixes {
846            header += &format!(
847                "\t{}{suffix}\t{}{suffix}\t{}{suffix}",
848                "coverage", "fire_coverage", "score",
849            );
850            if !pileup_opts.no_nuc {
851                header += &format!("\t{}{suffix}", "nuc_coverage");
852            }
853            if !pileup_opts.no_msp {
854                header += &format!("\t{}{suffix}", "msp_coverage");
855            }
856            if pileup_opts.m6a {
857                header += &format!("\t{}{suffix}", "m6a_coverage");
858            }
859            if pileup_opts.cpg {
860                header += &format!("\t{}{suffix}", "cpg_coverage");
861            }
862        }
863        if pileup_opts.rolling_max.is_some() {
864            header += "\trolling_max";
865        }
866        // Add name column at the end to minimize breaking downstream tools
867        if include_name {
868            header += "\tname";
869        }
870        header += "\n";
871        header
872    }
873
874    fn calculate_scores(&mut self) {
875        self.all_data
876            .calculate_scores(self.pileup_opts.min_fire_coverage);
877        // calculate rolling max
878        if self.pileup_opts.rolling_max.is_some() {
879            self.rolling_max = Some(self.all_data.calculate_rolling_max_score());
880        }
881        // scores for other tracks
882        if let Some(hap1_data) = &mut self.hap1_data {
883            hap1_data.calculate_scores(self.pileup_opts.min_fire_coverage);
884        }
885        if let Some(hap2_data) = &mut self.hap2_data {
886            hap2_data.calculate_scores(self.pileup_opts.min_fire_coverage);
887        }
888        if let Some(shuffled_data) = &mut self.shuffled_data {
889            shuffled_data.calculate_scores(self.pileup_opts.min_fire_coverage);
890        }
891    }
892
893    /// check if the ith row has the same data as the previous row
894    /// if it does, return true, otherwise return false
895    /// this is used to determine if the data should be written to the output
896    pub fn is_same_as_previous(&self, i: usize) -> bool {
897        if i == 0 {
898            true
899        } else {
900            let total_same = self.all_data.row(i) == self.all_data.row(i - 1);
901            let haps_same = if self.pileup_opts.haps {
902                let h1 = self.hap1_data.as_ref().unwrap();
903                let h2 = self.hap2_data.as_ref().unwrap();
904                let hap1_same = h1.row(i) == h1.row(i - 1);
905                let hap2_same = h2.row(i) == h2.row(i - 1);
906                hap1_same && hap2_same
907            } else {
908                true
909            };
910            let shuffled_same = if let Some(shuffled) = self.shuffled_data.as_ref() {
911                shuffled.row(i) == shuffled.row(i - 1)
912            } else {
913                true
914            };
915            total_same && haps_same && shuffled_same
916        }
917    }
918
919    fn wait_to_write(&self, i: usize) -> bool {
920        // write every row
921        if self.pileup_opts.per_base {
922            return false;
923        }
924        self.is_same_as_previous(i) && i != self.track_len - 1
925    }
926
927    pub fn log_stats(&self) {
928        let mut data_tracks = vec![&self.all_data];
929        if self.pileup_opts.haps {
930            data_tracks.push(self.hap1_data.as_ref().unwrap());
931            data_tracks.push(self.hap2_data.as_ref().unwrap());
932        }
933        if let Some(shuffled) = self.shuffled_data.as_ref() {
934            data_tracks.push(shuffled);
935        }
936        for data in data_tracks {
937            let total_coverage: i64 = data.coverage.iter().map(|x| *x as i64).sum();
938            let total_fire_coverage: i64 = data.fire_coverage.iter().map(|x| *x as i64).sum();
939            let total_score: f64 = data.scores.iter().map(|x| *x as f64).sum();
940            log::info!(
941                "Total coverage: {total_coverage}, Total fire coverage: {total_fire_coverage}, Total score: {total_score}"
942            );
943        }
944    }
945
946    pub fn write(&self, out: &mut Box<dyn Write>) -> Result<(), anyhow::Error> {
947        if !self.has_data {
948            return Ok(());
949        }
950        if self.shuffled_data.is_some() {
951            self.log_stats();
952        }
953        let mut write_start_index = 0;
954        let mut write_end_index = 1;
955        for i in 1..self.track_len {
956            // do we have the same data as the previous row?
957            if self.wait_to_write(i) {
958                write_end_index = i + 1;
959            } else {
960                let mut line = format!(
961                    "{}\t{}\t{}",
962                    self.chrom,
963                    write_start_index + self.chrom_start,
964                    write_end_index + self.chrom_start
965                );
966
967                let mut data_tracks = vec![&self.all_data];
968                if self.pileup_opts.haps {
969                    data_tracks.push(self.hap1_data.as_ref().unwrap());
970                    data_tracks.push(self.hap2_data.as_ref().unwrap());
971                }
972                if let Some(shuffled) = self.shuffled_data.as_ref() {
973                    data_tracks.push(shuffled);
974                }
975
976                for data in data_tracks {
977                    line += data.row(write_start_index).to_string().as_str();
978                }
979                if self.pileup_opts.rolling_max.is_some() {
980                    line += &format!(
981                        "\t{}",
982                        self.rolling_max.as_ref().unwrap()[write_start_index]
983                    );
984                }
985                // Add name column at the end to minimize breaking downstream tools
986                if let Some(name) = &self.region_name {
987                    line += &format!("\t{}", name);
988                }
989                // don't write empty lines unless keep_zeros is set
990                let mut cov = self.all_data.coverage[write_start_index];
991                if let Some(shuffled_data) = &self.shuffled_data {
992                    cov += shuffled_data.coverage[write_start_index];
993                }
994                if self.pileup_opts.keep_zeros || cov > 0 {
995                    line += "\n";
996                    bio_io::write_to_file(&line, out);
997                }
998                // reset the write indexes
999                write_start_index = i;
1000                write_end_index = i + 1;
1001            }
1002        }
1003        out.flush()?;
1004        Ok(())
1005    }
1006}
1007
1008/// split up a FetchDefinition into multiple regions of a certain size
1009/// TODO set up run_rgn to take a list of regions and multithread it
1010pub fn split_fetch_definition(
1011    rgn: &FetchDefinition,
1012    chrom_len: usize,
1013    window_size: usize,
1014) -> Vec<(i64, i64)> {
1015    let (start, end) = match rgn {
1016        FetchDefinition::RegionString(_chrom, start, end) => (*start, *end),
1017        _ => (0, chrom_len as i64),
1018    };
1019    let mut rgns = vec![];
1020    let mut cur_start = start;
1021    while cur_start < end {
1022        let cur_end = std::cmp::min(cur_start + window_size as i64, end);
1023        rgns.push((cur_start, cur_end));
1024        cur_start = cur_end;
1025    }
1026    rgns
1027}
1028
1029fn run_rgn(
1030    chrom: &str,
1031    rgn: FetchDefinition,
1032    bam: &mut IndexedReader,
1033    out: &mut Box<dyn Write>,
1034    pileup_opts: &PileupOptions,
1035    shuffled_fibers: &Option<ShuffledFibers>,
1036    region_name: Option<String>,
1037) -> Result<(), anyhow::Error> {
1038    let tid = bam.header().tid(chrom.as_bytes()).ok_or(anyhow::anyhow!(
1039        "Chromosome {} not found in BAM header",
1040        chrom
1041    ))?;
1042    let chrom_len = bam.header().target_len(tid).ok_or(anyhow::anyhow!(
1043        "Chromosome {} length not found in BAM header",
1044        chrom
1045    ))? as i64;
1046
1047    let window_size = if shuffled_fibers.is_some() {
1048        (chrom_len + 1) as usize
1049    } else {
1050        WINDOW_SIZE
1051    };
1052    log::info!("Window size on {chrom}: {window_size}");
1053
1054    let windows = split_fetch_definition(&rgn, chrom_len as usize, window_size);
1055    log::debug!("Splitting {} into {} windows", chrom, windows.len());
1056    for (chrom_start, mut chrom_end) in windows {
1057        if chrom_start >= chrom_len {
1058            continue;
1059        } else if chrom_end > chrom_len {
1060            chrom_end = chrom_len;
1061        }
1062
1063        // Fetch fibers from the region using the new iterator-based approach
1064        let fiber_iter =
1065            pileup_opts
1066                .input
1067                .fetch_fibers(bam, chrom, Some(chrom_start), Some(chrom_end))?;
1068
1069        // make the pileup
1070        log::debug!("Initializing pileup for {chrom}:{chrom_start}-{chrom_end}");
1071        let mut pileup = FiberseqPileup::new(
1072            chrom,
1073            chrom_start as usize,
1074            chrom_end as usize,
1075            pileup_opts.into(),
1076            shuffled_fibers,
1077            region_name.clone(),
1078        );
1079        pileup.add_fibers(fiber_iter);
1080
1081        // Only write if we have data
1082        if pileup.has_data() {
1083            pileup.write(out)?;
1084        }
1085    }
1086
1087    Ok(())
1088}
1089
1090/// extract existing fire calls into a bed9+ like file
1091pub fn pileup_track(pileup_opts: &mut PileupOptions) -> Result<(), anyhow::Error> {
1092    // read in the bam from stdin or from a file
1093    let mut bam = pileup_opts.input.indexed_bam_reader();
1094    let header = pileup_opts.input.header_view();
1095
1096    let mut out = bio_io::writer(&pileup_opts.out)?;
1097
1098    let shuffled_fibers = match &pileup_opts.shuffle {
1099        Some(file_path) => Some(ShuffledFibers::new(file_path)?),
1100        None => None,
1101    };
1102
1103    // Handle regions based on source (BED file, command-line args, or all chromosomes)
1104    // We process regions immediately rather than collecting them because FetchDefinition has lifetime constraints
1105    if let Some(bed_path) = &pileup_opts.bed {
1106        // Parse BED file
1107        let bed_records = bio_io::read_bed_regions(bed_path)?;
1108        let include_name = bed_records.iter().any(|r| r.name.is_some());
1109
1110        // add the header
1111        out.write_all(FiberseqPileup::header(pileup_opts, include_name).as_bytes())?;
1112
1113        // Process each BED record immediately
1114        for rec in bed_records {
1115            let fetch_def = FetchDefinition::RegionString(rec.chrom.as_bytes(), rec.start, rec.end);
1116            // If any record has a name, use "." for records without names to keep column count consistent
1117            let region_name = if include_name {
1118                Some(rec.name.unwrap_or_else(|| ".".to_string()))
1119            } else {
1120                None
1121            };
1122            run_rgn(
1123                &rec.chrom,
1124                fetch_def,
1125                &mut bam,
1126                &mut out,
1127                pileup_opts,
1128                &shuffled_fibers,
1129                region_name,
1130            )?;
1131        }
1132    } else if !pileup_opts.rgn.is_empty() {
1133        // Use command-line regions
1134        out.write_all(FiberseqPileup::header(pileup_opts, false).as_bytes())?;
1135
1136        for rgn_str in &pileup_opts.rgn {
1137            let (rgn, chrom) = region_parser(rgn_str);
1138            run_rgn(
1139                &chrom,
1140                rgn,
1141                &mut bam,
1142                &mut out,
1143                pileup_opts,
1144                &shuffled_fibers,
1145                None,
1146            )?;
1147        }
1148    } else {
1149        // Process all chromosomes
1150        out.write_all(FiberseqPileup::header(pileup_opts, false).as_bytes())?;
1151
1152        for chrom in header.target_names() {
1153            let chrom_str = String::from_utf8_lossy(chrom).to_string();
1154            let rgn = FetchDefinition::String(chrom);
1155            run_rgn(
1156                &chrom_str,
1157                rgn,
1158                &mut bam,
1159                &mut out,
1160                pileup_opts,
1161                &shuffled_fibers,
1162                None,
1163            )?;
1164        }
1165    }
1166
1167    Ok(())
1168}