use crate::error::SearchError;
use crate::hit::SearchHit;
use std::sync::atomic::{AtomicUsize, Ordering};
pub(crate) fn search_batch<F>(
queries: &[&[f32]],
workers: usize,
search: F,
) -> Result<Vec<Vec<SearchHit>>, SearchError>
where
F: Fn(&[f32]) -> Result<Vec<SearchHit>, SearchError> + Sync,
{
if queries.is_empty() {
return Ok(Vec::new());
}
let workers = workers.min(queries.len()).max(1);
let next = AtomicUsize::new(0);
let mut slots = Vec::new();
slots
.try_reserve_exact(queries.len())
.map_err(|_| SearchError::AllocationFailed)?;
slots.extend(std::iter::repeat_with(|| None).take(queries.len()));
let panicked = std::thread::scope(|scope| {
let handles = (0..workers)
.map(|_| {
let next = &next;
let search = &search;
scope.spawn(move || {
let mut local = Vec::new();
loop {
let index = next.fetch_add(1, Ordering::Relaxed);
let Some(query) = queries.get(index) else {
break;
};
local.push((index, search(query)));
}
local
})
})
.collect::<Vec<_>>();
let mut panicked = false;
for handle in handles {
match handle.join() {
Ok(local) => {
for (index, result) in local {
slots[index] = Some(result);
}
}
Err(_) => panicked = true,
}
}
panicked
});
if panicked {
return Err(SearchError::WorkerPanic);
}
slots
.into_iter()
.map(|slot| slot.ok_or(SearchError::WorkerPanic)?)
.collect()
}