Skip to main content

competitive/tools/
iterable.rs

1#[macro_export]
2macro_rules! comprehension {
3    ($it:expr; @$type:ty) => {
4        $it.collect::<$type>()
5    };
6    ($it:expr) => {
7        comprehension![$it; @Vec<_>]
8    };
9    ($it:expr; @$type:ty; $p:pat => $e:expr) => {
10        comprehension![$it.map(|$p| $e); @$type]
11    };
12    ($it:expr; $p:pat => $($t:tt)*) => {
13        comprehension![$it; @Vec<_>; $p => $($t)*]
14    };
15    ($it:expr; $p:pat, $($t:tt)*) => {
16        comprehension![$it; @Vec<_>; $p, $($t)*]
17    };
18    ($it:expr; @$type:ty; $p:pat => $e:expr) => {
19        comprehension![$it; @$type; $p => $e]
20    };
21    ($it:expr; @$type:ty; $p:pat, $b:expr) => {
22        comprehension![$it.filter(|$p| $b); @$type]
23    };
24    ($it:expr; @$type:ty; $p:pat => $e:expr, $b:expr) => {
25        comprehension![$it.filter_map(|$p| if $b { Some($e) } else { None }); @$type]
26    };
27    ($it:expr; @$type:ty; $p:pat => $e:expr, $b1:expr, $b2:expr) => {
28        comprehension![$it; @$type; $p => $e, $b1 & $b2]
29    };
30    ($it:expr; @$type:ty; $p:pat => $e:expr, $b1:expr, $b2:expr, $($t:tt)*) => {
31        comprehension![$it; @$type; $p => $e, $b1 & $b2, $($t)*]
32    };
33}
34
35#[cfg(test)]
36mod tests {
37    use crate::tools::Xorshift;
38    #[test]
39    fn test_comprehension() {
40        use std::collections::{HashMap, HashSet};
41        let mut rng = Xorshift::default();
42        for _ in 0..1000 {
43            let n = rng.random(0..=100);
44            let values: Vec<_> = rng.random_iter(-100i32..=100).take(n).collect();
45            assert_eq!(
46                comprehension!(values.iter().copied(); @HashSet<_>),
47                (values.iter().copied()).collect::<HashSet<_>>()
48            );
49            assert_eq!(comprehension!(values.iter().copied()), values.clone());
50            assert_eq!(
51                comprehension!(values.iter().copied(); @HashMap<_,_>; i => (i, i + i)),
52                (values.iter().copied())
53                    .map(|i| (i, i + i))
54                    .collect::<HashMap<_, _>>()
55            );
56            assert_eq!(
57                comprehension!(values.iter().copied(); i => i + i),
58                (values.iter().copied()).map(|i| i + i).collect::<Vec<_>>()
59            );
60            assert_eq!(
61                comprehension!(values.iter().copied(); &i, i % 2 == 0),
62                (values.iter().copied())
63                    .filter(|&i| i % 2 == 0)
64                    .collect::<Vec<_>>()
65            );
66            assert_eq!(
67                comprehension!(values.iter().copied(); i => i + i, i % 2 == 0),
68                (values.iter().copied())
69                    .filter_map(|i| if i % 2 == 0 { Some(i + i) } else { None })
70                    .collect::<Vec<_>>()
71            );
72            assert_eq!(
73                comprehension!(values.iter().copied(); i => i + i, i % 2 == 0, i % 3 == 0),
74                (values.iter().copied())
75                    .filter_map(|i| if i % 2 == 0 && i % 3 == 0 {
76                        Some(i + i)
77                    } else {
78                        None
79                    })
80                    .collect::<Vec<_>>()
81            );
82            assert_eq!(
83                comprehension!(values.iter().copied(); i => i + i, i % 2 == 0, i % 3 == 0, i % 4 == 0),
84                (values.iter().copied())
85                    .filter_map(|i| if i % 2 == 0 && i % 3 == 0 && i % 4 == 0 {
86                        Some(i + i)
87                    } else {
88                        None
89                    })
90                    .collect::<Vec<_>>()
91            );
92            assert_eq!(
93                comprehension!(values.iter().copied(); @HashMap<_,_>; i => (i / 24, i), i % 2 == 0, i % 3 == 0, i % 4 == 0),
94                (values.iter().copied())
95                    .filter_map(|i| if i % 2 == 0 && i % 3 == 0 && i % 4 == 0 {
96                        Some((i / 24, i))
97                    } else {
98                        None
99                    })
100                    .collect::<HashMap<_, _>>()
101            );
102        }
103    }
104}