Skip to main content

gtars_genomicdist/
stranded_region_set.rs

1//! Strand-aware set operations on [`StrandedRegionSet`].
2//!
3//! These operate on a `StrandedRegionSet` (regions paired with per-region
4//! strands) and match strands when combining regions.
5
6use std::collections::HashMap;
7
8use gtars_core::models::{Region, RegionSet};
9
10use crate::models::{Strand, StrandedRegionSet};
11
12// ── Strand-aware operations on StrandedRegionSet ─────────────────────────
13
14impl StrandedRegionSet {
15    /// Clip regions to chromosome boundaries, preserving strand information.
16    pub fn trim(&self, chrom_sizes: &HashMap<String, u32>) -> StrandedRegionSet {
17        let mut regions = Vec::new();
18        let mut strands = Vec::new();
19        for (r, s) in self.inner.regions.iter().zip(self.strands.iter()) {
20            if let Some(&chrom_size) = chrom_sizes.get(&r.chr) {
21                let start = r.start.min(chrom_size);
22                let end = r.end.min(chrom_size);
23                if start < end {
24                    regions.push(Region {
25                        chr: r.chr.clone(),
26                        start,
27                        end,
28                        rest: None,
29                    });
30                    strands.push(*s);
31                }
32                // Drop zero-width regions (start == end after clipping)
33            } else {
34                // Chromosome not in chrom_sizes — keep region as-is (no trimming)
35                regions.push(r.clone());
36                strands.push(*s);
37            }
38        }
39        StrandedRegionSet {
40            inner: RegionSet::from(regions),
41            strands,
42        }
43    }
44
45    /// Strand-aware promoter computation. Returns an unstranded `RegionSet`.
46    ///
47    /// - Plus / Unstranded: `[start - upstream, start + downstream)`
48    /// - Minus: `[end - downstream, end + upstream)`
49    pub fn promoters(&self, upstream: u32, downstream: u32) -> RegionSet {
50        self.promoters_stranded(upstream, downstream).inner
51    }
52
53    /// Like `promoters()` but preserves strand information, so a subsequent
54    /// strand-aware `reduce()` merges only same-strand promoters.
55    pub fn promoters_stranded(&self, upstream: u32, downstream: u32) -> StrandedRegionSet {
56        let regions: Vec<Region> = self
57            .inner
58            .regions
59            .iter()
60            .zip(self.strands.iter())
61            .map(|(r, strand)| match strand {
62                Strand::Minus => Region {
63                    chr: r.chr.clone(),
64                    start: r.end.saturating_sub(downstream),
65                    end: r.end.saturating_add(upstream),
66                    rest: None,
67                },
68                _ => Region {
69                    chr: r.chr.clone(),
70                    start: r.start.saturating_sub(upstream),
71                    end: r.start.saturating_add(downstream),
72                    rest: None,
73                },
74            })
75            .collect();
76        StrandedRegionSet {
77            inner: RegionSet::from(regions),
78            strands: self.strands.clone(),
79        }
80    }
81
82    /// Strand-aware reduce: merge overlapping/adjacent intervals only within
83    /// the same (chr, strand) group.
84    pub fn reduce(&self) -> StrandedRegionSet {
85        if self.inner.regions.is_empty() {
86            return StrandedRegionSet::new(
87                RegionSet::from(Vec::<Region>::new()),
88                Vec::new(),
89            );
90        }
91
92        // Build (region, strand) pairs and sort by (chr, strand, start)
93        let mut pairs: Vec<(&Region, &Strand)> = self
94            .inner
95            .regions
96            .iter()
97            .zip(self.strands.iter())
98            .collect();
99        pairs.sort_by(|(a, sa), (b, sb)| {
100            a.chr
101                .cmp(&b.chr)
102                .then(strand_ord(**sa).cmp(&strand_ord(**sb)))
103                .then(a.start.cmp(&b.start))
104        });
105
106        let mut merged_regions: Vec<Region> = Vec::new();
107        let mut merged_strands: Vec<Strand> = Vec::new();
108
109        let (mut cur_r, mut cur_s) = (pairs[0].0.clone(), *pairs[0].1);
110        for &(r, s) in &pairs[1..] {
111            if r.chr == cur_r.chr && *s == cur_s && r.start <= cur_r.end {
112                cur_r.end = cur_r.end.max(r.end);
113            } else {
114                merged_regions.push(Region {
115                    chr: cur_r.chr.clone(),
116                    start: cur_r.start,
117                    end: cur_r.end,
118                    rest: None,
119                });
120                merged_strands.push(cur_s);
121                cur_r = r.clone();
122                cur_s = *s;
123            }
124        }
125        merged_regions.push(Region {
126            chr: cur_r.chr,
127            start: cur_r.start,
128            end: cur_r.end,
129            rest: None,
130        });
131        merged_strands.push(cur_s);
132
133        StrandedRegionSet::new(RegionSet::from(merged_regions), merged_strands)
134    }
135
136    /// Strand-aware setdiff: subtract `other` from `self`, matching only
137    /// within the same (chr, strand) group.
138    pub fn setdiff(&self, other: &StrandedRegionSet) -> StrandedRegionSet {
139        let a = self.reduce();
140        let b = other.reduce();
141
142        // Group b by (chr, strand)
143        let mut b_map: HashMap<(String, Strand), Vec<&Region>> = HashMap::new();
144        for (r, s) in b.inner.regions.iter().zip(b.strands.iter()) {
145            b_map
146                .entry((r.chr.clone(), *s))
147                .or_default()
148                .push(r);
149        }
150
151        let mut result_regions: Vec<Region> = Vec::new();
152        let mut result_strands: Vec<Strand> = Vec::new();
153
154        // Process a by (chr, strand) groups
155        let mut i = 0;
156        while i < a.inner.regions.len() {
157            let chr = &a.inner.regions[i].chr;
158            let strand = a.strands[i];
159
160            // Find the extent of this (chr, strand) group
161            let mut j = i;
162            while j < a.inner.regions.len()
163                && a.inner.regions[j].chr == *chr
164                && a.strands[j] == strand
165            {
166                j += 1;
167            }
168
169            let empty_vec = vec![];
170            let b_chr_strand = b_map
171                .get(&(chr.clone(), strand))
172                .unwrap_or(&empty_vec);
173            let mut b_idx = 0;
174
175            for a_region in &a.inner.regions[i..j] {
176                while b_idx < b_chr_strand.len()
177                    && b_chr_strand[b_idx].end <= a_region.start
178                {
179                    b_idx += 1;
180                }
181
182                let mut pos = a_region.start;
183                let mut k = b_idx;
184
185                while k < b_chr_strand.len()
186                    && b_chr_strand[k].start < a_region.end
187                    && pos < a_region.end
188                {
189                    if b_chr_strand[k].start > pos {
190                        result_regions.push(Region {
191                            chr: chr.clone(),
192                            start: pos,
193                            end: b_chr_strand[k].start,
194                            rest: None,
195                        });
196                        result_strands.push(strand);
197                    }
198                    pos = pos.max(b_chr_strand[k].end);
199                    k += 1;
200                }
201
202                if pos < a_region.end {
203                    result_regions.push(Region {
204                        chr: chr.clone(),
205                        start: pos,
206                        end: a_region.end,
207                        rest: None,
208                    });
209                    result_strands.push(strand);
210                }
211            }
212
213            i = j;
214        }
215
216        StrandedRegionSet::new(RegionSet::from(result_regions), result_strands)
217    }
218}
219
220/// Ordering key for Strand so sorting groups by (chr, strand, start).
221fn strand_ord(s: Strand) -> u8 {
222    match s {
223        Strand::Plus => 0,
224        Strand::Minus => 1,
225        Strand::Unstranded => 2,
226    }
227}
228
229#[cfg(test)]
230mod tests {
231    use super::*;
232    use pretty_assertions::assert_eq;
233    use rstest::*;
234
235    fn make_region(chr: &str, start: u32, end: u32) -> Region {
236        Region {
237            chr: chr.to_string(),
238            start,
239            end,
240            rest: None,
241        }
242    }
243
244    fn make_regionset(regions: Vec<(&str, u32, u32)>) -> RegionSet {
245        let regions: Vec<Region> = regions
246            .into_iter()
247            .map(|(chr, start, end)| make_region(chr, start, end))
248            .collect();
249        RegionSet::from(regions)
250    }
251
252    // ── trim tests ──────────────────────────────────────────────────────
253
254
255    fn make_stranded(regions: Vec<(&str, u32, u32)>, strands: Vec<Strand>) -> StrandedRegionSet {
256        StrandedRegionSet::new(make_regionset(regions), strands)
257    }
258
259
260    fn test_stranded_promoters_plus() {
261        // Plus-strand gene at [1000, 5000): promoter 100bp upstream of start
262        let srs = make_stranded(vec![("chr1", 1000, 5000)], vec![Strand::Plus]);
263        let result = srs.promoters(100, 0);
264        assert_eq!(result.regions.len(), 1);
265        assert_eq!(result.regions[0], make_region("chr1", 900, 1000));
266    }
267
268    #[rstest]
269    fn test_stranded_promoters_minus() {
270        // Minus-strand gene at [3000, 8000): promoter 100bp upstream of end
271        let srs = make_stranded(vec![("chr2", 3000, 8000)], vec![Strand::Minus]);
272        let result = srs.promoters(100, 0);
273        assert_eq!(result.regions.len(), 1);
274        assert_eq!(result.regions[0], make_region("chr2", 8000, 8100));
275    }
276
277    #[rstest]
278    fn test_stranded_promoters_unstranded_matches_plus() {
279        // Unstranded should behave like plus-strand (same as original behavior)
280        let srs = make_stranded(vec![("chr1", 1000, 5000)], vec![Strand::Unstranded]);
281        let result = srs.promoters(100, 0);
282        assert_eq!(result.regions[0], make_region("chr1", 900, 1000));
283    }
284
285    #[rstest]
286    fn test_stranded_promoters_with_downstream() {
287        // Minus-strand with both upstream and downstream
288        let srs = make_stranded(vec![("chr1", 3000, 8000)], vec![Strand::Minus]);
289        let result = srs.promoters(200, 50);
290        // [end - downstream, end + upstream) = [7950, 8200)
291        assert_eq!(result.regions[0], make_region("chr1", 7950, 8200));
292    }
293
294    #[rstest]
295    fn test_stranded_promoters_saturating_at_zero() {
296        // Plus-strand gene near origin
297        let srs = make_stranded(vec![("chr1", 50, 500)], vec![Strand::Plus]);
298        let result = srs.promoters(200, 0);
299        assert_eq!(result.regions[0].start, 0); // saturates
300        assert_eq!(result.regions[0].end, 50);
301    }
302
303    // ── StrandedRegionSet reduce tests ───────────────────────────────
304
305    #[rstest]
306    fn test_stranded_reduce_keeps_opposite_strand_separate() {
307        // Two overlapping regions on opposite strands should NOT merge
308        let srs = make_stranded(
309            vec![("chr1", 100, 300), ("chr1", 200, 400)],
310            vec![Strand::Plus, Strand::Minus],
311        );
312        let reduced = srs.reduce();
313        assert_eq!(reduced.len(), 2);
314    }
315
316    #[rstest]
317    fn test_stranded_reduce_merges_same_strand() {
318        // Two overlapping regions on the same strand should merge
319        let srs = make_stranded(
320            vec![("chr1", 100, 300), ("chr1", 200, 400)],
321            vec![Strand::Plus, Strand::Plus],
322        );
323        let reduced = srs.reduce();
324        assert_eq!(reduced.len(), 1);
325        assert_eq!(reduced.inner.regions[0], make_region("chr1", 100, 400));
326    }
327
328    #[rstest]
329    fn test_stranded_reduce_empty() {
330        let srs = StrandedRegionSet::new(
331            RegionSet::from(Vec::<Region>::new()),
332            Vec::new(),
333        );
334        let reduced = srs.reduce();
335        assert_eq!(reduced.len(), 0);
336    }
337
338    // ── StrandedRegionSet setdiff tests ──────────────────────────────
339
340    #[rstest]
341    fn test_stranded_setdiff_same_strand_subtracts() {
342        let a = make_stranded(vec![("chr1", 0, 100)], vec![Strand::Plus]);
343        let b = make_stranded(vec![("chr1", 30, 70)], vec![Strand::Plus]);
344        let result = a.setdiff(&b);
345        assert_eq!(result.len(), 2);
346        assert_eq!(result.inner.regions[0], make_region("chr1", 0, 30));
347        assert_eq!(result.inner.regions[1], make_region("chr1", 70, 100));
348    }
349
350    #[rstest]
351    fn test_stranded_setdiff_different_strand_no_subtraction() {
352        // Opposite strand: should NOT subtract
353        let a = make_stranded(vec![("chr1", 0, 100)], vec![Strand::Plus]);
354        let b = make_stranded(vec![("chr1", 30, 70)], vec![Strand::Minus]);
355        let result = a.setdiff(&b);
356        assert_eq!(result.len(), 1);
357        assert_eq!(result.inner.regions[0], make_region("chr1", 0, 100));
358    }
359
360    // ── concat tests ────────────────────────────────────────────────────
361
362}