1use std::sync::{Arc, Mutex};
37use std::thread;
38
39use crate::error::{CoreError, CoreResult, ErrorContext, ErrorLocation};
40
41pub 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
54fn 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
86pub 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 let mut out: Vec<R> = vec![R::default(); n];
107
108 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 unsafe { std::slice::from_raw_parts(data_ptr as *const T, n) };
123 let out: &mut [R] =
124 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
141pub 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
179pub 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 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
252pub 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 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 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
322pub 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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
417pub enum ScanMode {
418 Inclusive,
420 Exclusive,
422}
423
424pub 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 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 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 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 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 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
532pub 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
539pub 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 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 let _ = remaining; let chunk = data[range.clone()].to_vec();
572 out.push(chunk);
573 offset = range.end;
574 }
575 out
576 };
577
578 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 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
617fn 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
638fn 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#[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}