use std::mem::size_of;
use std::sync::Mutex;
use std::sync::atomic::{AtomicUsize, Ordering};
use rudb_common::{Error, Memory, Reservation, Result, Spent, Stage, stage};
use crate::key::{mix, spread};
pub(crate) const PARTITIONS: usize = 16;
const EMPTY: u32 = u32::MAX;
const NOTHING: u64 = 0x9e37_79b9_7f4a_7c15;
#[derive(Debug, Clone, Copy)]
pub(crate) struct Record {
pub(crate) user: i64,
pub(crate) group: i32,
pub(crate) pair_hash: u32,
}
pub(crate) fn group_hash(group: i32, valid: bool) -> u32 {
let word = if valid { i64::from(group) as u64 } else { NOTHING };
let wide = spread(mix(0, word));
(wide ^ (wide >> 32)) as u32
}
#[derive(Debug, Default)]
pub(crate) struct Run {
pub(crate) rows: Vec<Record>,
pub(crate) validity: Vec<bool>,
}
impl Run {
pub(crate) fn push(&mut self, row: Record, valid: bool) {
self.rows.push(row);
if self.validity.is_empty() {
if valid {
return;
}
self.validity = vec![true; self.rows.len() - 1];
}
self.validity.push(valid);
}
fn valid_at(&self, row: usize) -> bool {
self.validity.is_empty() || self.validity[row]
}
pub(crate) fn footprint(&self) -> usize {
self.rows.capacity() * size_of::<Record>() + self.validity.capacity() * size_of::<bool>()
}
}
#[derive(Debug, Default)]
pub(crate) struct Held {
pub(crate) runs: Vec<Run>,
}
impl Held {
pub(crate) fn rows(&self) -> usize {
self.runs.iter().map(|run| run.rows.len()).sum()
}
}
#[inline]
pub(crate) fn scatter(partitions: &mut [Run], shift: u32, group: i32, valid: bool, user: i64) {
let group_word = if valid { i64::from(group) as u64 } else { NOTHING };
let wide_pair = spread(mix(spread(mix(0, group_word)), user as u64));
let pair_hash = (wide_pair ^ (wide_pair >> 32)) as u32;
partitions[(pair_hash >> shift) as usize].push(Record { user, group, pair_hash }, valid);
}
pub(crate) fn shift() -> u32 {
u32::BITS - PARTITIONS.ilog2()
}
pub(crate) fn in_parallel<T: Send>(
count: usize,
degree: usize,
what: &str,
run: impl Fn(usize) -> Result<T> + Sync,
) -> Result<Vec<T>> {
let next = AtomicUsize::new(0);
let slots: Vec<Mutex<Option<Result<T>>>> = (0..count).map(|_| Mutex::new(None)).collect();
let step = || {
loop {
let at = next.fetch_add(1, Ordering::Relaxed);
if at >= count {
return;
}
let done = run(at);
if let Ok(mut slot) = slots[at].lock() {
*slot = Some(done);
}
}
};
std::thread::scope(|scope| {
let degree = degree.min(count);
let mut handles = Vec::with_capacity(degree - 1);
for _ in 1..degree {
handles.push(scope.spawn(|| {
step();
stage::here()
}));
}
step();
let mut theirs = Spent::none();
for handle in handles {
let spent =
handle.join().map_err(|_| Error::internal("a radix pair worker panicked"))?;
theirs.add(spent);
}
stage::gained(theirs);
let mut out = Vec::with_capacity(count);
for (at, slot) in slots.iter().enumerate() {
out.push(
slot.lock()
.map_err(poisoned)?
.take()
.unwrap_or_else(|| Err(Error::internal(format!("nothing {what} {at}"))))?,
);
}
Ok::<_, Error>(out)
})
}
#[derive(Debug)]
pub(crate) struct Counted {
pub(crate) splits: Vec<Vec<Grouped>>,
pub(crate) held: Reservation,
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct Grouped {
pub(crate) group: i32,
pub(crate) group_hash: u32,
pub(crate) valid: bool,
}
#[inline]
pub(crate) fn split_of(group_hash: u32, splits: usize) -> usize {
((u64::from(group_hash) * splits as u64) >> u32::BITS) as usize
}
pub(crate) fn distinct_pairs(
partition: &mut Held,
splits: usize,
memory: &Memory,
) -> Result<Counted> {
let held_rows = partition.rows();
let pair_capacity = held_rows.saturating_mul(2).max(64).next_power_of_two();
let mut working = memory.reservation();
working.grow(width(pair_capacity * size_of::<u32>()))?;
let mut pair_buckets = vec![EMPTY; pair_capacity];
let pair_mask = pair_capacity - 1;
let all_valid = partition.runs.iter().all(|run| run.validity.is_empty());
let mut unique: Vec<Record> = Vec::new();
let mut unique_validity: Vec<bool> = Vec::new();
let timing = stage::Timing::start(Stage::Fold);
for run in &partition.runs {
for (source, &row) in run.rows.iter().enumerate() {
let valid = all_valid || run.valid_at(source);
let mut at = row.pair_hash as usize & pair_mask;
loop {
let slot = pair_buckets[at];
if slot == EMPTY {
pair_buckets[at] = u32::try_from(unique.len())
.map_err(|_| Error::out_of_memory("a radix pair partition is too large"))?;
unique.push(row);
if !all_valid {
unique_validity.push(valid);
}
break;
}
let slot = slot as usize;
let held = unique[slot];
let held_valid = all_valid || unique_validity[slot];
if held.pair_hash == row.pair_hash
&& held.group == row.group
&& held.user == row.user
&& held_valid == valid
{
break;
}
at = (at + 1) & pair_mask;
}
}
}
let pairs = unique.len();
partition.runs.clear();
working.grow(width(
unique.capacity() * size_of::<Record>() + unique_validity.capacity() * size_of::<bool>(),
))?;
let partition = Run { rows: unique, validity: unique_validity };
let even = pairs.div_ceil(splits);
let share = (even + even.isqrt() * 4).min(pairs);
let mut parts: Vec<Vec<Grouped>> = (0..splits).map(|_| Vec::with_capacity(share)).collect();
for row in 0..pairs {
let record = partition.rows[row];
let valid = all_valid || partition.validity[row];
let group_hash = group_hash(record.group, valid);
parts[split_of(group_hash, splits)].push(Grouped {
group: record.group,
group_hash,
valid,
});
}
let mut held = memory.reservation();
held.grow(width(
parts.iter().map(|split| split.capacity() * size_of::<Grouped>()).sum::<usize>(),
))?;
timing.stop(0);
Ok(Counted { splits: parts, held })
}
fn width(value: usize) -> u64 {
u64::try_from(value).unwrap_or(u64::MAX)
}
fn poisoned<T>(_: T) -> Error {
Error::internal("a radix pair lock was poisoned")
}
#[cfg(test)]
mod tests {
use super::{Held, Record, Run, distinct_pairs};
#[test]
fn a_run_that_opens_with_an_invalid_key_keeps_its_validity() {
let mut run = Run::default();
run.push(Record { user: 10, group: 0, pair_hash: 5 }, false);
run.push(Record { user: 11, group: 4, pair_hash: 5 }, true);
assert_eq!(run.validity, [false, true]);
assert!(!run.valid_at(0));
assert!(run.valid_at(1));
}
#[test]
fn a_null_key_does_not_join_the_group_whose_key_is_zero() {
let mut run = Run::default();
run.push(Record { user: 10, group: 0, pair_hash: 5 }, false);
run.push(Record { user: 10, group: 0, pair_hash: 5 }, true);
let mut partition = Held { runs: vec![run] };
let counted = distinct_pairs(&mut partition, 1, &rudb_common::Memory::unlimited())
.expect("a pair partition");
let mut found: Vec<bool> = counted.splits[0].iter().map(|pair| pair.valid).collect();
found.sort_unstable();
assert_eq!(found, [false, true]);
}
}