Skip to main content

fibertools_rs/utils/
input_bam.rs

1use crate::cli;
2use crate::fiber::FiberseqRecords;
3use crate::utils::bio_io;
4use clap::{Args, ValueHint};
5use rust_htslib::bam;
6use rust_htslib::bam::Read;
7use std::fmt::Debug;
8
9pub static MIN_ML_SCORE: &str = "125";
10
11/// This struct establishes the way a Fiber-seq bam should be filtered when it is read in.
12/// This struct is both used as an argument for building FiberSeq records and as a struct that is parsed into cli arguments.
13#[derive(Debug, Args, Clone)]
14pub struct FiberFilters {
15    /// BAM bit flags to filter on, equivalent to `-F` in samtools view
16    /// Defaults to 0 (no filtering)
17    #[clap(
18        global = true,
19        short = 'F',
20        long = "filter",
21        help_heading = "BAM-Options"
22    )]
23    pub bit_flag: Option<u16>,
24    /// Filtering expression to use for filtering records
25    /// Example: filter to nucleosomes with lengths greater than 150 bp
26    ///   -x "len(nuc)>150"
27    /// Example: filter to msps with lengths between 30 and 49 bp
28    ///   -x "len(msp)=30:50"
29    /// Example: combine 2+ filter expressions
30    ///   -x "len(nuc)<150,len(msp)=30:50"
31    /// Filtering expressions support len() and qual() functions over msp, nuc, m6a, cpg
32    #[clap(
33        global = true,
34        short = 'x',
35        long = "ftx",
36        alias = "ft-expression",
37        help_heading = "BAM-Options"
38    )]
39    pub filter_expression: Option<String>,
40    /// Minium score in the ML tag to use or include in the output
41    #[clap(long="ml", alias="min-ml-score", default_value = MIN_ML_SCORE, help_heading = "BAM-Options", env="FT_MIN_ML_SCORE")]
42    pub min_ml_score: u8,
43    /// Output uncompressed BAM files
44    #[clap(help_heading = "BAM-Options", short, long)]
45    pub uncompressed: bool,
46    /// strip basemods in the first or last X bp of the read
47    #[clap(
48        global = true,
49        long,
50        default_value = "0",
51        help_heading = "BAM-Options",
52        hide = true
53    )]
54    pub strip_starting_basemods: i64,
55    /// Convenience: apply the FIRE peak-calling pipeline's fiber-level filters
56    /// (`--skip-no-m6a`, `--min-msp 10`, `--min-ave-msp-size 10`). Individual
57    /// filter flags still override when both are set. Requires MSP/m6A
58    /// annotations on the input BAM, so it is a no-op for commands that run
59    /// before those annotations exist.
60    #[clap(global = true, long, help_heading = "FIRE-Filter")]
61    pub fire_filter: bool,
62    /// Drop fibers with no m6A calls. Off by default;
63    /// `--fire-filter` turns this on unless explicitly set to `false`.
64    /// Use `--skip-no-m6a=false` to override when `--fire-filter` is set.
65    #[clap(
66        global = true,
67        long,
68        num_args = 0..=1,
69        default_missing_value = "true",
70        require_equals = true,
71        help_heading = "FIRE-Filter"
72    )]
73    pub skip_no_m6a: Option<bool>,
74    /// Drop fibers with fewer than `N` MSP calls.
75    /// Off (0) by default; `--fire-filter` sets this to 10 unless overridden.
76    #[clap(global = true, long, env = "MIN_MSP", help_heading = "FIRE-Filter")]
77    pub min_msp: Option<usize>,
78    /// Drop fibers whose average MSP size is below `N`.
79    /// Off (0) by default; `--fire-filter` sets this to 10 unless overridden.
80    #[clap(
81        global = true,
82        long,
83        env = "MIN_AVE_MSP_SIZE",
84        help_heading = "FIRE-Filter"
85    )]
86    pub min_ave_msp_size: Option<i64>,
87}
88
89impl std::default::Default for FiberFilters {
90    fn default() -> Self {
91        Self {
92            bit_flag: Some(0),
93            min_ml_score: MIN_ML_SCORE.parse().unwrap(),
94            filter_expression: None,
95            uncompressed: false,
96            strip_starting_basemods: 0,
97            fire_filter: false,
98            skip_no_m6a: None,
99            min_msp: None,
100            min_ave_msp_size: None,
101        }
102    }
103}
104
105impl FiberFilters {
106    /// Get the bit flag value, using a default if not explicitly set
107    pub fn get_bit_flag(&self) -> u16 {
108        self.bit_flag.unwrap_or(0)
109    }
110
111    /// Resolved `--skip-no-m6a`, using `--fire-filter` as the fallback.
112    pub fn resolved_skip_no_m6a(&self) -> bool {
113        self.skip_no_m6a.unwrap_or(self.fire_filter)
114    }
115
116    /// Resolved `--min-msp`, using `--fire-filter` (10) as the fallback.
117    pub fn resolved_min_msp(&self) -> usize {
118        self.min_msp
119            .unwrap_or(if self.fire_filter { 10 } else { 0 })
120    }
121
122    /// Resolved `--min-ave-msp-size`, using `--fire-filter` (10) as the fallback.
123    pub fn resolved_min_ave_msp_size(&self) -> i64 {
124        self.min_ave_msp_size
125            .unwrap_or(if self.fire_filter { 10 } else { 0 })
126    }
127
128    /// True if any FIRE fiber-level filter is active.
129    pub fn fire_filter_active(&self) -> bool {
130        self.resolved_skip_no_m6a()
131            || self.resolved_min_msp() > 0
132            || self.resolved_min_ave_msp_size() > 0
133    }
134
135    /// True if `rec` passes the FIRE fiber-level filters (skip_no_m6a,
136    /// min_msp, min_ave_msp_size). Called by `FiberseqRecords::next` so
137    /// every downstream consumer sees only filtered fibers.
138    ///
139    /// Edge cases worth preserving:
140    /// - Fibers with zero MSPs are always rejected once any filter is
141    ///   active. The check also guards the divide-by-zero in the average
142    ///   MSP size below, so don't drop it when refactoring.
143    /// - The no-m6a rejection is gated on `resolved_skip_no_m6a()` so
144    ///   `--skip-no-m6a=false` actually disables it (e.g. when combined
145    ///   with `--fire-filter`).
146    pub fn passes_fire_filter(&self, rec: &crate::fiber::FiberseqData) -> bool {
147        if !self.fire_filter_active() {
148            return true;
149        }
150        let msp = rec.msp();
151        let n_msps = msp.len();
152        if n_msps == 0 {
153            return false;
154        }
155        if self.resolved_skip_no_m6a() && rec.m6a().is_empty() {
156            return false;
157        }
158        if n_msps < self.resolved_min_msp() {
159            return false;
160        }
161        let ave_msp_size = msp.lengths().iter().sum::<i64>() / n_msps as i64;
162        if ave_msp_size < self.resolved_min_ave_msp_size() {
163            return false;
164        }
165        true
166    }
167
168    /// This function accepts an iterator over bam records and filters them based on the bit flag.
169    pub fn filter_on_bit_flags<'a, I>(
170        &'a self,
171        records: I,
172    ) -> impl Iterator<Item = bam::Record> + 'a
173    where
174        I: IntoIterator<Item = Result<bam::Record, rust_htslib::errors::Error>> + 'a,
175    {
176        let bit_flag = self.get_bit_flag();
177        records
178            .into_iter()
179            .map(|r| r.expect("htslib is unable to read a record in the input."))
180            .filter(move |r| {
181                // filter by bit flag
182                // `move` is needed to capture bit_flag value in the closure
183                (r.flags() & bit_flag) == 0
184            })
185    }
186}
187
188/// This struct is used to parse the input bam file and the filters that should be applied to the bam file.
189/// This struct is parsed to create command line arguments and then passed to many functions.
190#[derive(Debug, Args)]
191pub struct InputBam {
192    /// Input BAM file. If no path is provided stdin is used. For m6A prediction, this should be a HiFi bam file with kinetics data. For other commands, this should be a bam file with m6A calls.
193    #[clap(default_value = "-", value_hint = ValueHint::AnyPath)]
194    pub bam: String,
195    #[clap(flatten)]
196    pub filters: FiberFilters,
197    #[clap(flatten)]
198    pub global: cli::GlobalOpts,
199    /// by skipping this field it is not parsed as a command line argument
200    #[clap(skip)]
201    pub header: Option<bam::Header>,
202}
203
204impl InputBam {
205    pub fn bam_reader(&mut self) -> bam::Reader {
206        let mut bam = bio_io::bam_reader(&self.bam);
207        bam.set_threads(self.global.threads)
208            .expect("unable to set threads for bam reader");
209        self.header = Some(bam::Header::from_template(bam.header()));
210        bam
211    }
212
213    pub fn indexed_bam_reader(&mut self) -> bam::IndexedReader {
214        if &self.bam == "-" {
215            panic!("Cannot use stdin (\"-\") for indexed bam reading. Please provide a file path for the bam file.");
216        }
217
218        let mut bam =
219            bam::IndexedReader::from_path(&self.bam).expect("unable to open indexed bam file");
220        self.header = Some(bam::Header::from_template(bam.header()));
221        bam.set_threads(self.global.threads).unwrap();
222        bam
223    }
224
225    pub fn fibers<'a>(&self, bam: &'a mut bam::Reader) -> FiberseqRecords<'a> {
226        FiberseqRecords::new(bam, self.filters.clone())
227    }
228
229    pub fn header_view(&self) -> bam::HeaderView {
230        bam::HeaderView::from_header(self.header.as_ref().expect(
231            "Input bam must be opened before opening the header or creating a writer with the input bam as a template.",
232        ))
233    }
234
235    pub fn header(&self) -> bam::Header {
236        bam::Header::from_template(&self.header_view())
237    }
238
239    pub fn bam_writer(&self, out: &str) -> bam::Writer {
240        let header = self.header();
241        let program_name = "fibertools-rs";
242        let program_id = "ft";
243        let program_version = crate::VERSION;
244        let mut out = crate::utils::bio_io::program_bam_writer_from_header(
245            out,
246            header,
247            program_name,
248            program_id,
249            program_version,
250        );
251        out.set_threads(self.global.threads)
252            .expect("unable to set threads for bam writer");
253        if self.filters.uncompressed {
254            out.set_compression_level(bam::CompressionLevel::Uncompressed)
255                .expect("Unable to set compression level to uncompressed");
256        }
257        out
258    }
259
260    /// Add panSN-spec prefix to all contig names in the stored header
261    /// This must be called after opening the BAM reader to have an effect
262    pub fn add_pansn_prefix(&mut self, pansn_prefix: &str) {
263        if let Some(ref mut header) = self.header {
264            *header = crate::utils::panspec::add_pan_spec_header(header, pansn_prefix);
265        }
266    }
267
268    /// Strip panSN-spec information from all contig names in the stored header
269    /// using the provided delimiter character (e.g., '#' for panSN format)
270    /// This must be called after opening the BAM reader to have an effect
271    pub fn strip_pansn_spec(&mut self, delimiter: char) {
272        if let Some(ref mut header) = self.header {
273            *header = crate::utils::panspec::strip_pan_spec_header(header, &delimiter);
274        }
275    }
276
277    /// Fetch fibers from a specific region with filters applied
278    /// Returns an iterator of FiberseqData records from the specified region
279    ///
280    /// # Arguments
281    /// * `bam` - Mutable reference to an IndexedReader
282    /// * `chrom` - Chromosome/contig name
283    /// * `start` - Optional start position (0-based)
284    /// * `end` - Optional end position (0-based, exclusive)
285    /// ```
286    pub fn fetch_fibers<'a>(
287        &'a self,
288        bam: &'a mut bam::IndexedReader,
289        chrom: &str,
290        start: Option<i64>,
291        end: Option<i64>,
292    ) -> Result<FiberseqRecords<'a, bam::IndexedReader>, rust_htslib::errors::Error> {
293        // Fetch the region
294        match (start, end) {
295            (Some(s), Some(e)) => bam.fetch((chrom, s, e))?,
296            (None, None) => bam.fetch(chrom.as_bytes())?,
297            _ => panic!("Both start and end must be specified, or neither"),
298        }
299
300        // Create FiberseqRecords iterator from the fetched records
301        let records = bam.records();
302        let header = self.header_view();
303        let fiber_iter = FiberseqRecords::from_rec_iterator(records, header, self.filters.clone());
304
305        Ok(fiber_iter)
306    }
307}
308
309impl std::default::Default for InputBam {
310    fn default() -> Self {
311        Self {
312            bam: "-".to_string(),
313            filters: FiberFilters::default(),
314            global: cli::GlobalOpts::default(),
315            header: None,
316        }
317    }
318}
319
320#[cfg(test)]
321mod tests {
322    use super::*;
323    use crate::fiber::FiberseqData;
324    use crate::utils::basemods::M6A_TYPE;
325    use molecular_annotation::{Encoding, MolecularAnnotations, QualitySpec, Strand};
326
327    /// Build a minimal `FiberseqData` shaped only for `passes_fire_filter`.
328    /// `msp_lengths` populates the `msp` annotation type (each entry
329    /// contributes to count and average size). `m6a_count` controls whether
330    /// the `m6a` annotation type is empty or present (only emptiness
331    /// matters for the filter).
332    fn make_fsd(msp_lengths: &[i64], m6a_count: usize) -> FiberseqData {
333        let mut annotations = MolecularAnnotations::new(1000);
334        if !msp_lengths.is_empty() {
335            let t = annotations.add_annotation_type("msp", QualitySpec::none(), Encoding::Ma);
336            for &len in msp_lengths {
337                t.add(0, len as u32, Strand::Forward, vec![], None);
338            }
339        }
340        if m6a_count > 0 {
341            let t = annotations.add_annotation_type(
342                M6A_TYPE,
343                "Q".parse().expect("Q parses"),
344                Encoding::mm_ml(),
345            );
346            for i in 0..m6a_count {
347                t.add(i as u32, 1, Strand::Forward, vec![0], None);
348            }
349        }
350        FiberseqData {
351            record: rust_htslib::bam::Record::new(),
352            annotations,
353            ec: 0.0,
354            target_name: ".".to_string(),
355            rg: ".".to_string(),
356            center_position: None,
357        }
358    }
359
360    fn filters() -> FiberFilters {
361        FiberFilters::default()
362    }
363
364    #[test]
365    fn passes_when_no_filter_active() {
366        // Default config: nothing rejected, even fibers with no m6a / no msps.
367        let f = filters();
368        assert!(!f.fire_filter_active());
369        assert!(f.passes_fire_filter(&make_fsd(&[], 0)));
370        assert!(f.passes_fire_filter(&make_fsd(&[100, 100, 100], 5)));
371    }
372
373    #[test]
374    fn skip_no_m6a_rejects_only_no_m6a_fibers() {
375        let mut f = filters();
376        f.skip_no_m6a = Some(true);
377        assert!(f.fire_filter_active());
378        // No m6a → reject (and no-msp also rejected as a side effect).
379        assert!(!f.passes_fire_filter(&make_fsd(&[100], 0)));
380        // Has m6a, even with zero msps, still rejected by the empty-msp guard.
381        assert!(!f.passes_fire_filter(&make_fsd(&[], 3)));
382        // Has both → passes (other thresholds default to 0).
383        assert!(f.passes_fire_filter(&make_fsd(&[5], 1)));
384    }
385
386    #[test]
387    fn min_msp_rejects_fibers_with_too_few_msps() {
388        let mut f = filters();
389        f.min_msp = Some(3);
390        assert!(f.fire_filter_active());
391        assert!(!f.passes_fire_filter(&make_fsd(&[100, 100], 5)));
392        assert!(f.passes_fire_filter(&make_fsd(&[100, 100, 100], 5)));
393    }
394
395    #[test]
396    fn min_ave_msp_size_rejects_low_average() {
397        let mut f = filters();
398        f.min_ave_msp_size = Some(50);
399        // Average = 30 → reject.
400        assert!(!f.passes_fire_filter(&make_fsd(&[10, 20, 60], 5)));
401        // Average = 60 → pass.
402        assert!(f.passes_fire_filter(&make_fsd(&[40, 60, 80], 5)));
403    }
404
405    #[test]
406    fn fire_filter_combo_applies_all_three_defaults() {
407        // `--fire-filter` alone should imply skip_no_m6a + min_msp=10 + min_ave_msp_size=10.
408        let mut f = filters();
409        f.fire_filter = true;
410        assert!(f.resolved_skip_no_m6a());
411        assert_eq!(f.resolved_min_msp(), 10);
412        assert_eq!(f.resolved_min_ave_msp_size(), 10);
413        // <10 msps → reject.
414        let nine_long_msps: Vec<i64> = vec![100; 9];
415        assert!(!f.passes_fire_filter(&make_fsd(&nine_long_msps, 5)));
416        // 10 msps but ave_size = 5 < 10 → reject.
417        let ten_short_msps: Vec<i64> = vec![5; 10];
418        assert!(!f.passes_fire_filter(&make_fsd(&ten_short_msps, 5)));
419        // 10 msps, ave_size = 100, has m6a → pass.
420        let ten_long_msps: Vec<i64> = vec![100; 10];
421        assert!(f.passes_fire_filter(&make_fsd(&ten_long_msps, 5)));
422        // Same fiber but no m6a → reject (skip_no_m6a is implied).
423        assert!(!f.passes_fire_filter(&make_fsd(&ten_long_msps, 0)));
424    }
425
426    #[test]
427    fn explicit_flag_overrides_fire_filter_default() {
428        // `--fire-filter --min-msp=5` should use 5, not 10.
429        let mut f = filters();
430        f.fire_filter = true;
431        f.min_msp = Some(5);
432        assert_eq!(f.resolved_min_msp(), 5);
433        // Other two still take fire-filter defaults.
434        assert!(f.resolved_skip_no_m6a());
435        assert_eq!(f.resolved_min_ave_msp_size(), 10);
436        // 5 msps would have failed under the 10 default, but passes under 5.
437        let five_long_msps: Vec<i64> = vec![100; 5];
438        assert!(f.passes_fire_filter(&make_fsd(&five_long_msps, 5)));
439    }
440
441    #[test]
442    fn explicit_skip_no_m6a_false_overrides_fire_filter_default() {
443        // `--fire-filter --skip-no-m6a=false` keeps the size/count thresholds
444        // but should turn the no-m6a guard off. The empty-msp guard still
445        // applies (it's a divide-by-zero protection), so we use a record
446        // with msps but no m6a.
447        let mut f = filters();
448        f.fire_filter = true;
449        f.skip_no_m6a = Some(false);
450        assert!(!f.resolved_skip_no_m6a());
451        let ten_long_msps: Vec<i64> = vec![100; 10];
452        assert!(f.passes_fire_filter(&make_fsd(&ten_long_msps, 0)));
453    }
454}