Skip to main content

competitive/tools/
array.rs

1#[macro_export]
2macro_rules! array {
3    [@inner $data:ident = [$init:expr; $len:expr]] => {{
4        use ::std::mem::{ManuallyDrop, MaybeUninit};
5        let mut $data: [MaybeUninit<_>; $len] = unsafe { MaybeUninit::uninit().assume_init() };
6        $init;
7        #[repr(C)]
8        union __Transmuter<const N: usize, T: Clone> {
9            src: ManuallyDrop<[MaybeUninit<T>; N]>,
10            dst: ManuallyDrop<[T; N]>,
11        }
12        ManuallyDrop::into_inner(unsafe { __Transmuter { src: ManuallyDrop::new($data) }.dst })
13    }};
14    [|| $e:expr; $len:expr] => {
15        $crate::array![@inner data = [data.iter_mut().for_each(|item| *item = MaybeUninit::new($e)); $len]]
16    };
17    [|$i:pat_param| $e:expr; $len:expr] => {
18        $crate::array![@inner data = [data.iter_mut().enumerate().for_each(|($i, item)| *item = MaybeUninit::new($e)); $len]]
19    };
20    [$e:expr; $len:expr] => {{
21        let e = $e;
22        $crate::array![|| Clone::clone(&e); $len]
23    }};
24}
25
26#[test]
27fn test_array() {
28    use crate::tools::Xorshift;
29    use std::array;
30    fn check<const N: usize>(start: i32, step: i32) {
31        let mut x = start;
32        assert_eq!(array![start; N], [start; N]);
33        assert_eq!(
34            array![|| { x += step; x }; N],
35            array::from_fn(|i| start + (i as i32 + 1) * step)
36        );
37        assert_eq!(x, start + N as i32 * step);
38        assert_eq!(
39            array![|i| start + i as i32 * step; N],
40            array::from_fn(|i| start + i as i32 * step)
41        );
42    }
43    let mut rng = Xorshift::default();
44    for (start, step) in (-5..=5)
45        .flat_map(|start| (-5..=5).map(move |step| (start, step)))
46        .chain(rng.random_iter((-1000..=1000, -1000..=1000)).take(1000))
47    {
48        check::<0>(start, step);
49        check::<1>(start, step);
50        check::<2>(start, step);
51        check::<3>(start, step);
52        check::<4>(start, step);
53        check::<8>(start, step);
54        check::<16>(start, step);
55        check::<32>(start, step);
56    }
57}