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    /// Detect the sequencing platform (PacBio vs ONT) for this read from its
187    /// aux tags, falling back to the read name if those were stripped.
188    pub fn platform(&self) -> crate::utils::platform::SeqPlatform {
189        crate::utils::platform::platform_from_record(&self.record)
190    }
191
192    pub fn get_hp(&self) -> String {
193        match self.record.aux(b"HP") {
194            Ok(Aux::U8(v)) => format!("H{v}"),
195            Ok(Aux::I8(v)) => format!("H{v}"),
196            Ok(Aux::U16(v)) => format!("H{v}"),
197            Ok(Aux::I16(v)) => format!("H{v}"),
198            Ok(Aux::U32(v)) => format!("H{v}"),
199            Ok(Aux::I32(v)) => format!("H{v}"),
200            _ => "UNK".to_string(),
201        }
202    }
203
204    //
205    //  WRITE BED12 FUNCTIONS
206    //
207    pub fn write_msp(&self, reference: bool) -> String {
208        let msp = self.msp();
209        let (starts, _ends, lengths) = if reference {
210            (
211                msp.reference_starts(),
212                msp.reference_ends(),
213                msp.reference_lengths(),
214            )
215        } else {
216            (msp.option_starts(), msp.option_ends(), msp.option_lengths())
217        };
218        self.to_bed12(reference, &starts, &lengths, LINKER_COLOR)
219    }
220
221    pub fn write_nuc(&self, reference: bool) -> String {
222        let nuc = self.nuc();
223        let (starts, _ends, lengths) = if reference {
224            (
225                nuc.reference_starts(),
226                nuc.reference_ends(),
227                nuc.reference_lengths(),
228            )
229        } else {
230            (nuc.option_starts(), nuc.option_ends(), nuc.option_lengths())
231        };
232        self.to_bed12(reference, &starts, &lengths, NUC_COLOR)
233    }
234
235    pub fn write_m6a(&self, reference: bool) -> String {
236        let m6a = self.m6a();
237        let starts = if reference {
238            m6a.reference_starts()
239        } else {
240            m6a.option_starts()
241        };
242        let lengths = vec![Some(1); starts.len()];
243        self.to_bed12(reference, &starts, &lengths, M6A_COLOR)
244    }
245
246    pub fn write_cpg(&self, reference: bool) -> String {
247        let cpg = self.cpg();
248        let starts = if reference {
249            cpg.reference_starts()
250        } else {
251            cpg.option_starts()
252        };
253        let lengths = vec![Some(1); starts.len()];
254        self.to_bed12(reference, &starts, &lengths, CPG_COLOR)
255    }
256
257    pub fn to_bed12(
258        &self,
259        reference: bool,
260        starts: &[Option<i64>],
261        lengths: &[Option<i64>],
262        color: &str,
263    ) -> String {
264        if starts.is_empty() {
265            return "".to_string();
266        }
267        // skip if no alignments are here
268        if self.record.is_unmapped() && reference {
269            return "".to_string();
270        }
271
272        let ct;
273        let start;
274        let end;
275        let name = String::from_utf8_lossy(self.record.qname()).to_string();
276        let mut rtn: String = String::with_capacity(0);
277        if reference {
278            ct = &self.target_name;
279            start = self.record.reference_start();
280            end = self.record.reference_end();
281        } else {
282            ct = &name;
283            start = 0;
284            end = self.record.seq_len() as i64;
285        }
286        let score = self.ec.round() as i64;
287        let strand = if self.record.is_reverse() { '-' } else { '+' };
288        // filter out positions that do not have an exact liftover
289        let (filtered_starts, filtered_lengths): (Vec<i64>, Vec<i64>) = starts
290            .iter()
291            .flatten()
292            .zip(lengths.iter().flatten())
293            .unzip();
294        // skip empty ones
295        if filtered_lengths.is_empty() || filtered_starts.is_empty() {
296            return "".to_string();
297        }
298        let b_ct = filtered_starts.len() + 2;
299        let b_ln: String = filtered_lengths
300            .iter()
301            .map(|&ln| ln.to_string() + ",")
302            .collect();
303        let b_st: String = filtered_starts
304            .iter()
305            .map(|&st| (st - start).to_string() + ",")
306            .collect();
307        assert_eq!(filtered_lengths.len(), filtered_starts.len());
308
309        rtn.push_str(ct);
310        rtn.push('\t');
311        rtn.push_str(&start.to_string());
312        rtn.push('\t');
313        rtn.push_str(&end.to_string());
314        rtn.push('\t');
315        rtn.push_str(&name);
316        rtn.push('\t');
317        rtn.push_str(&score.to_string());
318        rtn.push('\t');
319        rtn.push(strand);
320        rtn.push('\t');
321        rtn.push_str(&start.to_string());
322        rtn.push('\t');
323        rtn.push_str(&end.to_string());
324        rtn.push('\t');
325        rtn.push_str(color);
326        rtn.push('\t');
327        rtn.push_str(&b_ct.to_string());
328        rtn.push_str("\t0,"); // add a zero length start
329        rtn.push_str(&b_ln);
330        rtn.push_str("1\t0,"); // add a 1 base length and a 0 start point
331        rtn.push_str(&b_st);
332        write!(&mut rtn, "{}", format_args!("{}\n", end - start - 1)).unwrap();
333        rtn
334    }
335
336    //
337    // WRITE ALL FUNCTIONS
338    //
339
340    pub fn all_header(simplify: bool, quality: bool) -> String {
341        let mut x = format!(
342            "#{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t",
343            "ct", "st", "en", "fiber", "score", "strand", "sam_flag", "HP", "RG", "fiber_length",
344        );
345        if !simplify {
346            x.push_str("fiber_sequence\t")
347        }
348        if quality {
349            x.push_str("fiber_qual\t")
350        }
351        x.push_str(&format!(
352            "{}\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",
353            "ec",
354            "rq",
355            "total_AT_bp",
356            "total_m6a_bp",
357            "total_nuc_bp",
358            "total_msp_bp",
359            "total_5mC_bp",
360            "nuc_starts",
361            "nuc_lengths",
362            "ref_nuc_starts",
363            "ref_nuc_lengths",
364            "msp_starts",
365            "msp_lengths",
366            "ref_msp_starts",
367            "ref_msp_lengths",
368            "fire_starts",
369            "fire_lengths",
370            "fire_qual",
371            "ref_fire_starts",
372            "ref_fire_lengths",
373            "m6a",
374            "ref_m6a",
375            "m6a_qual",
376            "5mC",
377            "ref_5mC",
378            "5mC_qual"
379        ));
380        x
381    }
382
383    pub fn write_all(&self, simplify: bool, quality: bool) -> String {
384        // PB features
385        let name = std::str::from_utf8(self.record.qname()).unwrap();
386        let score = self.ec.round() as i64;
387        let q_len = self.record.seq_len() as i64;
388        let rq = match self.get_rq() {
389            Some(x) => format!("{x}"),
390            None => ".".to_string(),
391        };
392        // reference features
393        let ct;
394        let start;
395        let end;
396        let strand;
397        if self.record.is_unmapped() {
398            ct = ".";
399            start = 0;
400            end = 0;
401            strand = '.';
402        } else {
403            ct = &self.target_name;
404            start = self.record.reference_start();
405            end = self.record.reference_end();
406            strand = if self.record.is_reverse() { '-' } else { '+' };
407        }
408        let sam_flag = self.record.flags();
409        let hp = self.get_hp();
410
411        let at_count = self
412            .record
413            .seq()
414            .as_bytes()
415            .iter()
416            .filter(|&x| *x == b'A' || *x == b'T')
417            .count() as i64;
418
419        // get the info
420        let m6a = self.m6a();
421        let cpg = self.cpg();
422        let msp = self.msp();
423        let nuc = self.nuc();
424        let fire = self.fire();
425        let m6a_count = m6a.len();
426        let m6a_qual = m6a.qual().iter().map(|a| Some(*a as i64)).collect();
427        let cpg_count = cpg.len();
428        let cpg_qual = cpg.qual().iter().map(|a| Some(*a as i64)).collect();
429        let fire_qual: Vec<Option<i64>> = fire.qual().iter().map(|a| Some(*a as i64)).collect();
430
431        // write the features
432        let mut rtn = String::with_capacity(0);
433        // add first things 7
434        rtn.write_fmt(format_args!(
435            "{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t{}\t",
436            ct, start, end, name, score, strand, sam_flag, hp, self.rg, q_len
437        ))
438        .unwrap();
439        // add sequence
440        if !simplify {
441            rtn.write_fmt(format_args!(
442                "{}\t",
443                String::from_utf8_lossy(&self.record.seq().as_bytes()),
444            ))
445            .unwrap();
446        }
447        if quality {
448            // TODO add quality offset
449            rtn.write_fmt(format_args!(
450                "{}\t",
451                String::from_utf8_lossy(
452                    &self
453                        .record
454                        .qual()
455                        .iter()
456                        .map(|x| x + 33)
457                        .collect::<Vec<u8>>()
458                ),
459            ))
460            .unwrap();
461        }
462        // add PB features
463        let total_nuc_bp = nuc.lengths().iter().sum::<i64>();
464        let total_msp_bp = msp.lengths().iter().sum::<i64>();
465        rtn.write_fmt(format_args!(
466            "{}\t{}\t{}\t{}\t{}\t{}\t{}\t",
467            self.ec, rq, at_count, m6a_count, total_nuc_bp, total_msp_bp, cpg_count
468        ))
469        .unwrap();
470        // add fiber features. FIRE is its own coord group (parallel to
471        // nuc/m6a/cpg) — not wedged into the MSP columns.
472        let vecs = [
473            nuc.option_starts(),
474            nuc.option_lengths(),
475            nuc.reference_starts(),
476            nuc.reference_lengths(),
477            msp.option_starts(),
478            msp.option_lengths(),
479            msp.reference_starts(),
480            msp.reference_lengths(),
481            fire.option_starts(),
482            fire.option_lengths(),
483            fire_qual,
484            fire.reference_starts(),
485            fire.reference_lengths(),
486            m6a.option_starts(),
487            m6a.reference_starts(),
488            m6a_qual,
489            cpg.option_starts(),
490            cpg.reference_starts(),
491            cpg_qual,
492        ];
493        for vec in &vecs {
494            if vec.is_empty() {
495                rtn.push('.');
496                rtn.push('\t');
497            } else {
498                let z: String = vec
499                    .iter()
500                    .map(|x| match x {
501                        Some(y) => *y,
502                        None => -1,
503                    })
504                    .map(|x| x.to_string() + ",")
505                    .collect();
506                rtn.write_fmt(format_args!("{z}\t")).unwrap();
507            }
508        }
509        // replace the last tab with a newline
510        let len = rtn.len();
511        rtn.replace_range(len - 1..len, "\n");
512
513        rtn
514    }
515}
516
517pub struct FiberseqRecords<'a, R = bam::Reader>
518where
519    R: bam::Read,
520{
521    bam_chunk: BamChunk<'a, R>,
522    header: HeaderView,
523    filters: FiberFilters,
524    cur_chunk: Vec<FiberseqData>,
525}
526
527impl<'a> FiberseqRecords<'a, bam::Reader> {
528    pub fn new(bam: &'a mut bam::Reader, filters: FiberFilters) -> Self {
529        let header = bam.header().clone();
530        let bam_recs = bam.records();
531        let mut bam_chunk = BamChunk::new(bam_recs, None);
532        bam_chunk.set_bit_flag_filter(filters.get_bit_flag());
533        let cur_chunk: Vec<FiberseqData> = vec![];
534        FiberseqRecords {
535            bam_chunk,
536            header,
537            filters,
538            cur_chunk,
539        }
540    }
541}
542
543impl<'a> FiberseqRecords<'a, bam::IndexedReader> {
544    pub fn from_rec_iterator(
545        bam_recs: bam::Records<'a, bam::IndexedReader>,
546        header: HeaderView,
547        filters: FiberFilters,
548    ) -> Self {
549        let mut bam_chunk = BamChunk::new(bam_recs, None);
550        bam_chunk.set_bit_flag_filter(filters.get_bit_flag());
551        let cur_chunk: Vec<FiberseqData> = vec![];
552        FiberseqRecords {
553            bam_chunk,
554            header,
555            filters,
556            cur_chunk,
557        }
558    }
559}
560
561impl<R> Iterator for FiberseqRecords<'_, R>
562where
563    R: bam::Read,
564{
565    type Item = FiberseqData;
566
567    fn next(&mut self) -> Option<Self::Item> {
568        loop {
569            // if we are out of data check for another chunk in the bam
570            if self.cur_chunk.is_empty() {
571                match self.bam_chunk.next() {
572                    Some(recs) => {
573                        self.cur_chunk =
574                            FiberseqData::from_records(recs, &self.header, &self.filters);
575                        // we will be popping from this list so we want to remove the first element first, not the last
576                        self.cur_chunk.reverse();
577                    }
578                    None => return None,
579                }
580            }
581            let rec = self.cur_chunk.pop()?;
582            if self.filters.passes_fire_filter(&rec) {
583                return Some(rec);
584            }
585        }
586    }
587}