Skip to main content

moirai_iter/execution/
parallel.rs

1//! Parallel execution context.
2
3use super::base::ExecutionBase;
4use super::hybrid::owned_chunks;
5use crate::base::SendPtr;
6use std::fmt::Debug;
7
8/// Parallel execution context for CPU-bound work
9///
10/// Work runs on the process-wide scheduler rather than a context-owned pool,
11/// so several contexts share one worker set instead of over-subscribing the
12/// machine with a thread pool each.
13#[derive(Clone)]
14pub struct ParallelContext {
15    chunk_size: usize,
16}
17
18impl Default for ParallelContext {
19    fn default() -> Self {
20        Self::new()
21    }
22}
23
24impl Debug for ParallelContext {
25    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
26        f.debug_struct("ParallelContext")
27            .field("chunk_size", &self.chunk_size)
28            .finish()
29    }
30}
31
32impl ParallelContext {
33    /// Create a new parallel context with the default chunk size
34    pub fn new() -> Self {
35        Self { chunk_size: 1000 }
36    }
37
38    /// Create with specific chunk size
39    pub fn with_chunk_size(chunk_size: usize) -> Self {
40        Self { chunk_size }
41    }
42}
43
44impl ParallelContext {
45    /// Execute an iterator operation with parallel processing
46    pub fn execute_iter<T, F, R>(
47        &self,
48        items: Vec<T>,
49        func: F,
50    ) -> Result<Vec<R>, Box<dyn std::error::Error + Send + Sync>>
51    where
52        T: Send + 'static,
53        F: Fn(T) -> R + Send + Sync + 'static,
54        R: Send + 'static,
55    {
56        if items.is_empty() {
57            return Ok(Vec::new());
58        }
59
60        let chunk_size = self.chunk_size.max(1);
61
62        if items.len() <= chunk_size {
63            return Ok(items.into_iter().map(func).collect());
64        }
65
66        let item_count = items.len();
67        let chunks = owned_chunks(items, chunk_size);
68        let num_chunks = chunks.len();
69
70        // One owned input slot and one owned output slot per chunk. The chunk
71        // is taken out and the result written back through the same index, so
72        // ordering falls out of the index domain rather than a post-hoc sort of
73        // whatever arrived — the previous channel collect ended as soon as the
74        // senders dropped, so a panicking chunk silently returned a short `Vec`.
75        let mut chunks: Vec<Option<Vec<T>>> = chunks.into_iter().map(Some).collect();
76        let mut chunk_results: Vec<Option<Vec<R>>> = (0..num_chunks).map(|_| None).collect();
77        let chunks_ptr = SendPtr(chunks.as_mut_ptr());
78        let results_ptr = SendPtr(chunk_results.as_mut_ptr());
79
80        // SAFETY: the fan-out visits each index in `0..num_chunks` exactly once,
81        // so no two lanes touch the same input or output slot, and both vectors
82        // outlive the joined call.
83        let map_chunk = |idx: usize| unsafe {
84            let chunk = (*chunks_ptr.as_ptr().add(idx))
85                .take()
86                .expect("invariant: each chunk is claimed by exactly one index");
87            let mapped: Vec<R> = chunk.into_iter().map(&func).collect();
88            *results_ptr.as_ptr().add(idx) = Some(mapped);
89        };
90
91        let run_on_global = moirai_executor::global()
92            .for_each_indexed::<moirai_executor::schedule::SyncTask, _>(num_chunks, &map_chunk);
93
94        if crate::base::sequential_fallback_permitted(&run_on_global) {
95            (0..num_chunks).for_each(map_chunk);
96        }
97
98        let mut results = Vec::with_capacity(item_count);
99        for chunk in chunk_results {
100            results.extend(chunk.expect("invariant: every chunk index produced a result"));
101        }
102
103        Ok(results)
104    }
105
106    /// Execute a closure with the context
107    pub fn execute<F, R>(&self, func: F) -> Result<R, Box<dyn std::error::Error + Send + Sync>>
108    where
109        F: FnOnce() -> R + Send,
110        R: Send,
111    {
112        // Execute immediately in parallel context
113        Ok(func())
114    }
115}
116
117impl ExecutionBase for ParallelContext {
118    fn context_type(&self) -> &'static str {
119        "Parallel"
120    }
121}
122
123#[cfg(test)]
124mod tests {
125    use super::*;
126    use crate::test_support::panic_message;
127
128    const CHUNK: usize = 8;
129    const ITEMS: usize = CHUNK * 5;
130
131    #[test]
132    fn execute_iter_returns_every_item_in_input_order() {
133        let context = ParallelContext::with_chunk_size(CHUNK);
134        let items: Vec<usize> = (0..ITEMS).collect();
135
136        let doubled = context
137            .execute_iter(items.clone(), |item| item * 2)
138            .expect("chunked execution must succeed");
139
140        assert_eq!(
141            doubled,
142            items.iter().map(|item| item * 2).collect::<Vec<_>>(),
143            "results must follow input order, not completion order"
144        );
145    }
146
147    #[test]
148    fn execute_iter_propagates_a_chunk_panic_instead_of_truncating() {
149        // The previous channel-collect ended when the senders dropped, so a
150        // panicking chunk returned a short `Vec` and the caller could not tell.
151        // A missing chunk must surface, not shrink the result.
152        let context = ParallelContext::with_chunk_size(CHUNK);
153        let items: Vec<usize> = (0..ITEMS).collect();
154
155        let previous_hook = std::panic::take_hook();
156        std::panic::set_hook(Box::new(|_| {}));
157        let outcome = std::panic::catch_unwind(move || {
158            context.execute_iter(items, |item| {
159                assert_ne!(item, ITEMS - 1, "chunk panic");
160                item
161            })
162        });
163        std::panic::set_hook(previous_hook);
164
165        let payload = outcome
166            .expect_err("a panicking chunk must reach the caller rather than shorten the result");
167        // The chunk's own payload does not travel: the fan-out converts it to a
168        // spawn error and panics with its own invariant, which is what names
169        // the partial execution the caller must not retry.
170        let message = panic_message(&*payload);
171        assert!(
172            message.contains("indexed fan-out failed after partial execution"),
173            "the fan-out must report partial execution, got {message:?}"
174        );
175    }
176}