Skip to main content

fibertools_rs/
fiber.rs

1use super::subcommands::center::CenterPosition;
2use super::utils::input_bam::FiberFilters;
3use super::*;
4use crate::utils::bamannotations::*;
5use crate::utils::basemods::{CPG_TYPE, M6A_TYPE};
6use crate::utils::bio_io::*;
7use crate::utils::ftexpression::apply_filter_fsd;
8use crate::utils::ma_io::{FIRE_TYPE, MSP_TYPE, NUC_TYPE};
9use molecular_annotation::MolecularAnnotations;
10use rayon::prelude::*;
11use rust_htslib::bam::Read;
12use rust_htslib::{bam, bam::ext::BamRecordExtensions, bam::record::Aux, bam::HeaderView};
13use std::collections::HashMap;
14use std::fmt::Write;
15
16#[derive(Debug, Clone, PartialEq)]
17pub struct FiberseqData {
18    pub record: bam::Record,
19    pub annotations: MolecularAnnotations,
20    pub ec: f32,
21    pub target_name: String,
22    pub rg: String,
23    pub center_position: Option<CenterPosition>,
24}
25
26impl FiberseqData {
27    pub fn new(record: bam::Record, target_name: Option<&String>, filters: &FiberFilters) -> Self {
28        // read group
29        let rg = if let Ok(Aux::String(f)) = record.aux(b"RG") {
30            log::trace!("{f}");
31            f
32        } else {
33            "."
34        }
35        .to_string();
36        let mut annotations = crate::utils::ma_io::read_record(&record).unwrap_or_else(|e| {
37            log::warn!("Failed to read annotations: {e}");
38            MolecularAnnotations::from_record(&record)
39        });
40
41        // The library populates m6a/cpg from MM/ML on read; apply the
42        // fibertools read-side basemod filters (min ML score, end-strip
43        // distance) on top of that here.
44        if filters.min_ml_score > 0 || filters.strip_starting_basemods > 0 {
45            let seq_len = record.seq_len();
46            let strip = filters.strip_starting_basemods.max(0) as usize;
47            let upper = seq_len.saturating_sub(strip);
48            let min_ml = filters.min_ml_score;
49            for t in annotations.annotation_types.iter_mut() {
50                if !crate::utils::basemods::is_basemod_type(&t.name) {
51                    continue;
52                }
53                t.annotations.retain(|a| {
54                    if a.qualities.first().copied().unwrap_or(0) < min_ml {
55                        return false;
56                    }
57                    if strip > 0 {
58                        let p = a.start as usize;
59                        if p < strip || p >= upper {
60                            return false;
61                        }
62                    }
63                    true
64                });
65            }
66        }
67
68        // get the number of passes
69        let ec = if let Ok(Aux::Float(f)) = record.aux(b"ec") {
70            log::trace!("{f}");
71            f
72        } else {
73            0.0
74        };
75
76        let target_name = match target_name {
77            Some(t) => t.clone(),
78            None => ".".to_string(),
79        };
80
81        let mut fsd = FiberseqData {
82            record,
83            annotations,
84            ec,
85            target_name,
86            rg,
87            center_position: None,
88        };
89
90        apply_filter_fsd(&mut fsd, filters).expect("Failed to apply filter to FiberseqData");
91        fsd
92    }
93
94    pub fn dict_from_head_view(head_view: &HeaderView) -> HashMap<i32, String> {
95        if head_view.target_count() == 0 {
96            return HashMap::new();
97        }
98        let target_u8s = head_view.target_names();
99        let tids = target_u8s
100            .iter()
101            .map(|t| head_view.tid(t).expect("Unable to get tid"));
102        let target_names = target_u8s
103            .iter()
104            .map(|&a| String::from_utf8_lossy(a).to_string());
105
106        tids.zip(target_names)
107            .map(|(id, t)| (id as i32, t))
108            .collect()
109    }
110
111    pub fn target_name_from_tid(tid: i32, target_dict: &HashMap<i32, String>) -> Option<&String> {
112        target_dict.get(&tid)
113    }
114
115    pub fn from_records(
116        records: Vec<bam::Record>,
117        head_view: &HeaderView,
118        filters: &FiberFilters,
119    ) -> Vec<Self> {
120        let target_dict = Self::dict_from_head_view(head_view);
121        records
122            .into_par_iter()
123            .map(|r| {
124                let tid = r.tid();
125                (r, Self::target_name_from_tid(tid, &target_dict))
126            })
127            .map(|(r, target_name)| Self::new(r, target_name, filters))
128            .collect::<Vec<_>>()
129    }
130
131    //
132    // GET FUNCTIONS
133    //
134
135    /// View over `msp` annotations derived from `self.annotations`.
136    pub fn msp(&self) -> AnnotationTypeView<'_> {
137        AnnotationTypeView::new(&self.annotations, MSP_TYPE)
138    }
139
140    /// View over `nuc` annotations derived from `self.annotations`.
141    pub fn nuc(&self) -> AnnotationTypeView<'_> {
142        AnnotationTypeView::new(&self.annotations, NUC_TYPE)
143    }
144
145    /// View over `m6a` annotations derived from `self.annotations`.
146    pub fn m6a(&self) -> AnnotationTypeView<'_> {
147        AnnotationTypeView::new(&self.annotations, M6A_TYPE)
148    }
149
150    /// View over `cpg` annotations derived from `self.annotations`.
151    pub fn cpg(&self) -> AnnotationTypeView<'_> {
152        AnnotationTypeView::new(&self.annotations, CPG_TYPE)
153    }
154
155    /// View over `fire` annotations derived from `self.annotations`.
156    pub fn fire(&self) -> AnnotationTypeView<'_> {
157        AnnotationTypeView::new(&self.annotations, FIRE_TYPE)
158    }
159
160    /// Flush `self.annotations` onto the record's MA-family aux tags. The
161    /// single write path for subcommands that edit nuc/msp/fire annotations;
162    /// call this, then hand the record to the BAM writer.
163    ///
164    /// MM/ML are deliberately left untouched: this is an edit path, not a
165    /// basemod producer, so the record's original MM/ML bytes pass through
166    /// byte-identically. Basemod types (m6a/cpg) are `Encoding::MmMl` (set when
167    /// they were read), so `write_record` excludes them from the MA tag rather
168    /// than writing basemod calls there too. Producers that synthesize or
169    /// modify basemods must use `ma_io::write_record_with_basemods` instead.
170    pub fn serialize_annotations(&mut self) {
171        crate::utils::ma_io::write_record(&mut self.record, &self.annotations);
172    }
173
174    pub fn get_qname(&self) -> String {
175        String::from_utf8_lossy(self.record.qname()).to_string()
176    }
177
178    pub fn get_rq(&self) -> Option<f32> {
179        if let Ok(Aux::Float(f)) = self.record.aux(b"rq") {
180            Some(f)
181        } else {
182            None
183        }
184    }
185
186    pub fn get_hp(&self) -> String {
187        match self.record.aux(b"HP") {
188            Ok(Aux::U8(v)) => format!("H{v}"),
189            Ok(Aux::I8(v)) => format!("H{v}"),
190            Ok(Aux::U16(v)) => format!("H{v}"),
191            Ok(Aux::I16(v)) => format!("H{v}"),
192            Ok(Aux::U32(v)) => format!("H{v}"),
193            Ok(Aux::I32(v)) => format!("H{v}"),
194            _ => "UNK".to_string(),
195        }
196    }
197
198    //
199    //  WRITE BED12 FUNCTIONS
200    //
201    pub fn write_msp(&self, reference: bool) -> String {
202        let msp = self.msp();
203        let (starts, _ends, lengths) = if reference {
204            (
205                msp.reference_starts(),
206                msp.reference_ends(),
207                msp.reference_lengths(),
208            )
209        } else {
210            (msp.option_starts(), msp.option_ends(), msp.option_lengths())
211        };
212        self.to_bed12(reference, &starts, &lengths, LINKER_COLOR)
213    }
214
215    pub fn write_nuc(&self, reference: bool) -> String {
216        let nuc = self.nuc();
217        let (starts, _ends, lengths) = if reference {
218            (
219                nuc.reference_starts(),
220                nuc.reference_ends(),
221                nuc.reference_lengths(),
222            )
223        } else {
224            (nuc.option_starts(), nuc.option_ends(), nuc.option_lengths())
225        };
226        self.to_bed12(reference, &starts, &lengths, NUC_COLOR)
227    }
228
229    pub fn write_m6a(&self, reference: bool) -> String {
230        let m6a = self.m6a();
231        let starts = if reference {
232            m6a.reference_starts()
233        } else {
234            m6a.option_starts()
235        };
236        let lengths = vec![Some(1); starts.len()];
237        self.to_bed12(reference, &starts, &lengths, M6A_COLOR)
238    }
239
240    pub fn write_cpg(&self, reference: bool) -> String {
241        let cpg = self.cpg();
242        let starts = if reference {
243            cpg.reference_starts()
244        } else {
245            cpg.option_starts()
246        };
247        let lengths = vec![Some(1); starts.len()];
248        self.to_bed12(reference, &starts, &lengths, CPG_COLOR)
249    }
250
251    pub fn to_bed12(
252        &self,
253        reference: bool,
254        starts: &[Option<i64>],
255        lengths: &[Option<i64>],
256        color: &str,
257    ) -> String {
258        if starts.is_empty() {
259            return "".to_string();
260        }
261        // skip if no alignments are here
262        if self.record.is_unmapped() && reference {
263            return "".to_string();
264        }
265
266        let ct;
267        let start;
268        let end;
269        let name = String::from_utf8_lossy(self.record.qname()).to_string();
270        let mut rtn: String = String::with_capacity(0);
271        if reference {
272            ct = &self.target_name;
273            start = self.record.reference_start();
274            end = self.record.reference_end();
275        } else {
276            ct = &name;
277            start = 0;
278            end = self.record.seq_len() as i64;
279        }
280        let score = self.ec.round() as i64;
281        let strand = if self.record.is_reverse() { '-' } else { '+' };
282        // filter out positions that do not have an exact liftover
283        let (filtered_starts, filtered_lengths): (Vec<i64>, Vec<i64>) = starts
284            .iter()
285            .flatten()
286            .zip(lengths.iter().flatten())
287            .unzip();
288        // skip empty ones
289        if filtered_lengths.is_empty() || filtered_starts.is_empty() {
290            return "".to_string();
291        }
292        let b_ct = filtered_starts.len() + 2;
293        let b_ln: String = filtered_lengths
294            .iter()
295            .map(|&ln| ln.to_string() + ",")
296            .collect();
297        let b_st: String = filtered_starts
298            .iter()
299            .map(|&st| (st - start).to_string() + ",")
300            .collect();
301        assert_eq!(filtered_lengths.len(), filtered_starts.len());
302
303        rtn.push_str(ct);
304        rtn.push('\t');
305        rtn.push_str(&start.to_string());
306        rtn.push('\t');
307        rtn.push_str(&end.to_string());
308        rtn.push('\t');
309        rtn.push_str(&name);
310        rtn.push('\t');
311        rtn.push_str(&score.to_string());
312        rtn.push('\t');
313        rtn.push(strand);
314        rtn.push('\t');
315        rtn.push_str(&start.to_string());
316        rtn.push('\t');
317        rtn.push_str(&end.to_string());
318        rtn.push('\t');
319        rtn.push_str(color);
320        rtn.push('\t');
321        rtn.push_str(&b_ct.to_string());
322        rtn.push_str("\t0,"); // add a zero length start
323        rtn.push_str(&b_ln);
324        rtn.push_str("1\t0,"); // add a 1 base length and a 0 start point
325        rtn.push_str(&b_st);
326        write!(&mut rtn, "{}", format_args!("{}\n", end - start - 1)).unwrap();
327        rtn
328    }
329
330    //
331    // WRITE ALL FUNCTIONS
332    //
333
334    pub fn all_header(simplify: bool, quality: bool) -> String {
335        let mut x = format!(
336            "#{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t",
337            "ct", "st", "en", "fiber", "score", "strand", "sam_flag", "HP", "RG", "fiber_length",
338        );
339        if !simplify {
340            x.push_str("fiber_sequence\t")
341        }
342        if quality {
343            x.push_str("fiber_qual\t")
344        }
345        x.push_str(&format!(
346            "{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\n",
347            "ec",
348            "rq",
349            "total_AT_bp",
350            "total_m6a_bp",
351            "total_nuc_bp",
352            "total_msp_bp",
353            "total_5mC_bp",
354            "nuc_starts",
355            "nuc_lengths",
356            "ref_nuc_starts",
357            "ref_nuc_lengths",
358            "msp_starts",
359            "msp_lengths",
360            "ref_msp_starts",
361            "ref_msp_lengths",
362            "fire_starts",
363            "fire_lengths",
364            "fire_qual",
365            "ref_fire_starts",
366            "ref_fire_lengths",
367            "m6a",
368            "ref_m6a",
369            "m6a_qual",
370            "5mC",
371            "ref_5mC",
372            "5mC_qual"
373        ));
374        x
375    }
376
377    pub fn write_all(&self, simplify: bool, quality: bool) -> String {
378        // PB features
379        let name = std::str::from_utf8(self.record.qname()).unwrap();
380        let score = self.ec.round() as i64;
381        let q_len = self.record.seq_len() as i64;
382        let rq = match self.get_rq() {
383            Some(x) => format!("{x}"),
384            None => ".".to_string(),
385        };
386        // reference features
387        let ct;
388        let start;
389        let end;
390        let strand;
391        if self.record.is_unmapped() {
392            ct = ".";
393            start = 0;
394            end = 0;
395            strand = '.';
396        } else {
397            ct = &self.target_name;
398            start = self.record.reference_start();
399            end = self.record.reference_end();
400            strand = if self.record.is_reverse() { '-' } else { '+' };
401        }
402        let sam_flag = self.record.flags();
403        let hp = self.get_hp();
404
405        let at_count = self
406            .record
407            .seq()
408            .as_bytes()
409            .iter()
410            .filter(|&x| *x == b'A' || *x == b'T')
411            .count() as i64;
412
413        // get the info
414        let m6a = self.m6a();
415        let cpg = self.cpg();
416        let msp = self.msp();
417        let nuc = self.nuc();
418        let fire = self.fire();
419        let m6a_count = m6a.len();
420        let m6a_qual = m6a.qual().iter().map(|a| Some(*a as i64)).collect();
421        let cpg_count = cpg.len();
422        let cpg_qual = cpg.qual().iter().map(|a| Some(*a as i64)).collect();
423        let fire_qual: Vec<Option<i64>> = fire.qual().iter().map(|a| Some(*a as i64)).collect();
424
425        // write the features
426        let mut rtn = String::with_capacity(0);
427        // add first things 7
428        rtn.write_fmt(format_args!(
429            "{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t",
430            ct, start, end, name, score, strand, sam_flag, hp, self.rg, q_len
431        ))
432        .unwrap();
433        // add sequence
434        if !simplify {
435            rtn.write_fmt(format_args!(
436                "{}\t",
437                String::from_utf8_lossy(&self.record.seq().as_bytes()),
438            ))
439            .unwrap();
440        }
441        if quality {
442            // TODO add quality offset
443            rtn.write_fmt(format_args!(
444                "{}\t",
445                String::from_utf8_lossy(
446                    &self
447                        .record
448                        .qual()
449                        .iter()
450                        .map(|x| x + 33)
451                        .collect::<Vec<u8>>()
452                ),
453            ))
454            .unwrap();
455        }
456        // add PB features
457        let total_nuc_bp = nuc.lengths().iter().sum::<i64>();
458        let total_msp_bp = msp.lengths().iter().sum::<i64>();
459        rtn.write_fmt(format_args!(
460            "{}\t{}\t{}\t{}\t{}\t{}\t{}\t",
461            self.ec, rq, at_count, m6a_count, total_nuc_bp, total_msp_bp, cpg_count
462        ))
463        .unwrap();
464        // add fiber features. FIRE is its own coord group (parallel to
465        // nuc/m6a/cpg) — not wedged into the MSP columns.
466        let vecs = [
467            nuc.option_starts(),
468            nuc.option_lengths(),
469            nuc.reference_starts(),
470            nuc.reference_lengths(),
471            msp.option_starts(),
472            msp.option_lengths(),
473            msp.reference_starts(),
474            msp.reference_lengths(),
475            fire.option_starts(),
476            fire.option_lengths(),
477            fire_qual,
478            fire.reference_starts(),
479            fire.reference_lengths(),
480            m6a.option_starts(),
481            m6a.reference_starts(),
482            m6a_qual,
483            cpg.option_starts(),
484            cpg.reference_starts(),
485            cpg_qual,
486        ];
487        for vec in &vecs {
488            if vec.is_empty() {
489                rtn.push('.');
490                rtn.push('\t');
491            } else {
492                let z: String = vec
493                    .iter()
494                    .map(|x| match x {
495                        Some(y) => *y,
496                        None => -1,
497                    })
498                    .map(|x| x.to_string() + ",")
499                    .collect();
500                rtn.write_fmt(format_args!("{z}\t")).unwrap();
501            }
502        }
503        // replace the last tab with a newline
504        let len = rtn.len();
505        rtn.replace_range(len - 1..len, "\n");
506
507        rtn
508    }
509}
510
511pub struct FiberseqRecords<'a, R = bam::Reader>
512where
513    R: bam::Read,
514{
515    bam_chunk: BamChunk<'a, R>,
516    header: HeaderView,
517    filters: FiberFilters,
518    cur_chunk: Vec<FiberseqData>,
519}
520
521impl<'a> FiberseqRecords<'a, bam::Reader> {
522    pub fn new(bam: &'a mut bam::Reader, filters: FiberFilters) -> Self {
523        let header = bam.header().clone();
524        let bam_recs = bam.records();
525        let mut bam_chunk = BamChunk::new(bam_recs, None);
526        bam_chunk.set_bit_flag_filter(filters.get_bit_flag());
527        let cur_chunk: Vec<FiberseqData> = vec![];
528        FiberseqRecords {
529            bam_chunk,
530            header,
531            filters,
532            cur_chunk,
533        }
534    }
535}
536
537impl<'a> FiberseqRecords<'a, bam::IndexedReader> {
538    pub fn from_rec_iterator(
539        bam_recs: bam::Records<'a, bam::IndexedReader>,
540        header: HeaderView,
541        filters: FiberFilters,
542    ) -> Self {
543        let mut bam_chunk = BamChunk::new(bam_recs, None);
544        bam_chunk.set_bit_flag_filter(filters.get_bit_flag());
545        let cur_chunk: Vec<FiberseqData> = vec![];
546        FiberseqRecords {
547            bam_chunk,
548            header,
549            filters,
550            cur_chunk,
551        }
552    }
553}
554
555impl<R> Iterator for FiberseqRecords<'_, R>
556where
557    R: bam::Read,
558{
559    type Item = FiberseqData;
560
561    fn next(&mut self) -> Option<Self::Item> {
562        loop {
563            // if we are out of data check for another chunk in the bam
564            if self.cur_chunk.is_empty() {
565                match self.bam_chunk.next() {
566                    Some(recs) => {
567                        self.cur_chunk =
568                            FiberseqData::from_records(recs, &self.header, &self.filters);
569                        // we will be popping from this list so we want to remove the first element first, not the last
570                        self.cur_chunk.reverse();
571                    }
572                    None => return None,
573                }
574            }
575            let rec = self.cur_chunk.pop()?;
576            if self.filters.passes_fire_filter(&rec) {
577                return Some(rec);
578            }
579        }
580    }
581}