use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Mutex, OnceLock};
use crate::graph::session::QUERY_THREAD_STACK_SIZE;
pub(crate) const QUERY_THREADS_ENV: &str = "KGLITE_QUERY_THREADS";
pub(crate) const PARALLEL_POLL_INTERVAL: usize = 4096;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum CostClass {
Compiled,
Interpreted,
}
pub(crate) const PARALLEL_MIN_ROWS_COMPILED: usize = 20_000;
pub(crate) const PARALLEL_MIN_ROWS_INTERPRETED: usize = 5_000;
#[inline]
pub(crate) fn should_fan_out(rows: usize, cost: CostClass) -> bool {
rows >= match cost {
CostClass::Compiled => PARALLEL_MIN_ROWS_COMPILED,
CostClass::Interpreted => PARALLEL_MIN_ROWS_INTERPRETED,
}
}
pub(crate) const PROJECTION_MIN_ROWS: usize = 4096;
static QUERY_POOL: OnceLock<Option<rayon::ThreadPool>> = OnceLock::new();
fn configured_width() -> usize {
if let Some(raw) = std::env::var_os(QUERY_THREADS_ENV) {
if let Some(n) = raw
.to_str()
.and_then(|s| s.trim().parse::<usize>().ok())
.filter(|n| *n > 0)
{
return n;
}
}
std::thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(1)
}
fn pool() -> Option<&'static rayon::ThreadPool> {
QUERY_POOL
.get_or_init(|| {
rayon::ThreadPoolBuilder::new()
.num_threads(configured_width())
.stack_size(QUERY_THREAD_STACK_SIZE)
.thread_name(|i| format!("kglite-query-{i}"))
.build()
.ok()
})
.as_ref()
}
pub(crate) fn install<OP, R>(op: OP) -> R
where
OP: FnOnce() -> R + Send,
R: Send,
{
match pool() {
Some(p) => p.install(op),
None => op(),
}
}
#[cfg(test)]
pub(crate) static PARALLEL_SCANS: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(0);
#[cfg(test)]
pub(crate) fn parallel_scans() -> usize {
PARALLEL_SCANS.load(Ordering::Relaxed)
}
#[cfg(test)]
pub(crate) static PARALLEL_CANDIDATE_SCANS: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(0);
#[cfg(test)]
pub(crate) fn parallel_candidate_scans() -> usize {
PARALLEL_CANDIDATE_SCANS.load(Ordering::Relaxed)
}
#[cfg(test)]
pub(crate) static PARALLEL_AGGREGATIONS: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(0);
#[cfg(test)]
pub(crate) fn parallel_aggregations() -> usize {
PARALLEL_AGGREGATIONS.load(Ordering::Relaxed)
}
#[cfg(test)]
pub(crate) static PARALLEL_SORT_KEYS: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(0);
#[cfg(test)]
pub(crate) fn parallel_sort_keys() -> usize {
PARALLEL_SORT_KEYS.load(Ordering::Relaxed)
}
const UNKNOWN_REASON: &str = "parallel region failed";
pub(crate) struct ParallelInterrupt<F> {
had_error: AtomicBool,
first_error: Mutex<Option<String>>,
probe: F,
}
impl<F> ParallelInterrupt<F>
where
F: Fn() -> Option<String> + Sync,
{
pub(crate) fn new(probe: F) -> Self {
ParallelInterrupt {
had_error: AtomicBool::new(false),
first_error: Mutex::new(None),
probe,
}
}
#[inline]
pub(crate) fn check(&self, index: usize) -> Result<(), String> {
if index & (PARALLEL_POLL_INTERVAL - 1) == 0 {
return self.check_each();
}
if self.had_error.load(Ordering::Relaxed) {
return Err(self.recorded());
}
Ok(())
}
#[inline]
pub(crate) fn check_each(&self) -> Result<(), String> {
if self.had_error.load(Ordering::Relaxed) {
return Err(self.recorded());
}
if let Some(reason) = (self.probe)() {
self.fail(reason.clone());
return Err(reason);
}
Ok(())
}
pub(crate) fn fail(&self, reason: String) {
let mut slot = self.first_error.lock().unwrap();
if !self.had_error.swap(true, Ordering::Relaxed) {
*slot = Some(reason);
}
}
#[inline]
pub(crate) fn capture<T>(&self, outcome: Result<T, String>) -> Option<T> {
match outcome {
Ok(value) => Some(value),
Err(reason) => {
self.fail(reason);
None
}
}
}
fn recorded(&self) -> String {
self.first_error
.lock()
.unwrap()
.clone()
.unwrap_or_else(|| UNKNOWN_REASON.to_string())
}
pub(crate) fn finish(self) -> Result<(), String> {
if self.had_error.load(Ordering::Relaxed) {
return Err(self
.first_error
.into_inner()
.unwrap()
.unwrap_or_else(|| UNKNOWN_REASON.to_string()));
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use rayon::prelude::*;
#[test]
fn poll_interval_is_a_power_of_two() {
assert!(PARALLEL_POLL_INTERVAL.is_power_of_two());
}
#[test]
fn pool_workers_get_the_query_thread_stack_size() {
let width = install(rayon::current_num_threads);
assert!(width >= 1);
let name = install(|| {
(0..1024usize)
.into_par_iter()
.map(|_| {
std::thread::current()
.name()
.unwrap_or("<unnamed>")
.to_string()
})
.find_any(|n| n.starts_with("kglite-query-"))
});
assert!(
name.is_some(),
"parallel work inside install() did not run on a kglite-query worker"
);
}
#[test]
fn first_reason_wins_and_later_workers_see_the_latch() {
let guard = ParallelInterrupt::new(|| None);
guard.fail("first".to_string());
guard.fail("second".to_string());
assert_eq!(guard.check_each(), Err("first".to_string()));
assert_eq!(guard.finish(), Err("first".to_string()));
}
#[test]
fn check_polls_only_on_chunk_boundaries() {
let polls = AtomicBool::new(false);
let guard = ParallelInterrupt::new(|| {
polls.store(true, Ordering::Relaxed);
None
});
assert_eq!(guard.check(1), Ok(()));
assert!(!polls.load(Ordering::Relaxed), "off-boundary index polled");
assert_eq!(guard.check(PARALLEL_POLL_INTERVAL), Ok(()));
assert!(polls.load(Ordering::Relaxed), "boundary index did not poll");
}
#[test]
fn capture_latches_a_unit_error() {
let guard = ParallelInterrupt::new(|| None);
assert_eq!(guard.capture(Ok(7)), Some(7));
assert_eq!(guard.capture::<i32>(Err("boom".into())), None);
assert_eq!(guard.finish(), Err("boom".to_string()));
}
}