Skip to main content

competitive/math/
array_vec.rs

1use std::ops::{
2    Add, AddAssign, BitAnd, BitAndAssign, BitOr, BitOrAssign, BitXor, BitXorAssign, Div, DivAssign,
3    Index, IndexMut, Mul, MulAssign, Neg, Not, Rem, RemAssign, Shl, ShlAssign, Shr, ShrAssign, Sub,
4    SubAssign,
5};
6
7#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, PartialOrd, Ord, Hash)]
8pub struct ArrayVecScalar<T>(pub T);
9
10impl<T> From<T> for ArrayVecScalar<T> {
11    fn from(value: T) -> Self {
12        Self(value)
13    }
14}
15
16pub trait ToArrayVecScalar: Sized {
17    fn to_array_vec_scalar(self) -> ArrayVecScalar<Self>;
18}
19
20impl<T> ToArrayVecScalar for T {
21    fn to_array_vec_scalar(self) -> ArrayVecScalar<Self> {
22        ArrayVecScalar(self)
23    }
24}
25
26#[derive(Clone, Copy, Debug, PartialEq, Eq, PartialOrd, Ord, Hash)]
27pub struct ArrayVec<T, const N: usize>(pub [T; N]);
28
29pub trait ToArrayVec<T, const N: usize>: Sized {
30    fn to_array_vec(self) -> ArrayVec<T, N>;
31}
32
33impl<T, const N: usize> ToArrayVec<T, N> for [T; N] {
34    fn to_array_vec(self) -> ArrayVec<T, N> {
35        ArrayVec(self)
36    }
37}
38
39impl<T, const N: usize> Default for ArrayVec<T, N>
40where
41    T: Default,
42{
43    fn default() -> Self {
44        Self(std::array::from_fn(|_| T::default()))
45    }
46}
47
48impl<T, const N: usize> ArrayVec<T, N> {
49    pub fn new(data: [T; N]) -> Self {
50        Self(data)
51    }
52
53    pub fn map<U>(&self, transform: impl FnMut(&T) -> U) -> ArrayVec<U, N> {
54        ArrayVec(array_from_iter(self.0.iter().map(transform)))
55    }
56
57    pub fn zip_with<U, V>(
58        &self,
59        other: &ArrayVec<U, N>,
60        mut combine: impl FnMut(&T, &U) -> V,
61    ) -> ArrayVec<V, N> {
62        ArrayVec(array_from_iter(
63            self.0
64                .iter()
65                .zip(other.0.iter())
66                .map(|(left, right)| combine(left, right)),
67        ))
68    }
69}
70
71impl<T, const N: usize> From<[T; N]> for ArrayVec<T, N> {
72    fn from(data: [T; N]) -> Self {
73        Self(data)
74    }
75}
76
77impl<T, const N: usize> From<ArrayVec<T, N>> for [T; N] {
78    fn from(data: ArrayVec<T, N>) -> Self {
79        data.0
80    }
81}
82
83impl<T, const N: usize> AsRef<[T; N]> for ArrayVec<T, N> {
84    fn as_ref(&self) -> &[T; N] {
85        &self.0
86    }
87}
88
89impl<T, const N: usize> AsMut<[T; N]> for ArrayVec<T, N> {
90    fn as_mut(&mut self) -> &mut [T; N] {
91        &mut self.0
92    }
93}
94
95impl<T, const N: usize> Index<usize> for ArrayVec<T, N> {
96    type Output = T;
97    fn index(&self, index: usize) -> &Self::Output {
98        &self.0[index]
99    }
100}
101
102impl<T, const N: usize> IndexMut<usize> for ArrayVec<T, N> {
103    fn index_mut(&mut self, index: usize) -> &mut Self::Output {
104        &mut self.0[index]
105    }
106}
107
108#[inline]
109fn array_from_iter<T, I, const N: usize>(mut iter: I) -> [T; N]
110where
111    I: Iterator<Item = T>,
112{
113    std::array::from_fn(|_| iter.next().unwrap())
114}
115
116macro_rules! impl_arrayvec_binop {
117    ($imp:ident, $method:ident, $op:tt) => {
118        impl<T, U, V, const N: usize> $imp<ArrayVec<U, N>> for ArrayVec<T, N>
119        where
120            T: $imp<U, Output = V>,
121        {
122            type Output = ArrayVec<V, N>;
123            fn $method(self, rhs: ArrayVec<U, N>) -> Self::Output {
124                ArrayVec(array_from_iter(
125                    self.0
126                        .into_iter()
127                        .zip(rhs.0.into_iter())
128                        .map(|(left_value, right_value)| left_value $op right_value),
129                ))
130            }
131        }
132        impl<T, U, V, const N: usize> $imp<&ArrayVec<U, N>> for ArrayVec<T, N>
133        where
134            T: $imp<U, Output = V>,
135            U: Clone,
136        {
137            type Output = ArrayVec<V, N>;
138            fn $method(self, rhs: &ArrayVec<U, N>) -> Self::Output {
139                $imp::$method(self, rhs.clone())
140            }
141        }
142        impl<T, U, V, const N: usize> $imp<ArrayVec<U, N>> for &ArrayVec<T, N>
143        where
144            T: Clone + $imp<U, Output = V>,
145        {
146            type Output = ArrayVec<V, N>;
147            fn $method(self, rhs: ArrayVec<U, N>) -> Self::Output {
148                $imp::$method(self.clone(), rhs)
149            }
150        }
151        impl<T, U, V, const N: usize> $imp<&ArrayVec<U, N>> for &ArrayVec<T, N>
152        where
153            T: Clone + $imp<U, Output = V>,
154            U: Clone,
155        {
156            type Output = ArrayVec<V, N>;
157            fn $method(self, rhs: &ArrayVec<U, N>) -> Self::Output {
158                $imp::$method(self.clone(), rhs.clone())
159            }
160        }
161
162        impl<T, U, V, const N: usize> $imp<ArrayVecScalar<U>> for ArrayVec<T, N>
163        where
164            T: $imp<U, Output = V>,
165            U: Clone,
166        {
167            type Output = ArrayVec<V, N>;
168            fn $method(self, rhs: ArrayVecScalar<U>) -> Self::Output {
169                let scalar_value = rhs.0;
170                ArrayVec(array_from_iter(
171                    self.0
172                        .into_iter()
173                        .map(|value| value $op scalar_value.clone()),
174                ))
175            }
176        }
177        impl<T, U, V, const N: usize> $imp<&ArrayVecScalar<U>> for ArrayVec<T, N>
178        where
179            T: $imp<U, Output = V>,
180            U: Clone,
181        {
182            type Output = ArrayVec<V, N>;
183            fn $method(self, rhs: &ArrayVecScalar<U>) -> Self::Output {
184                $imp::$method(self, rhs.clone())
185            }
186        }
187        impl<T, U, V, const N: usize> $imp<ArrayVecScalar<U>> for &ArrayVec<T, N>
188        where
189            T: Clone + $imp<U, Output = V>,
190            U: Clone,
191        {
192            type Output = ArrayVec<V, N>;
193            fn $method(self, rhs: ArrayVecScalar<U>) -> Self::Output {
194                $imp::$method(self.clone(), rhs)
195            }
196        }
197        impl<T, U, V, const N: usize> $imp<&ArrayVecScalar<U>> for &ArrayVec<T, N>
198        where
199            T: Clone + $imp<U, Output = V>,
200            U: Clone,
201        {
202            type Output = ArrayVec<V, N>;
203            fn $method(self, rhs: &ArrayVecScalar<U>) -> Self::Output {
204                $imp::$method(self.clone(), rhs.clone())
205            }
206        }
207
208        impl<T, U, V, const N: usize> $imp<ArrayVec<T, N>> for ArrayVecScalar<U>
209        where
210            U: Clone + $imp<T, Output = V>,
211        {
212            type Output = ArrayVec<V, N>;
213            fn $method(self, rhs: ArrayVec<T, N>) -> Self::Output {
214                let scalar_value = self.0;
215                ArrayVec(array_from_iter(
216                    rhs.0
217                        .into_iter()
218                        .map(|value| scalar_value.clone() $op value),
219                ))
220            }
221        }
222        impl<T, U, V, const N: usize> $imp<&ArrayVec<T, N>> for ArrayVecScalar<U>
223        where
224            U: Clone + $imp<T, Output = V>,
225            T: Clone,
226        {
227            type Output = ArrayVec<V, N>;
228            fn $method(self, rhs: &ArrayVec<T, N>) -> Self::Output {
229                $imp::$method(self, rhs.clone())
230            }
231        }
232        impl<T, U, V, const N: usize> $imp<ArrayVec<T, N>> for &ArrayVecScalar<U>
233        where
234            U: Clone + $imp<T, Output = V>,
235        {
236            type Output = ArrayVec<V, N>;
237            fn $method(self, rhs: ArrayVec<T, N>) -> Self::Output {
238                $imp::$method(self.clone(), rhs)
239            }
240        }
241        impl<T, U, V, const N: usize> $imp<&ArrayVec<T, N>> for &ArrayVecScalar<U>
242        where
243            U: Clone + $imp<T, Output = V>,
244            T: Clone,
245        {
246            type Output = ArrayVec<V, N>;
247            fn $method(self, rhs: &ArrayVec<T, N>) -> Self::Output {
248                $imp::$method(self.clone(), rhs.clone())
249            }
250        }
251    };
252}
253
254macro_rules! impl_arrayvec_unop {
255    ($imp:ident, $method:ident, $op:tt) => {
256        impl<T, U, const N: usize> $imp for ArrayVec<T, N>
257        where
258            T: $imp<Output = U>,
259        {
260            type Output = ArrayVec<U, N>;
261            fn $method(self) -> Self::Output {
262                ArrayVec(array_from_iter(
263                    self.0.into_iter().map(|value| $op value),
264                ))
265            }
266        }
267        impl<T, U, const N: usize> $imp for &ArrayVec<T, N>
268        where
269            T: Clone + $imp<Output = U>,
270        {
271            type Output = ArrayVec<U, N>;
272            fn $method(self) -> Self::Output {
273                $imp::$method(self.clone())
274            }
275        }
276    };
277}
278
279macro_rules! impl_arrayvec_assign {
280    ($imp:ident, $method:ident) => {
281        impl<T, U, const N: usize> $imp<ArrayVec<U, N>> for ArrayVec<T, N>
282        where
283            T: $imp<U>,
284        {
285            fn $method(&mut self, rhs: ArrayVec<U, N>) {
286                for (left_value, right_value) in self.0.iter_mut().zip(rhs.0.into_iter()) {
287                    left_value.$method(right_value);
288                }
289            }
290        }
291        impl<T, U, const N: usize> $imp<&ArrayVec<U, N>> for ArrayVec<T, N>
292        where
293            T: $imp<U>,
294            U: Clone,
295        {
296            fn $method(&mut self, rhs: &ArrayVec<U, N>) {
297                for (left_value, right_value) in self.0.iter_mut().zip(rhs.0.iter()) {
298                    left_value.$method(right_value.clone());
299                }
300            }
301        }
302        impl<T, U, const N: usize> $imp<ArrayVecScalar<U>> for ArrayVec<T, N>
303        where
304            T: $imp<U>,
305            U: Clone,
306        {
307            fn $method(&mut self, rhs: ArrayVecScalar<U>) {
308                let scalar_value = rhs.0;
309                for value in self.0.iter_mut() {
310                    value.$method(scalar_value.clone());
311                }
312            }
313        }
314        impl<T, U, const N: usize> $imp<&ArrayVecScalar<U>> for ArrayVec<T, N>
315        where
316            T: $imp<U>,
317            U: Clone,
318        {
319            fn $method(&mut self, rhs: &ArrayVecScalar<U>) {
320                self.$method(rhs.clone());
321            }
322        }
323    };
324}
325
326impl_arrayvec_binop!(Add, add, +);
327impl_arrayvec_binop!(Sub, sub, -);
328impl_arrayvec_binop!(Mul, mul, *);
329impl_arrayvec_binop!(Div, div, /);
330impl_arrayvec_binop!(Rem, rem, %);
331impl_arrayvec_binop!(BitAnd, bitand, &);
332impl_arrayvec_binop!(BitOr, bitor, |);
333impl_arrayvec_binop!(BitXor, bitxor, ^);
334impl_arrayvec_binop!(Shl, shl, <<);
335impl_arrayvec_binop!(Shr, shr, >>);
336
337impl_arrayvec_unop!(Neg, neg, -);
338impl_arrayvec_unop!(Not, not, !);
339
340impl_arrayvec_assign!(AddAssign, add_assign);
341impl_arrayvec_assign!(SubAssign, sub_assign);
342impl_arrayvec_assign!(MulAssign, mul_assign);
343impl_arrayvec_assign!(DivAssign, div_assign);
344impl_arrayvec_assign!(RemAssign, rem_assign);
345impl_arrayvec_assign!(BitAndAssign, bitand_assign);
346impl_arrayvec_assign!(BitOrAssign, bitor_assign);
347impl_arrayvec_assign!(BitXorAssign, bitxor_assign);
348impl_arrayvec_assign!(ShlAssign, shl_assign);
349impl_arrayvec_assign!(ShrAssign, shr_assign);
350
351#[cfg(test)]
352mod tests {
353    use super::*;
354    use crate::tools::Xorshift;
355    use std::array;
356    use std::ops::Add;
357
358    #[derive(Clone, Copy, Debug, PartialEq, Eq)]
359    struct LeftValue(i32);
360
361    #[derive(Clone, Copy, Debug, PartialEq, Eq)]
362    struct RightValue(i32);
363
364    #[derive(Clone, Copy, Debug, PartialEq, Eq)]
365    struct SumValue(i32);
366
367    impl Add<RightValue> for LeftValue {
368        type Output = SumValue;
369        fn add(self, rhs: RightValue) -> Self::Output {
370            SumValue(self.0 + rhs.0)
371        }
372    }
373
374    #[derive(Clone, Copy, Debug, PartialEq, Eq)]
375    struct ScalarValue(i32);
376
377    impl Add<i32> for ScalarValue {
378        type Output = i64;
379        fn add(self, rhs: i32) -> Self::Output {
380            self.0 as i64 + rhs as i64
381        }
382    }
383
384    #[test]
385    fn test_array_operations() {
386        let mut rng = Xorshift::default();
387        for _ in 0..1000 {
388            let a: [i32; 8] = array::from_fn(|_| rng.random(-100..=100));
389            let b: [i32; 8] = array::from_fn(|_| rng.random(1..=100));
390            let x = rng.random(1..=100i32);
391            let left = a.to_array_vec();
392            let right = b.to_array_vec();
393            assert_eq!(
394                (a.map(LeftValue).to_array_vec() + b.map(RightValue).to_array_vec()).0,
395                array::from_fn(|i| SumValue(a[i] + b[i]))
396            );
397            assert_eq!(
398                (a.map(ScalarValue).to_array_vec() + x.to_array_vec_scalar()).0,
399                a.map(|a| i64::from(a) + i64::from(x))
400            );
401            assert_eq!((left + right).0, array::from_fn(|i| a[i] + b[i]));
402            assert_eq!((left + x.to_array_vec_scalar()).0, a.map(|a| a + x));
403            let mut actual = left;
404            actual += right;
405            assert_eq!(actual.0, array::from_fn(|i| a[i] + b[i]));
406            let mut actual = left;
407            actual += &right;
408            assert_eq!(actual.0, array::from_fn(|i| a[i] + b[i]));
409            let mut actual = left;
410            actual += x.to_array_vec_scalar();
411            assert_eq!(actual.0, a.map(|a| a + x));
412            assert_eq!((left - right).0, array::from_fn(|i| a[i] - b[i]));
413            assert_eq!((left - x.to_array_vec_scalar()).0, a.map(|a| a - x));
414            let mut actual = left;
415            actual -= right;
416            assert_eq!(actual.0, array::from_fn(|i| a[i] - b[i]));
417            let mut actual = left;
418            actual -= &right;
419            assert_eq!(actual.0, array::from_fn(|i| a[i] - b[i]));
420            let mut actual = left;
421            actual -= x.to_array_vec_scalar();
422            assert_eq!(actual.0, a.map(|a| a - x));
423            assert_eq!((left * right).0, array::from_fn(|i| a[i] * b[i]));
424            assert_eq!((left * x.to_array_vec_scalar()).0, a.map(|a| a * x));
425            let mut actual = left;
426            actual *= right;
427            assert_eq!(actual.0, array::from_fn(|i| a[i] * b[i]));
428            let mut actual = left;
429            actual *= &right;
430            assert_eq!(actual.0, array::from_fn(|i| a[i] * b[i]));
431            let mut actual = left;
432            actual *= x.to_array_vec_scalar();
433            assert_eq!(actual.0, a.map(|a| a * x));
434            assert_eq!((left / right).0, array::from_fn(|i| a[i] / b[i]));
435            assert_eq!((left / x.to_array_vec_scalar()).0, a.map(|a| a / x));
436            let mut actual = left;
437            actual /= right;
438            assert_eq!(actual.0, array::from_fn(|i| a[i] / b[i]));
439            let mut actual = left;
440            actual /= &right;
441            assert_eq!(actual.0, array::from_fn(|i| a[i] / b[i]));
442            let mut actual = left;
443            actual /= x.to_array_vec_scalar();
444            assert_eq!(actual.0, a.map(|a| a / x));
445            assert_eq!((left % right).0, array::from_fn(|i| a[i] % b[i]));
446            assert_eq!((left % x.to_array_vec_scalar()).0, a.map(|a| a % x));
447            let mut actual = left;
448            actual %= right;
449            assert_eq!(actual.0, array::from_fn(|i| a[i] % b[i]));
450            let mut actual = left;
451            actual %= &right;
452            assert_eq!(actual.0, array::from_fn(|i| a[i] % b[i]));
453            let mut actual = left;
454            actual %= x.to_array_vec_scalar();
455            assert_eq!(actual.0, a.map(|a| a % x));
456            assert_eq!((left & right).0, array::from_fn(|i| a[i] & b[i]));
457            assert_eq!((left & x.to_array_vec_scalar()).0, a.map(|a| a & x));
458            let mut actual = left;
459            actual &= right;
460            assert_eq!(actual.0, array::from_fn(|i| a[i] & b[i]));
461            let mut actual = left;
462            actual &= &right;
463            assert_eq!(actual.0, array::from_fn(|i| a[i] & b[i]));
464            let mut actual = left;
465            actual &= x.to_array_vec_scalar();
466            assert_eq!(actual.0, a.map(|a| a & x));
467            assert_eq!((left | right).0, array::from_fn(|i| a[i] | b[i]));
468            assert_eq!((left | x.to_array_vec_scalar()).0, a.map(|a| a | x));
469            let mut actual = left;
470            actual |= right;
471            assert_eq!(actual.0, array::from_fn(|i| a[i] | b[i]));
472            let mut actual = left;
473            actual |= &right;
474            assert_eq!(actual.0, array::from_fn(|i| a[i] | b[i]));
475            let mut actual = left;
476            actual |= x.to_array_vec_scalar();
477            assert_eq!(actual.0, a.map(|a| a | x));
478            assert_eq!((left ^ right).0, array::from_fn(|i| a[i] ^ b[i]));
479            assert_eq!((left ^ x.to_array_vec_scalar()).0, a.map(|a| a ^ x));
480            let mut actual = left;
481            actual ^= right;
482            assert_eq!(actual.0, array::from_fn(|i| a[i] ^ b[i]));
483            let mut actual = left;
484            actual ^= &right;
485            assert_eq!(actual.0, array::from_fn(|i| a[i] ^ b[i]));
486            let mut actual = left;
487            actual ^= x.to_array_vec_scalar();
488            assert_eq!(actual.0, a.map(|a| a ^ x));
489            assert_eq!((x.to_array_vec_scalar() * left).0, a.map(|a| x * a));
490            let shifts: [u32; 8] = array::from_fn(|_| rng.random(0..32));
491            let shift = rng.random(0..32u32);
492            assert_eq!(
493                (left << shifts.to_array_vec()).0,
494                array::from_fn(|i| a[i] << shifts[i])
495            );
496            assert_eq!(
497                (left << shift.to_array_vec_scalar()).0,
498                a.map(|a| a << shift)
499            );
500            let mut actual = left;
501            actual <<= shift.to_array_vec_scalar();
502            assert_eq!(actual.0, a.map(|a| a << shift));
503            assert_eq!(
504                (left >> shifts.to_array_vec()).0,
505                array::from_fn(|i| a[i] >> shifts[i])
506            );
507            assert_eq!(
508                (left >> shift.to_array_vec_scalar()).0,
509                a.map(|a| a >> shift)
510            );
511            let mut actual = left;
512            actual >>= shift.to_array_vec_scalar();
513            assert_eq!(actual.0, a.map(|a| a >> shift));
514        }
515    }
516}