1use std::collections::HashMap;
7
8use gtars_core::models::{Region, RegionSet};
9
10use crate::models::{Strand, StrandedRegionSet};
11
12impl StrandedRegionSet {
15 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 } else {
34 regions.push(r.clone());
36 strands.push(*s);
37 }
38 }
39 StrandedRegionSet {
40 inner: RegionSet::from(regions),
41 strands,
42 }
43 }
44
45 pub fn promoters(&self, upstream: u32, downstream: u32) -> RegionSet {
50 self.promoters_stranded(upstream, downstream).inner
51 }
52
53 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 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 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 pub fn setdiff(&self, other: &StrandedRegionSet) -> StrandedRegionSet {
139 let a = self.reduce();
140 let b = other.reduce();
141
142 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 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 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
220fn 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 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 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 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 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 let srs = make_stranded(vec![("chr1", 3000, 8000)], vec![Strand::Minus]);
289 let result = srs.promoters(200, 50);
290 assert_eq!(result.regions[0], make_region("chr1", 7950, 8200));
292 }
293
294 #[rstest]
295 fn test_stranded_promoters_saturating_at_zero() {
296 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); assert_eq!(result.regions[0].end, 50);
301 }
302
303 #[rstest]
306 fn test_stranded_reduce_keeps_opposite_strand_separate() {
307 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 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 #[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 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 }