use std::hash::{Hash, Hasher};
use std::sync::{Arc, Mutex, Weak};
use rustc_hash::FxHasher;
use radixdb_core::CompactArc;
use radixdb_core::{Row, Value};
pub trait JoinHashObserver {
fn insert_raw_hash(&mut self, hash: u64);
fn retained_bytes(&self) -> usize;
}
#[derive(Clone)]
pub struct JoinHashState {
build_rows: CompactArc<Vec<Row>>,
key_indices: CompactArc<[usize]>,
table: Arc<JoinHashTable>,
_memory_reservation: Option<Arc<JoinMemoryReservation>>,
}
#[derive(Default)]
struct JoinMemoryUsage {
retained_bytes: usize,
peak_bytes: usize,
}
#[derive(Default)]
pub(crate) struct JoinMemoryOwner {
usage: Mutex<JoinMemoryUsage>,
}
impl std::fmt::Debug for JoinMemoryOwner {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("JoinMemoryOwner")
.field("retained_bytes", &self.retained_bytes())
.field("peak_bytes", &self.peak_bytes())
.finish()
}
}
impl JoinMemoryOwner {
pub(crate) fn try_reserve(
self: &Arc<Self>,
bytes: usize,
max_bytes: usize,
) -> Option<JoinMemoryReservation> {
let mut usage = self
.usage
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let next = usage.retained_bytes.checked_add(bytes)?;
if next > max_bytes {
return None;
}
usage.retained_bytes = next;
usage.peak_bytes = usage.peak_bytes.max(next);
Some(JoinMemoryReservation {
owner: Arc::downgrade(self),
bytes,
})
}
pub(crate) fn retained_bytes(&self) -> usize {
self.usage
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.retained_bytes
}
pub(crate) fn peak_bytes(&self) -> usize {
self.usage
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.peak_bytes
}
}
#[doc(hidden)]
pub struct JoinMemoryReservation {
owner: Weak<JoinMemoryOwner>,
bytes: usize,
}
impl JoinMemoryReservation {
#[doc(hidden)]
pub fn try_resize(&mut self, bytes: usize, max_bytes: usize) -> bool {
let Some(owner) = self.owner.upgrade() else {
self.bytes = bytes;
return true;
};
let mut usage = owner
.usage
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let without_self = usage.retained_bytes.saturating_sub(self.bytes);
let Some(next) = without_self.checked_add(bytes) else {
return false;
};
if next > max_bytes {
return false;
}
usage.retained_bytes = next;
usage.peak_bytes = usage.peak_bytes.max(next);
self.bytes = bytes;
true
}
}
impl std::fmt::Debug for JoinMemoryReservation {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("JoinMemoryReservation")
.field("bytes", &self.bytes)
.finish()
}
}
impl Drop for JoinMemoryReservation {
fn drop(&mut self) {
let Some(owner) = self.owner.upgrade() else {
return;
};
let mut usage = owner
.usage
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
usage.retained_bytes = usage.retained_bytes.saturating_sub(self.bytes);
}
}
impl JoinHashState {
pub fn build(build_rows: CompactArc<Vec<Row>>, key_indices: &[usize]) -> Self {
let table = JoinHashTable::build(&build_rows, key_indices);
Self::from_table(build_rows, key_indices, table, None)
}
pub fn try_build(
build_rows: CompactArc<Vec<Row>>,
key_indices: &[usize],
max_bytes: usize,
) -> Option<Self> {
JoinHashTable::fits_retained_budget(build_rows.len(), max_bytes)
.then(|| Self::build(build_rows, key_indices))
}
pub(crate) fn build_reserved(
build_rows: CompactArc<Vec<Row>>,
key_indices: &[usize],
reservation: JoinMemoryReservation,
) -> Self {
let table = JoinHashTable::build(&build_rows, key_indices);
Self::from_table(build_rows, key_indices, table, Some(reservation))
}
pub fn build_with_bloom(
build_rows: CompactArc<Vec<Row>>,
key_indices: &[usize],
observer: &mut impl JoinHashObserver,
) -> Self {
let table = JoinHashTable::build_with_observer(&build_rows, key_indices, observer);
Self::from_table(build_rows, key_indices, table, None)
}
pub(crate) fn build_with_bloom_reserved(
build_rows: CompactArc<Vec<Row>>,
key_indices: &[usize],
observer: &mut impl JoinHashObserver,
reservation: JoinMemoryReservation,
) -> Self {
let table = JoinHashTable::build_with_observer(&build_rows, key_indices, observer);
Self::from_table(build_rows, key_indices, table, Some(reservation))
}
fn from_table(
build_rows: CompactArc<Vec<Row>>,
key_indices: &[usize],
table: JoinHashTable,
reservation: Option<JoinMemoryReservation>,
) -> Self {
assert_eq!(
table.len(),
build_rows.len(),
"join hash state must contain one entry per build row"
);
Self {
build_rows,
key_indices: CompactArc::from(key_indices.to_vec()),
table: Arc::new(table),
_memory_reservation: reservation.map(Arc::new),
}
}
pub fn matches(&self, build_rows: &CompactArc<Vec<Row>>, key_indices: &[usize]) -> bool {
CompactArc::ptr_eq(&self.build_rows, build_rows) && self.key_indices.as_ref() == key_indices
}
#[inline]
pub fn build_rows(&self) -> &CompactArc<Vec<Row>> {
&self.build_rows
}
#[inline]
pub fn table(&self) -> &Arc<JoinHashTable> {
&self.table
}
#[cfg(test)]
fn shares_allocations_with(&self, other: &Self) -> bool {
CompactArc::ptr_eq(&self.build_rows, &other.build_rows)
&& Arc::ptr_eq(&self.table, &other.table)
}
}
const EMPTY: u32 = u32::MAX;
const MIN_BUCKETS: usize = 16;
#[doc(hidden)]
pub const DEFAULT_JOIN_HASH_STATE_MAX_BYTES: usize = 256 * 1024 * 1024;
#[repr(C)]
#[derive(Debug, Clone, Copy)]
struct HashEntry {
hash: u64,
row_idx: u32,
next: u32,
}
impl HashEntry {
#[inline]
fn new(hash: u64, row_idx: u32, next: u32) -> Self {
Self {
hash,
row_idx,
next,
}
}
}
pub struct JoinHashTable {
bucket_heads: Vec<i32>,
entries: Vec<HashEntry>,
bucket_mask: u64,
len: usize,
}
impl JoinHashTable {
fn bucket_count_for_rows(row_count: usize) -> Option<usize> {
row_count
.checked_mul(4)
.map(|scaled| scaled / 3)
.map(|scaled| scaled.max(MIN_BUCKETS))?
.checked_next_power_of_two()
}
pub fn estimated_retained_bytes(row_count: usize) -> Option<usize> {
if row_count > u32::MAX as usize {
return None;
}
let buckets = Self::bucket_count_for_rows(row_count)?;
buckets
.checked_mul(std::mem::size_of::<i32>())?
.checked_add(row_count.checked_mul(std::mem::size_of::<HashEntry>())?)
}
#[inline]
pub fn fits_retained_budget(row_count: usize, max_bytes: usize) -> bool {
Self::estimated_retained_bytes(row_count).is_some_and(|bytes| bytes <= max_bytes)
}
pub fn with_capacity(row_count: usize) -> Self {
let bucket_count = Self::bucket_count_for_rows(row_count)
.expect("join hash table capacity exceeds addressable range");
assert!(
row_count <= u32::MAX as usize,
"join hash table row index exceeds u32"
);
let bucket_mask = (bucket_count - 1) as u64;
Self {
bucket_heads: vec![-1; bucket_count],
entries: Vec::with_capacity(row_count),
bucket_mask,
len: 0,
}
}
pub fn empty() -> Self {
Self {
bucket_heads: vec![-1; MIN_BUCKETS],
entries: Vec::new(),
bucket_mask: (MIN_BUCKETS - 1) as u64,
len: 0,
}
}
pub fn build(rows: &[Row], key_indices: &[usize]) -> Self {
#[cfg(feature = "bench-harness")]
let started = std::time::Instant::now();
let mut table = Self::with_capacity(rows.len());
for (idx, row) in rows.iter().enumerate() {
let hash = hash_row_keys(row, key_indices);
table.insert(hash, idx as u32);
}
#[cfg(feature = "bench-harness")]
radixdb_storage::instrumentation::record_hash_build(rows.len(), started.elapsed());
table
}
pub fn build_with_observer(
rows: &[Row],
key_indices: &[usize],
observer: &mut impl JoinHashObserver,
) -> Self {
#[cfg(feature = "bench-harness")]
let started = std::time::Instant::now();
let mut table = Self::with_capacity(rows.len());
for (idx, row) in rows.iter().enumerate() {
let hash = hash_row_keys(row, key_indices);
table.insert(hash, idx as u32);
observer.insert_raw_hash(hash);
}
#[cfg(feature = "bench-harness")]
radixdb_storage::instrumentation::record_hash_build(rows.len(), started.elapsed());
table
}
#[inline]
pub fn insert(&mut self, hash: u64, row_idx: u32) {
let bucket = (hash & self.bucket_mask) as usize;
let old_head = self.bucket_heads[bucket];
let entry_idx = self.len as u32;
let next = if old_head >= 0 {
old_head as u32
} else {
EMPTY
};
self.entries.push(HashEntry::new(hash, row_idx, next));
self.bucket_heads[bucket] = entry_idx as i32;
self.len += 1;
}
#[inline]
pub fn probe(&self, hash: u64) -> ProbeIter<'_> {
let bucket = (hash & self.bucket_mask) as usize;
let first = self.bucket_heads[bucket];
ProbeIter {
table: self,
hash,
current: first,
}
}
#[inline]
pub fn probe_cursor(&self, hash: u64) -> ProbeCursor {
let bucket = (hash & self.bucket_mask) as usize;
ProbeCursor {
hash,
current: self.bucket_heads[bucket],
}
}
#[inline]
pub fn probe_next(&self, cursor: &mut ProbeCursor) -> Option<usize> {
while cursor.current >= 0 {
let entry = &self.entries[cursor.current as usize];
cursor.current = if entry.next == EMPTY {
-1
} else {
entry.next as i32
};
if entry.hash == cursor.hash {
return Some(entry.row_idx as usize);
}
}
None
}
#[inline]
pub fn len(&self) -> usize {
self.len
}
#[inline]
pub fn is_empty(&self) -> bool {
self.len == 0
}
#[inline]
pub fn bucket_count(&self) -> usize {
self.bucket_heads.len()
}
#[inline]
pub fn load_factor(&self) -> f64 {
self.len as f64 / self.bucket_heads.len() as f64
}
}
pub struct ProbeIter<'a> {
table: &'a JoinHashTable,
hash: u64,
current: i32,
}
#[derive(Debug, Clone, Copy, Default)]
pub struct ProbeCursor {
hash: u64,
current: i32,
}
impl Iterator for ProbeIter<'_> {
type Item = usize;
#[inline]
fn next(&mut self) -> Option<usize> {
while self.current >= 0 {
let entry = &self.table.entries[self.current as usize];
self.current = if entry.next == EMPTY {
-1
} else {
entry.next as i32
};
if entry.hash == self.hash {
return Some(entry.row_idx as usize);
}
}
None
}
}
#[inline]
pub fn hash_keys_with<'a, F>(key_indices: &[usize], get_value: F) -> u64
where
F: Fn(usize) -> Option<&'a Value>,
{
let mut hasher = FxHasher::default();
for &idx in key_indices {
if let Some(value) = get_value(idx) {
hash_value(&mut hasher, value);
} else {
0xDEADBEEF_u64.hash(&mut hasher);
}
}
hasher.finish()
}
#[inline]
pub fn hash_row_keys(row: &Row, key_indices: &[usize]) -> u64 {
let mut hasher = FxHasher::default();
for &idx in key_indices {
if let Some(value) = row.get(idx) {
hash_value(&mut hasher, value);
} else {
0xDEADBEEF_u64.hash(&mut hasher);
}
}
hasher.finish()
}
#[inline]
fn hash_value<H: Hasher>(hasher: &mut H, value: &Value) {
value.hash(hasher);
}
#[inline]
pub fn verify_key_equality(row1: &Row, row2: &Row, indices1: &[usize], indices2: &[usize]) -> bool {
debug_assert_eq!(indices1.len(), indices2.len());
for (&idx1, &idx2) in indices1.iter().zip(indices2.iter()) {
let (Some(value1), Some(value2)) = (row1.get(idx1), row2.get(idx2)) else {
return false;
};
if value1.is_null() || value2.is_null() || value1 != value2 {
return false;
}
}
true
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn hash_entry_retains_only_hash_and_compact_row_reference() {
assert_eq!(std::mem::size_of::<HashEntry>(), 16);
}
#[test]
fn immutable_hash_state_reuses_only_exact_batch_and_key_layout() {
let rows = CompactArc::new(vec![make_row(vec![1, 10]), make_row(vec![2, 20])]);
let state = JoinHashState::build(CompactArc::clone(&rows), &[0]);
let reused = state.clone();
assert!(state.matches(&rows, &[0]));
assert!(!state.matches(&rows, &[1]));
assert!(!state.matches(&CompactArc::new((*rows).clone()), &[0]));
assert!(state.shares_allocations_with(&reused));
}
#[test]
fn hash_state_admission_counts_buckets_and_entries_before_allocation() {
let retained = JoinHashTable::estimated_retained_bytes(100).unwrap();
assert_eq!(retained, 256 * std::mem::size_of::<i32>() + 100 * 16);
assert!(JoinHashTable::fits_retained_budget(100, retained));
assert!(!JoinHashTable::fits_retained_budget(100, retained - 1));
assert!(JoinHashTable::estimated_retained_bytes(u32::MAX as usize + 1).is_none());
}
use radixdb_core::DataType;
fn make_row(values: Vec<i64>) -> Row {
Row::from_values(values.into_iter().map(Value::integer).collect())
}
#[test]
fn test_basic_insert_and_probe() {
let mut table = JoinHashTable::with_capacity(4);
table.insert(100, 0);
table.insert(200, 1);
table.insert(100, 2); table.insert(300, 3);
assert_eq!(table.len(), 4);
let matches: Vec<_> = table.probe(100).collect();
assert_eq!(matches.len(), 2);
assert!(matches.contains(&0));
assert!(matches.contains(&2));
let matches: Vec<_> = table.probe(200).collect();
assert_eq!(matches, vec![1]);
let matches: Vec<_> = table.probe(999).collect();
assert!(matches.is_empty());
}
#[test]
fn test_build_from_rows() {
let rows = vec![
make_row(vec![1, 10]),
make_row(vec![2, 20]),
make_row(vec![1, 30]), make_row(vec![3, 40]),
];
let key_indices = vec![0]; let table = JoinHashTable::build(&rows, &key_indices);
assert_eq!(table.len(), 4);
let hash = hash_row_keys(&rows[0], &key_indices);
let matches: Vec<_> = table.probe(hash).collect();
assert_eq!(matches.len(), 2);
}
#[cfg(feature = "bench-harness")]
#[test]
fn runtime_profile_observes_bulk_hash_build_under_parallel_tests() {
let before = radixdb_storage::instrumentation::snapshot().runtime_profile;
let rows = vec![make_row(vec![1]), make_row(vec![2]), make_row(vec![3])];
let _table = JoinHashTable::build(&rows, &[0]);
let after = radixdb_storage::instrumentation::snapshot().runtime_profile;
let delta = after.delta(before);
assert!(delta.hash_build_calls >= 1);
assert!(delta.hash_build_rows >= 3);
assert!(delta.hash_build_nanos > 0);
}
#[test]
fn test_empty_table() {
let table = JoinHashTable::empty();
assert!(table.is_empty());
assert_eq!(table.len(), 0);
let matches: Vec<_> = table.probe(100).collect();
assert!(matches.is_empty());
}
#[test]
fn test_load_factor() {
let mut table = JoinHashTable::with_capacity(100);
for i in 0..100 {
table.insert(i as u64, i as u32);
}
let load = table.load_factor();
assert!(
load > 0.3 && load <= 1.0,
"Load factor {} out of expected range",
load
);
assert_eq!(table.len(), 100);
}
#[test]
fn test_verify_key_equality() {
let row1 = Row::from_values(vec![Value::integer(1), Value::text("hello")]);
let row2 = Row::from_values(vec![Value::integer(1), Value::text("hello")]);
let row3 = Row::from_values(vec![Value::integer(2), Value::text("hello")]);
assert!(verify_key_equality(&row1, &row2, &[0, 1], &[0, 1]));
assert!(!verify_key_equality(&row1, &row3, &[0, 1], &[0, 1]));
}
#[test]
fn test_hash_row_keys() {
let row1 = make_row(vec![1, 2, 3]);
let row2 = make_row(vec![1, 2, 3]);
let row3 = make_row(vec![1, 2, 4]);
let indices = vec![0, 1];
assert_eq!(
hash_row_keys(&row1, &indices),
hash_row_keys(&row2, &indices)
);
let row4 = make_row(vec![1, 2, 999]);
assert_eq!(
hash_row_keys(&row1, &indices),
hash_row_keys(&row4, &indices)
);
assert_ne!(hash_row_keys(&row1, &[0, 2]), hash_row_keys(&row3, &[0, 2]));
}
fn assert_join_key_equal(left: Value, right: Value) {
let left_row = Row::from_values(vec![left]);
let right_row = Row::from_values(vec![right]);
assert_eq!(
hash_row_keys(&left_row, &[0]),
hash_row_keys(&right_row, &[0])
);
assert!(verify_key_equality(&left_row, &right_row, &[0], &[0]));
}
#[test]
fn test_canonical_numeric_join_keys() {
const TWO_POW_53: i64 = 9_007_199_254_740_992;
assert_join_key_equal(Value::Integer(TWO_POW_53), Value::Float(TWO_POW_53 as f64));
assert_join_key_equal(
Value::Integer(TWO_POW_53 + 2),
Value::Float((TWO_POW_53 + 2) as f64),
);
assert_join_key_equal(Value::Float(0.0), Value::Float(-0.0));
assert_join_key_equal(Value::Integer(0), Value::Float(-0.0));
}
#[test]
fn test_rounded_integer_neighbor_does_not_join() {
const TWO_POW_53: i64 = 9_007_199_254_740_992;
let build_rows = [Row::from_values(vec![Value::Integer(TWO_POW_53 + 1)])];
let probe_row = Row::from_values(vec![Value::Float((TWO_POW_53 + 1) as f64)]);
let probe_hash = hash_row_keys(&probe_row, &[0]);
let mut table = JoinHashTable::with_capacity(1);
table.insert(probe_hash, 0);
assert!(!verify_key_equality(&probe_row, &build_rows[0], &[0], &[0]));
let verified_matches = table
.probe(probe_hash)
.filter(|&build_idx| {
verify_key_equality(&probe_row, &build_rows[build_idx], &[0], &[0])
})
.count();
assert_eq!(verified_matches, 0);
}
#[test]
fn test_nan_payloads_share_join_key_contract() {
let nan1 = f64::from_bits(0x7ff8_0000_0000_0001);
let nan2 = f64::from_bits(0x7ff8_0000_0000_0002);
assert_join_key_equal(Value::Float(nan1), Value::Float(nan2));
}
#[test]
fn test_null_keys_never_join() {
let left = Row::from_values(vec![Value::Null(DataType::Integer)]);
let right = Row::from_values(vec![Value::Null(DataType::Float)]);
assert_eq!(hash_row_keys(&left, &[0]), hash_row_keys(&right, &[0]));
assert!(!verify_key_equality(&left, &right, &[0], &[0]));
}
#[test]
fn test_composite_keys_use_canonical_value_contract() {
const TWO_POW_53: i64 = 9_007_199_254_740_992;
let integer_key =
Row::from_values(vec![Value::Integer(TWO_POW_53), Value::Text("same".into())]);
let float_key = Row::from_values(vec![
Value::Float(TWO_POW_53 as f64),
Value::Text("same".into()),
]);
let different_tail = Row::from_values(vec![
Value::Float(TWO_POW_53 as f64),
Value::Text("different".into()),
]);
let null_tail = Row::from_values(vec![
Value::Float(TWO_POW_53 as f64),
Value::Null(DataType::Text),
]);
assert_eq!(
hash_row_keys(&integer_key, &[0, 1]),
hash_row_keys(&float_key, &[0, 1])
);
assert!(verify_key_equality(
&integer_key,
&float_key,
&[0, 1],
&[0, 1]
));
assert!(!verify_key_equality(
&integer_key,
&different_tail,
&[0, 1],
&[0, 1]
));
assert!(!verify_key_equality(
&integer_key,
&null_tail,
&[0, 1],
&[0, 1]
));
}
#[test]
fn test_row_and_indexed_get_hash_paths_are_identical() {
let row = Row::from_values(vec![
Value::Integer(9_007_199_254_740_992),
Value::Float(-0.0),
Value::Text("key".into()),
]);
let indices = [0, 1, 2];
assert_eq!(
hash_row_keys(&row, &indices),
hash_keys_with(&indices, |idx| row.get(idx))
);
}
#[test]
fn test_chain_collision() {
let mut table = JoinHashTable {
bucket_heads: vec![-1; 4], entries: Vec::new(),
bucket_mask: 3,
len: 0,
};
table.insert(0, 0);
table.insert(4, 1);
table.insert(8, 2);
table.insert(12, 3);
assert_eq!(table.probe(0).count(), 1);
assert_eq!(table.probe(4).count(), 1);
assert_eq!(table.probe(8).count(), 1);
assert_eq!(table.probe(12).count(), 1);
}
}