Skip to main content

p3_maybe_rayon/
serial.rs

1use core::iter::{FlatMap, IntoIterator, Iterator};
2use core::marker::{Send, Sized, Sync};
3use core::ops::{Fn, FnOnce};
4use core::slice::{
5    Chunks, ChunksExact, ChunksExactMut, ChunksMut, RChunks, RChunksExact, RChunksExactMut,
6    RChunksMut, Split, SplitMut, Windows,
7};
8
9pub trait IntoParallelIterator {
10    type Iter: Iterator<Item = Self::Item>;
11    type Item: Send;
12
13    fn into_par_iter(self) -> Self::Iter;
14}
15impl<T: IntoIterator> IntoParallelIterator for T
16where
17    T::Item: Send,
18{
19    type Iter = T::IntoIter;
20    type Item = T::Item;
21
22    fn into_par_iter(self) -> Self::Iter {
23        self.into_iter()
24    }
25}
26
27pub trait IntoParallelRefIterator<'data> {
28    type Iter: Iterator<Item = Self::Item>;
29    type Item: Send + 'data;
30
31    fn par_iter(&'data self) -> Self::Iter;
32}
33
34impl<'data, I: 'data + ?Sized> IntoParallelRefIterator<'data> for I
35where
36    &'data I: IntoParallelIterator,
37{
38    type Iter = <&'data I as IntoParallelIterator>::Iter;
39    type Item = <&'data I as IntoParallelIterator>::Item;
40
41    fn par_iter(&'data self) -> Self::Iter {
42        self.into_par_iter()
43    }
44}
45
46pub trait IntoParallelRefMutIterator<'data> {
47    type Iter: Iterator<Item = Self::Item>;
48    type Item: Send + 'data;
49
50    fn par_iter_mut(&'data mut self) -> Self::Iter;
51}
52
53impl<'data, I: 'data + ?Sized> IntoParallelRefMutIterator<'data> for I
54where
55    &'data mut I: IntoParallelIterator,
56{
57    type Iter = <&'data mut I as IntoParallelIterator>::Iter;
58    type Item = <&'data mut I as IntoParallelIterator>::Item;
59
60    fn par_iter_mut(&'data mut self) -> Self::Iter {
61        self.into_par_iter()
62    }
63}
64
65pub trait ParallelSlice<T: Sync> {
66    /// Returns a plain slice, which is used to implement the rest of the
67    /// parallel methods.
68    fn as_parallel_slice(&self) -> &[T];
69
70    fn par_split<P>(&self, separator: P) -> Split<'_, T, P>
71    where
72        P: Fn(&T) -> bool + Sync + Send,
73    {
74        self.as_parallel_slice().split(separator)
75    }
76
77    fn par_windows(&self, window_size: usize) -> Windows<'_, T> {
78        self.as_parallel_slice().windows(window_size)
79    }
80
81    fn par_chunks(&self, chunk_size: usize) -> Chunks<'_, T> {
82        self.as_parallel_slice().chunks(chunk_size)
83    }
84
85    fn par_chunks_exact(&self, chunk_size: usize) -> ChunksExact<'_, T> {
86        self.as_parallel_slice().chunks_exact(chunk_size)
87    }
88
89    fn par_rchunks(&self, chunk_size: usize) -> RChunks<'_, T> {
90        self.as_parallel_slice().rchunks(chunk_size)
91    }
92
93    fn par_rchunks_exact(&self, chunk_size: usize) -> RChunksExact<'_, T> {
94        self.as_parallel_slice().rchunks_exact(chunk_size)
95    }
96}
97
98impl<T: Sync> ParallelSlice<T> for [T] {
99    #[inline]
100    fn as_parallel_slice(&self) -> &[T] {
101        self
102    }
103}
104
105pub trait ParallelSliceMut<T: Send> {
106    /// Returns a plain mutable slice, which is used to implement the rest of
107    /// the parallel methods.
108    fn as_parallel_slice_mut(&mut self) -> &mut [T];
109
110    fn par_split_mut<P>(&mut self, separator: P) -> SplitMut<'_, T, P>
111    where
112        P: Fn(&T) -> bool + Sync + Send,
113    {
114        self.as_parallel_slice_mut().split_mut(separator)
115    }
116
117    fn par_chunks_mut(&mut self, chunk_size: usize) -> ChunksMut<'_, T> {
118        self.as_parallel_slice_mut().chunks_mut(chunk_size)
119    }
120
121    fn par_chunks_exact_mut(&mut self, chunk_size: usize) -> ChunksExactMut<'_, T> {
122        self.as_parallel_slice_mut().chunks_exact_mut(chunk_size)
123    }
124
125    fn par_rchunks_mut(&mut self, chunk_size: usize) -> RChunksMut<'_, T> {
126        self.as_parallel_slice_mut().rchunks_mut(chunk_size)
127    }
128
129    fn par_rchunks_exact_mut(&mut self, chunk_size: usize) -> RChunksExactMut<'_, T> {
130        self.as_parallel_slice_mut().rchunks_exact_mut(chunk_size)
131    }
132}
133
134impl<T: Send> ParallelSliceMut<T> for [T] {
135    #[inline]
136    fn as_parallel_slice_mut(&mut self) -> &mut [T] {
137        self
138    }
139}
140
141pub trait ParIterExt: Iterator {
142    /// Serially, this returns the first matching item, unlike rayon's `find_any`, which may
143    /// return any matching item found by any thread.
144    fn find_any<P>(self, predicate: P) -> Option<Self::Item>
145    where
146        P: Fn(&Self::Item) -> bool + Sync + Send;
147
148    fn find_map_any<P, R>(self, predicate: P) -> Option<R>
149    where
150        P: Fn(Self::Item) -> Option<R> + Sync + Send;
151
152    fn flat_map_iter<U, F>(self, map_op: F) -> FlatMap<Self, U, F>
153    where
154        Self: Sized,
155        U: IntoIterator,
156        F: Fn(Self::Item) -> U;
157
158    /// Initialize one serial state and reuse it for every mapped item.
159    fn map_init<INIT, OP, T, R>(self, init: INIT, op: OP) -> impl Iterator<Item = R>
160    where
161        Self: Sized,
162        INIT: Fn() -> T + Sync + Send,
163        OP: Fn(&mut T, Self::Item) -> R + Sync + Send,
164    {
165        let mut state = init();
166        self.map(move |item| op(&mut state, item))
167    }
168
169    /// Serially, `init` is called once and the single state is reused across every item,
170    /// unlike rayon's `for_each_init`, which calls `init` once per split.
171    fn for_each_init<OP, INIT, T>(self, init: INIT, op: OP)
172    where
173        Self: Sized,
174        OP: Fn(&mut T, Self::Item) + Sync + Send,
175        INIT: Fn() -> T + Sync + Send;
176
177    /// No-op: there is only one split in a serial build, so a minimum split length has
178    /// nothing to constrain.
179    fn with_min_len(self, min_len: usize) -> Self
180    where
181        Self: Sized,
182    {
183        let _ = min_len;
184        self
185    }
186
187    /// No-op: there is only one split in a serial build, so a maximum split length has
188    /// nothing to constrain.
189    fn with_max_len(self, max_len: usize) -> Self
190    where
191        Self: Sized,
192    {
193        let _ = max_len;
194        self
195    }
196}
197
198impl<I: Iterator> ParIterExt for I {
199    fn find_any<P>(mut self, predicate: P) -> Option<Self::Item>
200    where
201        P: Fn(&Self::Item) -> bool + Sync + Send,
202    {
203        self.find(predicate)
204    }
205
206    fn find_map_any<P, R>(mut self, predicate: P) -> Option<R>
207    where
208        P: Fn(Self::Item) -> Option<R> + Sync + Send,
209    {
210        self.find_map(predicate)
211    }
212
213    fn flat_map_iter<U, F>(self, map_op: F) -> FlatMap<Self, U, F>
214    where
215        Self: Sized,
216        U: IntoIterator,
217        F: Fn(Self::Item) -> U,
218    {
219        self.flat_map(map_op)
220    }
221
222    fn for_each_init<OP, INIT, T>(self, init: INIT, op: OP)
223    where
224        Self: Sized,
225        OP: Fn(&mut T, Self::Item) + Sync + Send,
226        INIT: Fn() -> T + Sync + Send,
227    {
228        let mut state = init();
229        self.for_each(|item| op(&mut state, item));
230    }
231}
232
233/// Runs `oper_a` then `oper_b` in sequence, unlike rayon's `join`, which may run them
234/// concurrently on separate threads.
235pub fn join<A, B, RA, RB>(oper_a: A, oper_b: B) -> (RA, RB)
236where
237    A: FnOnce() -> RA,
238    B: FnOnce() -> RB,
239{
240    let result_a = oper_a();
241    let result_b = oper_b();
242    (result_a, result_b)
243}
244
245/// Always 1: there is only one thread of execution in a serial build.
246pub const fn current_num_threads() -> usize {
247    1
248}
249
250#[cfg(test)]
251mod tests {
252    use core::sync::atomic::{AtomicUsize, Ordering};
253
254    use super::ParIterExt;
255
256    #[test]
257    fn map_init_reuses_state_in_order_and_accepts_empty_input() {
258        let initialized = AtomicUsize::new(0);
259        let mut values = (1..4).map_init(
260            || {
261                initialized.fetch_add(1, Ordering::Relaxed);
262                0
263            },
264            |state, item| {
265                *state += item;
266                *state
267            },
268        );
269        assert_eq!(values.next(), Some(1));
270        assert_eq!(values.next(), Some(3));
271        assert_eq!(values.next(), Some(6));
272        assert_eq!(values.next(), None);
273        assert_eq!(initialized.load(Ordering::Relaxed), 1);
274        assert_eq!(
275            (0..0)
276                .map_init(
277                    || 0,
278                    |state, item| {
279                        *state += item;
280                        *state
281                    }
282                )
283                .next(),
284            None
285        );
286    }
287}