Skip to main content

moirai_iter/parallel/adapters/
map.rs

1use super::super::{fallible, Consumer, MapConsumer, ParallelIterator, TryStreamItem, VecParIter};
2use super::chunks::Chunks;
3use super::pair::Interleave;
4use super::ref_ops::Enumerate;
5use super::stride::StepBy;
6
7/// Map adapter for parallel iterators.
8pub struct Map<I, F> {
9    pub(super) base: I,
10    pub(super) map_fn: F,
11}
12
13impl<I, F> Map<I, F> {
14    pub(crate) fn new(base: I, map_fn: F) -> Self {
15        Self { base, map_fn }
16    }
17
18    /// Reduce a mapped fallible stream without materializing mapped items first.
19    pub fn try_reduce_with<ReduceFn, R>(self, reduce_fn: ReduceFn) -> Option<R>
20    where
21        I: ParallelIterator,
22        F: Fn(I::Item) -> R + Send + Sync + Clone,
23        R: TryStreamItem,
24        ReduceFn: Fn(<R as TryStreamItem>::Output, <R as TryStreamItem>::Output) -> R
25            + Send
26            + Sync
27            + Clone,
28    {
29        fallible::try_reduce_with_items(
30            self.base.seq_items().into_iter().map(self.map_fn),
31            reduce_fn,
32        )
33    }
34
35    /// Return mapped logical indices without materializing the mapped stream.
36    pub fn positions<Predicate, R>(
37        self,
38        predicate: Predicate,
39    ) -> super::position::MapPositions<I, F, Predicate>
40    where
41        I: ParallelIterator,
42        F: Fn(I::Item) -> R + Send + Sync + Clone,
43        Predicate: Fn(R) -> bool + Send + Sync + Clone,
44        R: Send,
45    {
46        super::position::MapPositions::new(self.base, self.map_fn, predicate)
47    }
48}
49
50impl<I, F, R> ParallelIterator for Map<I, F>
51where
52    I: ParallelIterator,
53    F: Fn(I::Item) -> R + Send + Sync + Clone,
54    R: Send,
55{
56    type Item = R;
57
58    fn seq_items(self) -> Vec<Self::Item> {
59        self.base.seq_items().into_iter().map(self.map_fn).collect()
60    }
61
62    fn seq_items_window(self, skip: usize, take: Option<usize>) -> Vec<Self::Item> {
63        self.base
64            .seq_items_window(skip, take)
65            .into_iter()
66            .map(self.map_fn)
67            .collect()
68    }
69
70    fn seq_items_reversed(self) -> Vec<Self::Item> {
71        self.base
72            .seq_items_reversed()
73            .into_iter()
74            .map(self.map_fn)
75            .collect()
76    }
77
78    fn seq_items_reversed_prefix(self, count: usize) -> Vec<Self::Item> {
79        self.base
80            .seq_items_reversed_prefix(count)
81            .into_iter()
82            .map(self.map_fn)
83            .collect()
84    }
85
86    fn drive<C, R2>(self, consumer: C) -> R2
87    where
88        C: Consumer<Self::Item, Result = R2> + Send + Sync,
89        R2: Send,
90    {
91        self.base.drive(MapConsumer::new(consumer, self.map_fn))
92    }
93
94    fn position_first<P>(self, predicate: P) -> Option<usize>
95    where
96        P: Fn(Self::Item) -> bool + Send + Sync + Clone,
97    {
98        self.base
99            .seq_items()
100            .into_iter()
101            .map(self.map_fn)
102            .position(predicate)
103    }
104
105    fn position_any<P>(self, predicate: P) -> Option<usize>
106    where
107        P: Fn(Self::Item) -> bool + Send + Sync + Clone,
108    {
109        self.position_first(predicate)
110    }
111
112    fn position_last<P>(self, predicate: P) -> Option<usize>
113    where
114        P: Fn(Self::Item) -> bool + Send + Sync + Clone,
115    {
116        self.base
117            .seq_items()
118            .into_iter()
119            .map(self.map_fn)
120            .rposition(predicate)
121    }
122
123    fn try_reduce<Identity, ReduceFn, T, E>(
124        self,
125        identity: Identity,
126        reduce_fn: ReduceFn,
127    ) -> Result<T, E>
128    where
129        Identity: Fn() -> T + Send + Sync + Clone,
130        ReduceFn: Fn(T, T) -> Result<T, E> + Send + Sync + Clone,
131        R: Into<Result<T, E>>,
132        T: Send,
133        E: Send,
134    {
135        let mut accumulator = identity();
136        for item in self.base.seq_items().into_iter().map(self.map_fn) {
137            accumulator = reduce_fn(accumulator, item.into()?)?;
138        }
139        Ok(accumulator)
140    }
141}
142
143impl<I, MapFn, Mapped> Map<Chunks<I>, MapFn>
144where
145    I: ParallelIterator,
146    I::Item: Sync + 'static,
147    MapFn: Fn(Vec<I::Item>) -> Mapped + Send + Sync + Clone,
148    Mapped: Send,
149{
150    /// Sum mapped chunk outputs without materializing the chunk-output stream.
151    pub fn sum<S>(self) -> S
152    where
153        S: std::iter::Sum<Mapped> + Send,
154    {
155        let map_fn = self.map_fn;
156        let (base, chunk_size) = self.base.into_parts();
157        let mut items = base.seq_items().into_iter();
158
159        std::iter::from_fn(move || {
160            let chunk: Vec<_> = items.by_ref().take(chunk_size).collect();
161            (!chunk.is_empty()).then(|| map_fn(chunk))
162        })
163        .sum()
164    }
165}
166
167impl<T, MapFn, Mapped> Map<Enumerate<Interleave<StepBy<VecParIter<T>>, VecParIter<T>>>, MapFn>
168where
169    T: Send + Sync + 'static,
170    MapFn: Fn((usize, T)) -> Mapped + Send + Sync + Clone,
171    Mapped: Send,
172{
173    /// Sum mapped vector-backed interleaved index/value pairs without building pair streams.
174    pub fn sum<S>(self) -> S
175    where
176        S: std::iter::Sum<Mapped> + Send,
177    {
178        let map_fn = self.map_fn;
179        let interleave = self.base.base;
180        let step = interleave.left.step();
181        let left = interleave.left.base.into_vec();
182        let right = interleave.right.into_vec();
183        let left_count = if left.is_empty() {
184            0
185        } else {
186            ((left.len() - 1) / step) + 1
187        };
188        let right_count = right.len();
189        let paired_count = left_count.min(right_count);
190        let tail_start = paired_count
191            .checked_mul(2)
192            .expect("interleave index overflow");
193
194        if left_count <= right_count {
195            let mut left = left.into_iter().step_by(step);
196            let mut right = right.into_iter();
197            let mut index = 0usize;
198            let mut pending_right = None;
199            let mut paired_done = false;
200
201            std::iter::from_fn(move || {
202                if let Some(mapped) = pending_right.take() {
203                    return Some(mapped);
204                }
205
206                if !paired_done {
207                    if let Some(left_value) = left.next() {
208                        let right_value = right
209                            .next()
210                            .expect("right side must cover paired interleave item");
211                        let left_index = index;
212                        let right_index = index.checked_add(1).expect("interleave index overflow");
213                        index = index.checked_add(2).expect("interleave index overflow");
214                        pending_right = Some(map_fn((right_index, right_value)));
215                        return Some(map_fn((left_index, left_value)));
216                    }
217                    paired_done = true;
218                    index = tail_start;
219                }
220
221                right.next().map(|value| {
222                    let mapped = map_fn((index, value));
223                    index = index.checked_add(1).expect("interleave index overflow");
224                    mapped
225                })
226            })
227            .sum()
228        } else {
229            let mut left = left.into_iter().step_by(step);
230            let mut right = right.into_iter();
231            let mut index = 0usize;
232            let mut pending_right = None;
233            let mut paired_done = false;
234
235            std::iter::from_fn(move || {
236                if let Some(mapped) = pending_right.take() {
237                    return Some(mapped);
238                }
239
240                if !paired_done {
241                    if let Some(right_value) = right.next() {
242                        let left_value = left
243                            .next()
244                            .expect("left side must cover paired interleave item");
245                        let left_index = index;
246                        let right_index = index.checked_add(1).expect("interleave index overflow");
247                        index = index.checked_add(2).expect("interleave index overflow");
248                        pending_right = Some(map_fn((right_index, right_value)));
249                        return Some(map_fn((left_index, left_value)));
250                    }
251                    paired_done = true;
252                    index = tail_start;
253                }
254
255                left.next().map(|value| {
256                    let mapped = map_fn((index, value));
257                    index = index.checked_add(1).expect("interleave index overflow");
258                    mapped
259                })
260            })
261            .sum()
262        }
263    }
264}
265
266/// Map adapter with cloned per-operation state.
267pub struct MapWith<I, T, F> {
268    pub(super) base: I,
269    pub(super) init: T,
270    pub(super) map_fn: F,
271}
272
273impl<I, T, F> MapWith<I, T, F> {
274    pub(crate) fn new(base: I, init: T, map_fn: F) -> Self {
275        Self { base, init, map_fn }
276    }
277}
278
279impl<I, T, F, R> ParallelIterator for MapWith<I, T, F>
280where
281    I: ParallelIterator,
282    T: Send + Clone,
283    F: Fn(&mut T, I::Item) -> R + Send + Sync + Clone,
284    R: Send + Sync + 'static,
285{
286    type Item = R;
287
288    fn seq_items(self) -> Vec<Self::Item> {
289        let mut state = self.init;
290        self.base
291            .seq_items()
292            .into_iter()
293            .map(|item| (self.map_fn)(&mut state, item))
294            .collect()
295    }
296
297    fn drive<C, R2>(self, consumer: C) -> R2
298    where
299        C: Consumer<Self::Item, Result = R2> + Send + Sync,
300        R2: Send,
301    {
302        consumer.consume(VecParIter::new(self.seq_items()))
303    }
304}
305
306/// Map adapter with lazily initialized state.
307pub struct MapInit<I, Init, F> {
308    pub(super) base: I,
309    pub(super) init: Init,
310    pub(super) map_fn: F,
311}
312
313impl<I, Init, F> MapInit<I, Init, F> {
314    pub(crate) fn new(base: I, init: Init, map_fn: F) -> Self {
315        Self { base, init, map_fn }
316    }
317}
318
319impl<I, Init, T, F, R> ParallelIterator for MapInit<I, Init, F>
320where
321    I: ParallelIterator,
322    Init: Fn() -> T + Send + Sync + Clone,
323    T: Send,
324    F: Fn(&mut T, I::Item) -> R + Send + Sync + Clone,
325    R: Send + Sync + 'static,
326{
327    type Item = R;
328
329    fn seq_items(self) -> Vec<Self::Item> {
330        let mut state = (self.init)();
331        self.base
332            .seq_items()
333            .into_iter()
334            .map(|item| (self.map_fn)(&mut state, item))
335            .collect()
336    }
337
338    fn drive<C, R2>(self, consumer: C) -> R2
339    where
340        C: Consumer<Self::Item, Result = R2> + Send + Sync,
341        R2: Send,
342    {
343        consumer.consume(VecParIter::new(self.seq_items()))
344    }
345}
346
347/// Update adapter that mutates each item before yielding it.
348pub struct Update<I, F> {
349    pub(super) base: I,
350    pub(super) update_fn: F,
351}
352
353impl<I, F> Update<I, F> {
354    pub(crate) fn new(base: I, update_fn: F) -> Self {
355        Self { base, update_fn }
356    }
357}
358
359impl<I, F> ParallelIterator for Update<I, F>
360where
361    I: ParallelIterator,
362    F: Fn(&mut I::Item) + Send + Sync + Clone,
363    I::Item: Sync + 'static,
364{
365    type Item = I::Item;
366
367    fn seq_items(self) -> Vec<Self::Item> {
368        self.base
369            .seq_items()
370            .into_iter()
371            .map(|mut item| {
372                (self.update_fn)(&mut item);
373                item
374            })
375            .collect()
376    }
377
378    fn drive<C, R>(self, consumer: C) -> R
379    where
380        C: Consumer<Self::Item, Result = R> + Send + Sync,
381        R: Send,
382    {
383        consumer.consume(VecParIter::new(self.seq_items()))
384    }
385}