#[cfg(feature = "parallel")]
use crate::AnamnesisError;
pub(crate) const MIN_PARALLEL_BYTES: u64 = 4 * 1024 * 1024;
pub(crate) fn map_indexed<T, R, F, P>(
items: &[T],
threads: usize,
work_bytes: u64,
f: F,
mut on_result: P,
) -> crate::Result<Vec<R>>
where
T: Sync,
R: Send,
F: Fn(usize, &T) -> crate::Result<R> + Sync,
P: FnMut(&R),
{
let worth_spawning = threads > 1 && items.len() > 1 && work_bytes >= MIN_PARALLEL_BYTES;
#[cfg(feature = "parallel")]
if worth_spawning {
return map_indexed_parallel(items, threads, &f, &mut on_result);
}
#[cfg(not(feature = "parallel"))]
let _ = worth_spawning;
let mut out = Vec::with_capacity(items.len());
for (idx, item) in items.iter().enumerate() {
let result = f(idx, item)?;
on_result(&result);
out.push(result);
}
Ok(out)
}
#[cfg(feature = "parallel")]
fn map_indexed_parallel<T, R, F, P>(
items: &[T],
threads: usize,
f: &F,
on_result: &mut P,
) -> crate::Result<Vec<R>>
where
T: Sync,
R: Send,
F: Fn(usize, &T) -> crate::Result<R> + Sync,
P: FnMut(&R),
{
use std::sync::atomic::{AtomicUsize, Ordering};
let n_workers = threads.min(items.len());
let cursor = AtomicUsize::new(0);
let mut collected: Vec<(usize, R)> = Vec::with_capacity(items.len());
let mut failure: Option<(usize, AnamnesisError)> = None;
let mut panicked = false;
std::thread::scope(|scope| {
let handles: Vec<_> = (0..n_workers)
.map(|_| {
let cursor = &cursor;
scope.spawn(move || {
let mut local: Vec<(usize, R)> = Vec::new();
loop {
let idx = cursor.fetch_add(1, Ordering::Relaxed);
let Some(item) = items.get(idx) else { break };
match f(idx, item) {
Ok(result) => local.push((idx, result)),
Err(err) => return Err((idx, err)),
}
}
Ok(local)
})
})
.collect();
for handle in handles {
let joined = handle.join();
match joined {
Ok(Ok(local)) => {
for (idx, result) in local {
on_result(&result);
collected.push((idx, result));
}
}
Ok(Err((idx, err))) => {
if failure.as_ref().is_none_or(|&(seen, _)| idx < seen) {
failure = Some((idx, err));
}
}
Err(_) => panicked = true,
}
}
});
if panicked {
return Err(AnamnesisError::Parse {
reason: "parallel dequant worker thread panicked".into(),
});
}
if let Some((_, err)) = failure {
return Err(err);
}
collected.sort_by_key(|&(idx, _)| idx);
Ok(collected.into_iter().map(|(_, result)| result).collect())
}
#[cfg(test)]
#[allow(
clippy::panic,
clippy::unwrap_used,
clippy::expect_used,
clippy::indexing_slicing
)]
mod tests {
use super::{MIN_PARALLEL_BYTES, map_indexed};
use crate::AnamnesisError;
const BIG: u64 = MIN_PARALLEL_BYTES * 2;
#[test]
fn results_are_in_input_order_for_every_budget() {
let items: Vec<usize> = (0..97).collect();
let baseline =
map_indexed(&items, 1, BIG, |_, &v| Ok(v * 3), |_| {}).expect("sequential map");
assert_eq!(baseline, items.iter().map(|v| v * 3).collect::<Vec<_>>());
for threads in [1usize, 2, 4, 8, 16] {
let out =
map_indexed(&items, threads, BIG, |_, &v| Ok(v * 3), |_| {}).expect("parallel map");
assert_eq!(out, baseline, "order must not depend on thread count");
}
}
#[test]
fn closure_receives_the_input_index() {
let items: Vec<usize> = (0..64).map(|i| i * 10).collect();
let out = map_indexed(&items, 8, BIG, |idx, &v| Ok((idx, v)), |_| {}).expect("map");
for (expected, &(idx, v)) in out.iter().enumerate() {
assert_eq!(idx, expected);
assert_eq!(v, expected * 10);
}
}
#[test]
fn on_result_fires_once_per_item() {
let items: Vec<usize> = (0..50).collect();
for threads in [1usize, 4] {
let mut seen = 0usize;
let out = map_indexed(&items, threads, BIG, |_, &v| Ok(v), |_| seen += 1)
.expect("map with callback");
assert_eq!(seen, out.len());
assert_eq!(seen, items.len());
}
}
#[test]
fn error_selection_is_deterministic_across_budgets() {
let items: Vec<usize> = (0..128).collect();
let failing = |_: usize, &v: &usize| -> crate::Result<usize> {
if v == 37 || v == 80 || v == 127 {
Err(AnamnesisError::Parse {
reason: format!("item {v} rejected"),
})
} else {
Ok(v)
}
};
for threads in [1usize, 2, 4, 8, 16] {
let err =
map_indexed(&items, threads, BIG, failing, |_| {}).expect_err("the map must fail");
assert_eq!(
err.to_string(),
AnamnesisError::Parse {
reason: "item 37 rejected".into(),
}
.to_string(),
"the lowest failing index must win at {threads} threads"
);
}
}
#[test]
fn degenerate_inputs() {
let empty: Vec<usize> = Vec::new();
let out = map_indexed(&empty, 8, BIG, |_, &v| Ok(v), |_| {}).expect("empty map");
assert!(out.is_empty());
let one = vec![42usize];
let out = map_indexed(&one, 8, BIG, |_, &v| Ok(v + 1), |_| {}).expect("single map");
assert_eq!(out, vec![43]);
}
#[test]
fn below_the_size_threshold_still_maps_correctly() {
let items: Vec<usize> = (0..40).collect();
let out = map_indexed(&items, 8, 1024, |_, &v| Ok(v * 2), |_| {}).expect("small map");
assert_eq!(out, items.iter().map(|v| v * 2).collect::<Vec<_>>());
}
}