Skip to main content

competitive/tools/
iterator_ext.rs

1use std::iter::Peekable;
2
3pub trait IteratorExt: Iterator {
4    fn merge_by<I, F>(self, other: I, is_first: F) -> MergeBy<Self, I, F>
5    where
6        Self: Sized,
7        I: Iterator<Item = Self::Item>,
8        F: FnMut(&Self::Item, &Self::Item) -> bool,
9    {
10        MergeBy {
11            left: self.peekable(),
12            right: other.peekable(),
13            is_first,
14        }
15    }
16}
17
18impl<I> IteratorExt for I where I: Iterator {}
19
20pub struct MergeBy<I, J, F>
21where
22    I: Iterator,
23    J: Iterator<Item = I::Item>,
24{
25    left: Peekable<I>,
26    right: Peekable<J>,
27    is_first: F,
28}
29
30impl<I, J, F> Iterator for MergeBy<I, J, F>
31where
32    I: Iterator,
33    J: Iterator<Item = I::Item>,
34    F: FnMut(&I::Item, &I::Item) -> bool,
35{
36    type Item = I::Item;
37
38    fn next(&mut self) -> Option<Self::Item> {
39        match (self.left.peek(), self.right.peek()) {
40            (Some(l), Some(r)) => {
41                if (self.is_first)(l, r) {
42                    self.left.next()
43                } else {
44                    self.right.next()
45                }
46            }
47            (Some(_), None) => self.left.next(),
48            (None, Some(_)) => self.right.next(),
49            (None, None) => None,
50        }
51    }
52}
53
54#[cfg(test)]
55mod tests {
56    use super::*;
57    use crate::tools::Xorshift;
58
59    #[test]
60    fn test_merge_by() {
61        let mut rng = Xorshift::default();
62        for _ in 0..1000 {
63            let n = rng.random(0..=100);
64            let m = rng.random(0..=100);
65            let mut a: Vec<_> = rng.random_iter(-20..=20).take(n).collect();
66            let mut b: Vec<_> = rng.random_iter(-20..=20).take(m).collect();
67            let mut expected: Vec<_> = a.iter().chain(&b).copied().collect();
68            a.sort();
69            b.sort();
70            expected.sort();
71            assert_eq!(
72                a.iter()
73                    .merge_by(b.iter(), |x, y| x < y)
74                    .copied()
75                    .collect::<Vec<_>>(),
76                expected
77            );
78            expected.reverse();
79            assert_eq!(
80                a.iter()
81                    .rev()
82                    .merge_by(b.iter().rev(), |x, y| x > y)
83                    .copied()
84                    .collect::<Vec<_>>(),
85                expected
86            );
87        }
88    }
89}