Skip to main content

scirs2_core/concurrent/
parallel_iter.rs

1//! Parallel iterators over slices and owned `Vec`s.
2//!
3//! All functions in this module use OS threads (no Rayon dependency) and
4//! fall back to sequential execution for small inputs or when only one CPU
5//! is available.
6//!
7//! # Provided operations
8//!
9//! | Function | Description |
10//! |----------|-------------|
11//! | [`parallel_map`] | Apply `f` to every element; preserves order. |
12//! | [`parallel_reduce`] | Reduce with a commutative binary op; associativity required. |
13//! | [`parallel_filter`] | Retain elements matching a predicate. |
14//! | [`parallel_scan`] | Inclusive/exclusive prefix scan. |
15//! | [`parallel_merge_sort`] | In-place parallel merge sort. |
16//! | [`parallel_for_each`] | Execute a closure for each element (side-effects). |
17//! | [`parallel_partition`] | Partition elements into two `Vec`s based on a predicate. |
18//! | [`parallel_prefix_sum`] | Specialised f64 prefix-sum using Blelloch's algorithm. |
19//!
20//! # Example
21//!
22//! ```rust
23//! use scirs2_core::concurrent::parallel_iter::{parallel_map, parallel_reduce, parallel_scan, ScanMode};
24//!
25//! let data: Vec<i64> = (1..=8).collect();
26//! let doubled = parallel_map(&data, |&x| x * 2, 0).expect("map");
27//! assert_eq!(doubled, vec![2, 4, 6, 8, 10, 12, 14, 16]);
28//!
29//! let sum = parallel_reduce(&data, 0i64, |a, b| a + b, |a, b| a + b, 0).expect("reduce");
30//! assert_eq!(sum, 36);
31//!
32//! let prefix = parallel_scan(&data, 0i64, |a, b| a + b, ScanMode::Inclusive, 0).expect("scan");
33//! assert_eq!(prefix, vec![1, 3, 6, 10, 15, 21, 28, 36]);
34//! ```
35
36use std::sync::{Arc, Mutex};
37use std::thread;
38
39use crate::error::{CoreError, CoreResult, ErrorContext, ErrorLocation};
40
41// ── helpers ──────────────────────────────────────────────────────────────────
42
43/// Resolve `n_threads` to a concrete thread count (0 → hardware concurrency).
44pub fn resolve_threads(n_threads: usize) -> usize {
45    if n_threads == 0 {
46        thread::available_parallelism()
47            .map(|p| p.get())
48            .unwrap_or(1)
49    } else {
50        n_threads
51    }
52}
53
54/// Split `len` items into `n_chunks` index ranges.
55fn chunk_ranges(len: usize, n_chunks: usize) -> Vec<std::ops::Range<usize>> {
56    let n = n_chunks.max(1);
57    let base = len / n;
58    let rem = len % n;
59    let mut ranges = Vec::with_capacity(n);
60    let mut start = 0;
61    for i in 0..n {
62        let extra = if i < rem { 1 } else { 0 };
63        let end = (start + base + extra).min(len);
64        if start < len {
65            ranges.push(start..end);
66        }
67        start = start + base + extra;
68    }
69    ranges
70}
71
72fn spawn_err(e: impl std::fmt::Display) -> CoreError {
73    CoreError::SchedulerError(
74        ErrorContext::new(format!("failed to spawn thread: {e}"))
75            .with_location(ErrorLocation::new(file!(), line!())),
76    )
77}
78
79fn join_err(label: &'static str) -> CoreError {
80    CoreError::SchedulerError(
81        ErrorContext::new(format!("{label}: worker thread panicked"))
82            .with_location(ErrorLocation::new(file!(), line!())),
83    )
84}
85
86// ── parallel_map ─────────────────────────────────────────────────────────────
87
88/// Apply `f` to every element of `data` in parallel, returning results in the
89/// same order as the input.
90///
91/// `n_threads = 0` uses hardware concurrency.
92pub fn parallel_map<T, R, F>(data: &[T], f: F, n_threads: usize) -> CoreResult<Vec<R>>
93where
94    T: Sync + 'static,
95    R: Send + Default + Clone + 'static,
96    F: Fn(&T) -> R + Send + Sync + 'static,
97{
98    let n = data.len();
99    if n == 0 {
100        return Ok(Vec::new());
101    }
102
103    let n_threads = resolve_threads(n_threads).min(n);
104
105    // Allocate output vector upfront.
106    let mut out: Vec<R> = vec![R::default(); n];
107
108    // SAFETY: each chunk has a disjoint range of `out`.  We use raw pointers
109    // to hand non-overlapping slices to threads.
110    let data_ptr = data.as_ptr() as usize;
111    let out_ptr = out.as_mut_ptr() as usize;
112    let f = Arc::new(f);
113    let ranges = chunk_ranges(n, n_threads);
114
115    let mut handles = Vec::with_capacity(ranges.len());
116    for range in ranges {
117        let f2 = Arc::clone(&f);
118        let handle = thread::Builder::new()
119            .spawn(move || {
120                let data: &[T] =
121                    // SAFETY: range is within `data`, unique per thread.
122                    unsafe { std::slice::from_raw_parts(data_ptr as *const T, n) };
123                let out: &mut [R] =
124                    // SAFETY: range is within `out`, unique per thread.
125                    unsafe { std::slice::from_raw_parts_mut(out_ptr as *mut R, n) };
126                for i in range {
127                    out[i] = f2(&data[i]);
128                }
129            })
130            .map_err(spawn_err)?;
131        handles.push(handle);
132    }
133
134    for h in handles {
135        h.join().map_err(|_| join_err("parallel_map"))?;
136    }
137
138    Ok(out)
139}
140
141// ── parallel_for_each ────────────────────────────────────────────────────────
142
143/// Execute `f` for each element of `data` in parallel (fire-and-forget,
144/// side-effects only).  Order of execution is not guaranteed.
145pub fn parallel_for_each<T, F>(data: &[T], f: F, n_threads: usize) -> CoreResult<()>
146where
147    T: Sync + 'static,
148    F: Fn(&T) + Send + Sync + 'static,
149{
150    let n = data.len();
151    if n == 0 {
152        return Ok(());
153    }
154
155    let n_threads = resolve_threads(n_threads).min(n);
156    let data_ptr = data.as_ptr() as usize;
157    let f = Arc::new(f);
158    let mut handles = Vec::new();
159
160    for range in chunk_ranges(n, n_threads) {
161        let f2 = Arc::clone(&f);
162        let handle = thread::Builder::new()
163            .spawn(move || {
164                let data: &[T] = unsafe { std::slice::from_raw_parts(data_ptr as *const T, n) };
165                for i in range {
166                    f2(&data[i]);
167                }
168            })
169            .map_err(spawn_err)?;
170        handles.push(handle);
171    }
172
173    for h in handles {
174        h.join().map_err(|_| join_err("parallel_for_each"))?;
175    }
176    Ok(())
177}
178
179// ── parallel_reduce ───────────────────────────────────────────────────────────
180
181/// Parallel reduction using a chunk-local reduce followed by a sequential
182/// combine across chunks.
183///
184/// - `fold`: combine one element `T` into accumulator `R` (called per thread).
185/// - `combine`: merge two accumulators (called sequentially to merge chunks).
186/// - `identity`: neutral element for `combine`.
187pub fn parallel_reduce<T, R, Fold, Combine>(
188    data: &[T],
189    identity: R,
190    fold: Fold,
191    combine: Combine,
192    n_threads: usize,
193) -> CoreResult<R>
194where
195    T: Sync + 'static,
196    R: Send + Clone + 'static,
197    Fold: Fn(R, &T) -> R + Send + Sync + 'static,
198    Combine: Fn(R, R) -> R,
199{
200    let n = data.len();
201    if n == 0 {
202        return Ok(identity);
203    }
204
205    let n_threads = resolve_threads(n_threads).min(n);
206    let data_ptr = data.as_ptr() as usize;
207    let fold = Arc::new(fold);
208    let results: Arc<Mutex<Vec<(usize, R)>>> = Arc::new(Mutex::new(Vec::new()));
209    let ranges = chunk_ranges(n, n_threads);
210    let mut handles = Vec::with_capacity(ranges.len());
211
212    for (chunk_id, range) in ranges.into_iter().enumerate() {
213        let f2 = Arc::clone(&fold);
214        let results2 = Arc::clone(&results);
215        let id = identity.clone();
216        let handle = thread::Builder::new()
217            .spawn(move || {
218                let data: &[T] = unsafe { std::slice::from_raw_parts(data_ptr as *const T, n) };
219                let local = data[range].iter().fold(id, |acc, x| f2(acc, x));
220                if let Ok(mut g) = results2.lock() {
221                    g.push((chunk_id, local));
222                }
223            })
224            .map_err(spawn_err)?;
225        handles.push(handle);
226    }
227
228    for h in handles {
229        h.join().map_err(|_| join_err("parallel_reduce"))?;
230    }
231
232    let mut partials = Arc::try_unwrap(results)
233        .map_err(|_| {
234            CoreError::SchedulerError(ErrorContext::new("parallel_reduce: Arc still held"))
235        })?
236        .into_inner()
237        .map_err(|e| {
238            CoreError::SchedulerError(
239                ErrorContext::new(format!("parallel_reduce: mutex poisoned: {e}"))
240                    .with_location(ErrorLocation::new(file!(), line!())),
241            )
242        })?;
243
244    // Sort by chunk_id to ensure deterministic combine order.
245    partials.sort_by_key(|(id, _)| *id);
246    let result = partials
247        .into_iter()
248        .fold(identity, |acc, (_, r)| combine(acc, r));
249    Ok(result)
250}
251
252// ── parallel_filter ───────────────────────────────────────────────────────────
253
254/// Retain elements matching `pred` in parallel.
255///
256/// The order of elements in the output matches the order in the input.
257pub fn parallel_filter<T, F>(data: Vec<T>, pred: F, n_threads: usize) -> CoreResult<Vec<T>>
258where
259    T: Send + 'static,
260    F: Fn(&T) -> bool + Send + Sync + 'static,
261{
262    let n = data.len();
263    if n == 0 {
264        return Ok(Vec::new());
265    }
266
267    let n_threads = resolve_threads(n_threads).min(n);
268    let pred = Arc::new(pred);
269    let ranges = chunk_ranges(n, n_threads);
270    let n_chunks = ranges.len();
271
272    // Move data into Arc<Mutex<Vec<Option<T>>>> so threads can take elements.
273    let data: Vec<Option<T>> = data.into_iter().map(Some).collect();
274    let shared: Arc<Mutex<Vec<Option<T>>>> = Arc::new(Mutex::new(data));
275    let chunk_results: Arc<Mutex<Vec<(usize, Vec<T>)>>> = Arc::new(Mutex::new(Vec::new()));
276    let mut handles = Vec::with_capacity(n_chunks);
277
278    for (chunk_id, range) in ranges.into_iter().enumerate() {
279        let p2 = Arc::clone(&pred);
280        let sh = Arc::clone(&shared);
281        let cr = Arc::clone(&chunk_results);
282
283        let handle = thread::Builder::new()
284            .spawn(move || {
285                // Extract our chunk's items.
286                let items: Vec<T> = {
287                    if let Ok(mut g) = sh.lock() {
288                        range.filter_map(|i| g[i].take()).collect()
289                    } else {
290                        Vec::new()
291                    }
292                };
293                let kept: Vec<T> = items.into_iter().filter(|x| p2(x)).collect();
294                if let Ok(mut g) = cr.lock() {
295                    g.push((chunk_id, kept));
296                }
297            })
298            .map_err(spawn_err)?;
299        handles.push(handle);
300    }
301
302    for h in handles {
303        h.join().map_err(|_| join_err("parallel_filter"))?;
304    }
305
306    let mut partials = Arc::try_unwrap(chunk_results)
307        .map_err(|_| {
308            CoreError::SchedulerError(ErrorContext::new("parallel_filter: Arc still held"))
309        })?
310        .into_inner()
311        .map_err(|e| {
312            CoreError::SchedulerError(
313                ErrorContext::new(format!("parallel_filter: mutex poisoned: {e}"))
314                    .with_location(ErrorLocation::new(file!(), line!())),
315            )
316        })?;
317
318    partials.sort_by_key(|(id, _)| *id);
319    Ok(partials.into_iter().flat_map(|(_, v)| v).collect())
320}
321
322// ── parallel_partition ────────────────────────────────────────────────────────
323
324/// Partition `data` into `(matching, non_matching)` in parallel.
325///
326/// The order within each partition preserves the original input order.
327pub fn parallel_partition<T, F>(
328    data: Vec<T>,
329    pred: F,
330    n_threads: usize,
331) -> CoreResult<(Vec<T>, Vec<T>)>
332where
333    T: Send + 'static,
334    F: Fn(&T) -> bool + Send + Sync + 'static,
335{
336    let n = data.len();
337    if n == 0 {
338        return Ok((Vec::new(), Vec::new()));
339    }
340
341    let n_threads = resolve_threads(n_threads).min(n);
342    let pred = Arc::new(pred);
343    let ranges = chunk_ranges(n, n_threads);
344
345    let shared: Arc<Mutex<Vec<Option<T>>>> =
346        Arc::new(Mutex::new(data.into_iter().map(Some).collect()));
347    let yes_chunks: Arc<Mutex<Vec<(usize, Vec<T>)>>> = Arc::new(Mutex::new(Vec::new()));
348    let no_chunks: Arc<Mutex<Vec<(usize, Vec<T>)>>> = Arc::new(Mutex::new(Vec::new()));
349    let mut handles = Vec::new();
350
351    for (chunk_id, range) in ranges.into_iter().enumerate() {
352        let p2 = Arc::clone(&pred);
353        let sh = Arc::clone(&shared);
354        let yc = Arc::clone(&yes_chunks);
355        let nc = Arc::clone(&no_chunks);
356
357        let handle = thread::Builder::new()
358            .spawn(move || {
359                let items: Vec<T> = {
360                    if let Ok(mut g) = sh.lock() {
361                        range.filter_map(|i| g[i].take()).collect()
362                    } else {
363                        Vec::new()
364                    }
365                };
366                let (yes, no): (Vec<T>, Vec<T>) = items.into_iter().partition(|x| p2(x));
367                if let Ok(mut g) = yc.lock() {
368                    g.push((chunk_id, yes));
369                }
370                if let Ok(mut g) = nc.lock() {
371                    g.push((chunk_id, no));
372                }
373            })
374            .map_err(spawn_err)?;
375        handles.push(handle);
376    }
377
378    for h in handles {
379        h.join().map_err(|_| join_err("parallel_partition"))?;
380    }
381
382    let mut yes = Arc::try_unwrap(yes_chunks)
383        .map_err(|_| {
384            CoreError::SchedulerError(ErrorContext::new("parallel_partition: yes Arc held"))
385        })?
386        .into_inner()
387        .map_err(|e| {
388            CoreError::SchedulerError(
389                ErrorContext::new(format!("parallel_partition: mutex poisoned: {e}"))
390                    .with_location(ErrorLocation::new(file!(), line!())),
391            )
392        })?;
393    let mut no = Arc::try_unwrap(no_chunks)
394        .map_err(|_| {
395            CoreError::SchedulerError(ErrorContext::new("parallel_partition: no Arc held"))
396        })?
397        .into_inner()
398        .map_err(|e| {
399            CoreError::SchedulerError(
400                ErrorContext::new(format!("parallel_partition: no mutex poisoned: {e}"))
401                    .with_location(ErrorLocation::new(file!(), line!())),
402            )
403        })?;
404
405    yes.sort_by_key(|(id, _)| *id);
406    no.sort_by_key(|(id, _)| *id);
407
408    let yes_flat: Vec<T> = yes.into_iter().flat_map(|(_, v)| v).collect();
409    let no_flat: Vec<T> = no.into_iter().flat_map(|(_, v)| v).collect();
410    Ok((yes_flat, no_flat))
411}
412
413// ── parallel_scan ────────────────────────────────────────────────────────────
414
415/// Scan mode (inclusive or exclusive).
416#[derive(Debug, Clone, Copy, PartialEq, Eq)]
417pub enum ScanMode {
418    /// `output[i] = op(input[0..=i])` — identity-free.
419    Inclusive,
420    /// `output[i] = op(input[0..i])` — `output[0]` equals `identity`.
421    Exclusive,
422}
423
424/// Parallel prefix scan (generalised prefix sum) using Blelloch's algorithm.
425///
426/// The `op` must be *associative* (but need not be commutative).
427///
428/// # Arguments
429/// * `data`      – input slice
430/// * `identity`  – identity element for `op` (needed for `Exclusive` scans)
431/// * `op`        – associative binary operator
432/// * `mode`      – [`ScanMode::Inclusive`] or [`ScanMode::Exclusive`]
433/// * `n_threads` – degree of parallelism (0 = auto)
434pub fn parallel_scan<T, F>(
435    data: &[T],
436    identity: T,
437    op: F,
438    mode: ScanMode,
439    n_threads: usize,
440) -> CoreResult<Vec<T>>
441where
442    T: Clone + Send + 'static,
443    F: Fn(T, T) -> T + Send + Sync + 'static,
444{
445    let n = data.len();
446    if n == 0 {
447        return Ok(Vec::new());
448    }
449
450    let n_threads = resolve_threads(n_threads).min(n);
451    let op = Arc::new(op);
452
453    // Step 1: compute local prefix sums per chunk.
454    let ranges = chunk_ranges(n, n_threads);
455    let n_chunks = ranges.len();
456    let data_ptr = data.as_ptr() as usize;
457
458    let chunk_sums: Arc<Mutex<Vec<(usize, T, Vec<T>)>>> =
459        Arc::new(Mutex::new(Vec::with_capacity(n_chunks)));
460    let mut handles = Vec::with_capacity(n_chunks);
461
462    for (chunk_id, range) in ranges.into_iter().enumerate() {
463        let op2 = Arc::clone(&op);
464        let cs = Arc::clone(&chunk_sums);
465        let id2 = identity.clone();
466
467        let handle = thread::Builder::new()
468            .spawn(move || {
469                let data: &[T] = unsafe { std::slice::from_raw_parts(data_ptr as *const T, n) };
470                let chunk = &data[range.clone()];
471                let mut local_prefix = Vec::with_capacity(range.len());
472                let mut acc = id2;
473                for x in chunk {
474                    acc = op2(acc, x.clone());
475                    local_prefix.push(acc.clone());
476                }
477                // The last element is the chunk's total.
478                let chunk_total = local_prefix.last().cloned().unwrap_or(acc);
479                if let Ok(mut g) = cs.lock() {
480                    g.push((chunk_id, chunk_total, local_prefix));
481                }
482            })
483            .map_err(spawn_err)?;
484        handles.push(handle);
485    }
486
487    for h in handles {
488        h.join().map_err(|_| join_err("parallel_scan local"))?;
489    }
490
491    let mut chunk_data = Arc::try_unwrap(chunk_sums)
492        .map_err(|_| CoreError::SchedulerError(ErrorContext::new("parallel_scan: Arc held")))?
493        .into_inner()
494        .map_err(|e| {
495            CoreError::SchedulerError(
496                ErrorContext::new(format!("parallel_scan: mutex poisoned: {e}"))
497                    .with_location(ErrorLocation::new(file!(), line!())),
498            )
499        })?;
500    chunk_data.sort_by_key(|(id, _, _)| *id);
501
502    // Step 2: compute global chunk offsets sequentially.
503    let mut offsets = Vec::with_capacity(n_chunks);
504    let mut running = identity.clone();
505    for (_, chunk_total, _) in &chunk_data {
506        offsets.push(running.clone());
507        running = op(running, chunk_total.clone());
508    }
509
510    // Step 3: apply offsets to local prefixes (always compute inclusive first).
511    let mut result = vec![identity.clone(); n];
512    let mut start = 0;
513    for (chunk_idx, (_, _, local_prefix)) in chunk_data.into_iter().enumerate() {
514        let offset = offsets[chunk_idx].clone();
515        let len = local_prefix.len();
516        for (j, lv) in local_prefix.into_iter().enumerate() {
517            result[start + j] = op(offset.clone(), lv);
518        }
519        start += len;
520    }
521
522    // For exclusive mode, shift right by 1 and prepend identity.
523    if mode == ScanMode::Exclusive {
524        let mut out = vec![identity; n];
525        out[1..n].clone_from_slice(&result[..(n - 1)]);
526        return Ok(out);
527    }
528
529    Ok(result)
530}
531
532/// Specialised parallel prefix sum for `f64` slices.
533///
534/// Returns a vector of length `n` where `result[i] = sum(input[0..=i])`.
535pub fn parallel_prefix_sum(data: &[f64], n_threads: usize) -> CoreResult<Vec<f64>> {
536    parallel_scan(data, 0.0f64, |a, b| a + b, ScanMode::Inclusive, n_threads)
537}
538
539// ── parallel_merge_sort ───────────────────────────────────────────────────────
540
541/// Parallel merge sort.
542///
543/// Splits the slice into chunks, sorts each chunk in a thread, then merges
544/// sequentially.  For small slices or `n_threads == 1` this degrades to
545/// `sort_unstable`.
546pub fn parallel_merge_sort<T>(data: &mut Vec<T>, n_threads: usize) -> CoreResult<()>
547where
548    T: Ord + Send + Clone + 'static,
549{
550    let n = data.len();
551    if n <= 1 {
552        return Ok(());
553    }
554
555    let n_threads = resolve_threads(n_threads).min(n);
556    if n_threads <= 1 {
557        data.sort_unstable();
558        return Ok(());
559    }
560
561    let ranges = chunk_ranges(n, n_threads);
562    // Split data into per-chunk owned vecs.
563    let mut chunks: Vec<Vec<T>> = {
564        let mut remaining = data.clone();
565        let mut out = Vec::with_capacity(ranges.len());
566        let mut offset = 0;
567        for range in &ranges {
568            let chunk: Vec<T> = remaining[offset..range.end - offset + offset].to_vec();
569            // Simpler: build chunks directly from data.
570            let _ = remaining; // avoid unused var warning
571            let chunk = data[range.clone()].to_vec();
572            out.push(chunk);
573            offset = range.end;
574        }
575        out
576    };
577
578    // Sort each chunk in a thread.
579    let sorted_chunks: Arc<Mutex<Vec<(usize, Vec<T>)>>> = Arc::new(Mutex::new(Vec::new()));
580    let mut handles = Vec::new();
581
582    for (id, mut chunk) in chunks.drain(..).enumerate() {
583        let sc = Arc::clone(&sorted_chunks);
584        let handle = thread::Builder::new()
585            .spawn(move || {
586                chunk.sort_unstable();
587                if let Ok(mut g) = sc.lock() {
588                    g.push((id, chunk));
589                }
590            })
591            .map_err(spawn_err)?;
592        handles.push(handle);
593    }
594
595    for h in handles {
596        h.join().map_err(|_| join_err("parallel_merge_sort"))?;
597    }
598
599    let mut sorted_chunks = Arc::try_unwrap(sorted_chunks)
600        .map_err(|_| CoreError::SchedulerError(ErrorContext::new("parallel_merge_sort: Arc held")))?
601        .into_inner()
602        .map_err(|e| {
603            CoreError::SchedulerError(
604                ErrorContext::new(format!("parallel_merge_sort: mutex poisoned: {e}"))
605                    .with_location(ErrorLocation::new(file!(), line!())),
606            )
607        })?;
608    sorted_chunks.sort_by_key(|(id, _)| *id);
609
610    // Sequential k-way merge.
611    let sorted: Vec<Vec<T>> = sorted_chunks.into_iter().map(|(_, v)| v).collect();
612    let merged = k_way_merge(sorted);
613    *data = merged;
614    Ok(())
615}
616
617/// K-way merge of sorted vectors into a single sorted vector.
618fn k_way_merge<T: Ord>(mut sorted: Vec<Vec<T>>) -> Vec<T> {
619    while sorted.len() > 1 {
620        let mut next = Vec::with_capacity(sorted.len() / 2 + 1);
621        let mut i = 0;
622        while i + 1 < sorted.len() {
623            let merged = merge_two(
624                std::mem::take(&mut sorted[i]),
625                std::mem::take(&mut sorted[i + 1]),
626            );
627            next.push(merged);
628            i += 2;
629        }
630        if i < sorted.len() {
631            next.push(std::mem::take(&mut sorted[i]));
632        }
633        sorted = next;
634    }
635    sorted.into_iter().next().unwrap_or_default()
636}
637
638/// Merge two sorted vectors.
639fn merge_two<T: Ord>(a: Vec<T>, b: Vec<T>) -> Vec<T> {
640    let mut result = Vec::with_capacity(a.len() + b.len());
641    let mut ai = a.into_iter();
642    let mut bi = b.into_iter();
643    let mut ahead = ai.next();
644    let mut bhead = bi.next();
645    loop {
646        match (ahead, bhead) {
647            (Some(av), Some(bv)) => {
648                if av <= bv {
649                    result.push(av);
650                    ahead = ai.next();
651                    bhead = Some(bv);
652                } else {
653                    result.push(bv);
654                    bhead = bi.next();
655                    ahead = Some(av);
656                }
657            }
658            (Some(av), None) => {
659                result.push(av);
660                result.extend(ai);
661                break;
662            }
663            (None, Some(bv)) => {
664                result.push(bv);
665                result.extend(bi);
666                break;
667            }
668            (None, None) => break,
669        }
670    }
671    result
672}
673
674// ── Tests ─────────────────────────────────────────────────────────────────────
675
676#[cfg(test)]
677mod tests {
678    use super::*;
679
680    #[test]
681    fn parallel_map_basic() {
682        let data: Vec<i32> = (1..=10).collect();
683        let result = parallel_map(&data, |&x| x * x, 0).expect("parallel_map");
684        assert_eq!(result, vec![1, 4, 9, 16, 25, 36, 49, 64, 81, 100]);
685    }
686
687    #[test]
688    fn parallel_map_empty() {
689        let data: Vec<i32> = Vec::new();
690        let result = parallel_map(&data, |&x| x * 2, 0).expect("map empty");
691        assert!(result.is_empty());
692    }
693
694    #[test]
695    fn parallel_map_single_thread() {
696        let data: Vec<u64> = (0..100).collect();
697        let result = parallel_map(&data, |&x| x + 1, 1).expect("single thread");
698        assert_eq!(result.len(), 100);
699        assert_eq!(result[99], 100);
700    }
701
702    #[test]
703    fn parallel_reduce_sum() {
704        let data: Vec<i64> = (1..=100).collect();
705        let sum =
706            parallel_reduce(&data, 0i64, |acc, &x| acc + x, |a, b| a + b, 4).expect("reduce sum");
707        assert_eq!(sum, 5050);
708    }
709
710    #[test]
711    fn parallel_reduce_empty() {
712        let data: Vec<i64> = Vec::new();
713        let sum = parallel_reduce(&data, 42i64, |acc, &x| acc + x, |a, b| a + b, 2)
714            .expect("reduce empty");
715        assert_eq!(sum, 42);
716    }
717
718    #[test]
719    fn parallel_filter_basic() {
720        let data: Vec<i32> = (1..=20).collect();
721        let evens = parallel_filter(data, |&x| x % 2 == 0, 4).expect("filter");
722        assert_eq!(evens, vec![2, 4, 6, 8, 10, 12, 14, 16, 18, 20]);
723    }
724
725    #[test]
726    fn parallel_filter_empty() {
727        let data: Vec<i32> = Vec::new();
728        let result = parallel_filter(data, |_| true, 2).expect("filter empty");
729        assert!(result.is_empty());
730    }
731
732    #[test]
733    fn parallel_scan_inclusive_sum() {
734        let data: Vec<i64> = (1..=8).collect();
735        let prefix =
736            parallel_scan(&data, 0i64, |a, b| a + b, ScanMode::Inclusive, 4).expect("scan inc");
737        assert_eq!(prefix, vec![1, 3, 6, 10, 15, 21, 28, 36]);
738    }
739
740    #[test]
741    fn parallel_scan_exclusive_sum() {
742        let data: Vec<i64> = (1..=5).collect();
743        let prefix =
744            parallel_scan(&data, 0i64, |a, b| a + b, ScanMode::Exclusive, 2).expect("scan exc");
745        assert_eq!(prefix, vec![0, 1, 3, 6, 10]);
746    }
747
748    #[test]
749    fn parallel_prefix_sum_basic() {
750        let data: Vec<f64> = (1..=5).map(|x| x as f64).collect();
751        let prefix = parallel_prefix_sum(&data, 2).expect("prefix sum");
752        let expected = [1.0, 3.0, 6.0, 10.0, 15.0];
753        for (a, b) in prefix.iter().zip(expected.iter()) {
754            assert!((a - b).abs() < 1e-10, "{a} vs {b}");
755        }
756    }
757
758    #[test]
759    fn parallel_merge_sort_basic() {
760        let mut data: Vec<i32> = vec![9, 3, 7, 1, 5, 2, 8, 4, 6, 0];
761        parallel_merge_sort(&mut data, 4).expect("merge sort");
762        assert_eq!(data, vec![0, 1, 2, 3, 4, 5, 6, 7, 8, 9]);
763    }
764
765    #[test]
766    fn parallel_merge_sort_single_element() {
767        let mut data = vec![42i32];
768        parallel_merge_sort(&mut data, 4).expect("sort single");
769        assert_eq!(data, vec![42]);
770    }
771
772    #[test]
773    fn parallel_merge_sort_already_sorted() {
774        let mut data: Vec<i32> = (0..50).collect();
775        parallel_merge_sort(&mut data, 4).expect("sort sorted");
776        assert_eq!(data, (0..50).collect::<Vec<_>>());
777    }
778
779    #[test]
780    fn parallel_partition_basic() {
781        let data: Vec<i32> = (1..=10).collect();
782        let (evens, odds) = parallel_partition(data, |&x| x % 2 == 0, 4).expect("partition");
783        assert_eq!(evens, vec![2, 4, 6, 8, 10]);
784        assert_eq!(odds, vec![1, 3, 5, 7, 9]);
785    }
786
787    #[test]
788    fn parallel_for_each_basic() {
789        use std::sync::atomic::{AtomicI64, Ordering};
790        let data: Vec<i64> = (1..=100).collect();
791        let sum = Arc::new(AtomicI64::new(0));
792        let s = Arc::clone(&sum);
793        parallel_for_each(
794            &data,
795            move |&x| {
796                s.fetch_add(x, Ordering::Relaxed);
797            },
798            4,
799        )
800        .expect("for_each");
801        assert_eq!(sum.load(Ordering::Relaxed), 5050);
802    }
803}