diff --git a/Cargo.lock b/Cargo.lock index 4f7b8bb279c20..84bab3245210f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -586,9 +586,9 @@ dependencies = [ [[package]] name = "aws-lc-rs" -version = "1.16.3" +version = "1.18.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0ec6fb3fe69024a75fa7e1bfb48aa6cf59706a101658ea01bfd33b2b248a038f" +checksum = "b281d307588d634de920874890732659e2e7672f72b5e10e81badc1a8a83621e" dependencies = [ "aws-lc-sys", "zeroize", @@ -596,14 +596,15 @@ dependencies = [ [[package]] name = "aws-lc-sys" -version = "0.40.0" +version = "0.45.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f50037ee5e1e41e7b8f9d161680a725bd1626cb6f8c7e901f91f942850852fe7" +checksum = "9bff6c3b54fad79a2e60b8102caf565819711497c1f5f092f49508e2f5c31b27" dependencies = [ "cc", "cmake", "dunce", "fs_extra", + "pkg-config", ] [[package]] @@ -2513,7 +2514,6 @@ dependencies = [ "insta", "itertools 0.15.0", "log", - "num-traits", "parking_lot", "pin-project-lite", "rand 0.9.4", @@ -5379,9 +5379,9 @@ dependencies = [ [[package]] name = "rustls" -version = "0.23.39" +version = "0.23.45" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7c2c118cb077cca2822033836dfb1b975355dfb784b5e8da48f7b6c5db74e60e" +checksum = "0d41d731c7d2f962d1ccc364cec258de3c0e93b38c2fb3ba97ac74513048d634" dependencies = [ "aws-lc-rs", "log", @@ -5417,9 +5417,9 @@ dependencies = [ [[package]] name = "rustls-webpki" -version = "0.103.13" +version = "0.103.15" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e" +checksum = "f3c3cf1d8b1e7d4927e2d154c3fcb02979afb9939629c62cd9048d4f07b60ac2" dependencies = [ "aws-lc-rs", "ring", diff --git a/datafusion/physical-plan/Cargo.toml b/datafusion/physical-plan/Cargo.toml index 0f72b74840d01..6bde9ca74d076 100644 --- a/datafusion/physical-plan/Cargo.toml +++ b/datafusion/physical-plan/Cargo.toml @@ -83,7 +83,6 @@ hashbrown = { workspace = true } indexmap = { workspace = true } itertools = { workspace = true, features = ["use_std"] } log = { workspace = true } -num-traits = { workspace = true } parking_lot = { workspace = true } pin-project-lite = { workspace = true } serde_json = { workspace = true, features = ["preserve_order"] } diff --git a/datafusion/physical-plan/src/aggregates/group_values/multi_group_by/primitive.rs b/datafusion/physical-plan/src/aggregates/group_values/multi_group_by/primitive.rs index 148c5697dea3b..fb3a96d7e7d30 100644 --- a/datafusion/physical-plan/src/aggregates/group_values/multi_group_by/primitive.rs +++ b/datafusion/physical-plan/src/aggregates/group_values/multi_group_by/primitive.rs @@ -27,7 +27,6 @@ use arrow::array::{ }; use arrow::buffer::ScalarBuffer; use arrow::datatypes::DataType; -use arrow::util::bit_util::apply_bitwise_binary_op; use datafusion_common::Result; use datafusion_common::utils::split_vec_min_alloc; use datafusion_execution::memory_pool::proxy::VecAllocExt; @@ -75,12 +74,8 @@ where "called with nullable input" ); let array_values = array.as_primitive::().values(); - let n = lhs_rows.len(); - - // Build a packed comparison bitmask, then AND it into equal_to_results - let num_bytes = n.div_ceil(8); - let mut cmp_buf = vec![0u8; num_bytes]; - + // Most hash-selected candidates are equal. Preserve those bits and + // clear only mismatches, avoiding a temporary mask and a second pass. for (i, (&lhs_row, &rhs_row)) in lhs_rows.iter().zip(rhs_rows.iter()).enumerate() { if !equal_to_results.get_bit(i) { @@ -98,20 +93,10 @@ where }; // `left` was already canonicalized on append; canonicalize the // input so ±0 (and any future equivalence class) compares equal. - if left.is_eq(right.canonicalize()) { - cmp_buf[i / 8] |= 1 << (i % 8); + if !left.is_eq(right.canonicalize()) { + equal_to_results.set_bit(i, false); } } - - // AND the comparison result into the existing equal_to_results bitmask - apply_bitwise_binary_op( - equal_to_results.as_slice_mut(), - 0, - &cmp_buf, - 0, - n, - |a, b| a & b, - ); } pub fn vectorized_equal_nullable( @@ -494,6 +479,89 @@ mod tests { test_not_nullable_primitive_equal_to_internal(append, equal_to); } + #[test] + fn primitive_vectorized_equality_preserves_prior_masks_and_tail_bits() { + use arrow::array::Decimal128Array; + use arrow::datatypes::Decimal128Type; + + let array: ArrayRef = Arc::new( + Decimal128Array::from(vec![ + -(10_i128.pow(38) - 1), + -1, + 0, + 1, + 10_i128.pow(38) - 1, + ]) + .with_precision_and_scale(38, 0) + .unwrap(), + ); + let mut builder = PrimitiveGroupValueBuilder::::new( + DataType::Decimal128(38, 0), + ); + builder.vectorized_append(&array, &[0, 1, 2, 3, 4]).unwrap(); + for n in [0, 1, 7, 8, 9, 63, 64, 65, 8193] { + let mut mask = make_true_buffer(n + 5); + let mut lhs = Vec::new(); + let mut rhs = Vec::new(); + let mut expected = Vec::new(); + for i in 0..n { + if i % 3 == 0 { + mask.set_bit(i, false); + // Already-false candidates must not dereference either index. + lhs.push(usize::MAX); + rhs.push(usize::MAX); + expected.push(false); + } else { + lhs.push(i * 17 % 5); + rhs.push(i * 13 % 5); + expected.push(builder.equal_to(lhs[i], &array, rhs[i])); + } + } + expected.extend([true; 5]); + builder.vectorized_equal_to(&lhs, &array, &rhs, &mut mask); + assert_eq!(to_vec(&mask), expected, "n={n}"); + } + } + + #[test] + fn primitive_vectorized_equality_preserves_float_canonicalization() { + use arrow::array::Float64Array; + use arrow::datatypes::Float64Type; + + let stored: ArrayRef = Arc::new(Float64Array::from(vec![ + -0.0, + 0.0, + f64::NAN, + f64::INFINITY, + f64::NEG_INFINITY, + ])); + let input: ArrayRef = Arc::new(Float64Array::from(vec![ + 0.0, + -0.0, + f64::from_bits(0x7ff8_0000_0000_0001), + f64::INFINITY, + f64::NEG_INFINITY, + ])); + let mut builder = + PrimitiveGroupValueBuilder::::new(DataType::Float64); + builder + .vectorized_append(&stored, &[0, 1, 2, 3, 4]) + .unwrap(); + let mut mask = make_true_buffer(5); + builder.vectorized_equal_to( + &[0, 1, 2, 3, 4], + &input, + &[0, 1, 2, 3, 4], + &mut mask, + ); + assert_eq!( + to_vec(&mask), + (0..5) + .map(|i| builder.equal_to(i, &input, i)) + .collect::>() + ); + } + fn test_not_nullable_primitive_equal_to_internal(mut append: A, mut equal_to: E) where A: FnMut(&mut PrimitiveGroupValueBuilder, &ArrayRef, &[usize]), diff --git a/datafusion/physical-plan/src/joins/array_map.rs b/datafusion/physical-plan/src/joins/array_map.rs index 4e56cf013c8f7..77eb190eefea9 100644 --- a/datafusion/physical-plan/src/joins/array_map.rs +++ b/datafusion/physical-plan/src/joins/array_map.rs @@ -16,7 +16,6 @@ // under the License. use arrow_schema::DataType; -use num_traits::AsPrimitive; use std::mem::size_of; use crate::joins::MapOffset; @@ -26,12 +25,40 @@ use arrow::buffer::BooleanBuffer; use arrow::datatypes::ArrowNumericType; use datafusion_common::{Result, ScalarValue, internal_err}; -/// A macro to downcast only supported integer types (up to 64-bit) and invoke a generic function. +/// Conversion to the dense map's index domain. Native <=64-bit integers retain +/// their branch-free wrapping representation. Decimal128 values must fit i64: +/// truncating arbitrary i128 probes would alias distinct keys separated by 2^64. +trait ArrayMapKey: Copy { + fn map_key(self) -> Option; +} + +macro_rules! native_map_keys { + ($($t:ty),*) => {$( + impl ArrayMapKey for $t { + #[inline] + fn map_key(self) -> Option { + Some(self as u64) + } + } + )*}; +} + +native_map_keys!(i8, i16, i32, i64, u8, u16, u32, u64); + +impl ArrayMapKey for i128 { + #[inline] + fn map_key(self) -> Option { + i64::try_from(self).ok().map(|value| value as u64) + } +} + +/// Downcast supported native integers and Decimal128 without materializing a +/// narrowed Arrow array. Decimal range eligibility is checked from build bounds. /// /// Usage: `downcast_supported_integer!(data_type => (Method, arg1, arg2, ...))` /// /// The `Method` must be an associated method of [`ArrayMap`] that is generic over -/// `` and allow `T::Native: AsPrimitive`. +/// `` and allow `T::Native: ArrayMapKey`. macro_rules! downcast_supported_integer { ($DATA_TYPE:expr => ($METHOD:ident $(, $ARGS:expr)*)) => { match $DATA_TYPE { @@ -43,6 +70,7 @@ macro_rules! downcast_supported_integer { arrow::datatypes::DataType::UInt16 => ArrayMap::$METHOD::($($ARGS),*), arrow::datatypes::DataType::UInt32 => ArrayMap::$METHOD::($($ARGS),*), arrow::datatypes::DataType::UInt64 => ArrayMap::$METHOD::($($ARGS),*), + arrow::datatypes::DataType::Decimal128(..) => ArrayMap::$METHOD::($($ARGS),*), _ => { return internal_err!( "Unsupported type for ArrayMap: {:?}", @@ -124,6 +152,7 @@ impl ArrayMap { | DataType::UInt16 | DataType::UInt32 | DataType::UInt64 + | DataType::Decimal128(..) ) } @@ -137,6 +166,7 @@ impl ArrayMap { ScalarValue::UInt16(Some(v)) => Some(*v as u64), ScalarValue::UInt32(Some(v)) => Some(*v as u64), ScalarValue::UInt64(Some(v)) => Some(*v), + ScalarValue::Decimal128(Some(v), _, _) => v.map_key(), _ => None, } } @@ -210,13 +240,17 @@ impl ArrayMap { num_of_distinct_key: &mut usize, ) -> Result<()> where - T::Native: AsPrimitive, + T::Native: ArrayMapKey, { let arr = array.as_primitive::(); // Iterate in reverse to maintain FIFO order when there are duplicate keys. for (i, val) in arr.iter().enumerate().rev() { if let Some(val) = val { - let key: u64 = val.as_(); + let Some(key) = val.map_key() else { + return internal_err!( + "ArrayMap build key exceeds its supported domain" + ); + }; let Some(idx) = Self::key_to_index(key, offset_val, data.len()) else { return internal_err!("failed build Array idx >= data.len()"); }; @@ -292,7 +326,7 @@ impl ArrayMap { build_indices: &mut Vec, ) -> Result> where - T::Native: Copy + AsPrimitive, + T::Native: ArrayMapKey, { probe_indices.clear(); build_indices.clear(); @@ -312,7 +346,10 @@ impl ArrayMap { continue; } // SAFETY: prob_idx is guaranteed to be within bounds by the loop range. - let prob_val: u64 = unsafe { arr.value_unchecked(prob_idx) }.as_(); + let Some(prob_val) = unsafe { arr.value_unchecked(prob_idx) }.map_key() + else { + continue; + }; let Some(build_value) = self.get_value(prob_val) else { continue; }; @@ -359,7 +396,11 @@ impl ArrayMap { let is_last = prob_side_idx == arr.len() - 1; // SAFETY: prob_idx is guaranteed to be within bounds by the loop range. - let prob_val: u64 = unsafe { arr.value_unchecked(prob_side_idx) }.as_(); + let Some(prob_val) = + unsafe { arr.value_unchecked(prob_side_idx) }.map_key() + else { + continue; + }; let Some(build_idx) = self.get_value(prob_val) else { continue; }; @@ -403,7 +444,7 @@ impl ArrayMap { array: &ArrayRef, ) -> Result where - T::Native: AsPrimitive, + T::Native: ArrayMapKey, { let arr = array.as_primitive::(); let buffer = BooleanBuffer::collect_bool(arr.len(), |i| { @@ -411,8 +452,9 @@ impl ArrayMap { return false; } // SAFETY: i is within bounds [0, arr.len()) - let key: u64 = unsafe { arr.value_unchecked(i) }.as_(); - self.get_value(key).is_some() + unsafe { arr.value_unchecked(i) } + .map_key() + .is_some_and(|key| self.get_value(key).is_some()) }); Ok(BooleanArray::new(buffer, None)) } @@ -426,6 +468,121 @@ mod tests { use arrow::array::UInt64Array; use std::sync::Arc; + #[test] + fn decimal_bounds_require_a_lossless_index_domain() { + for scale in [0, 2, 19] { + for value in [i128::from(i64::MIN), -1, 0, i128::from(i64::MAX)] { + assert_eq!( + ArrayMap::key_to_u64(&ScalarValue::Decimal128( + Some(value), + 38, + scale + )), + Some(value as u64) + ); + } + for value in [ + i128::from(i64::MIN) - 1, + i128::from(i64::MAX) + 1, + 1_i128 << 64, + ] { + assert_eq!( + ArrayMap::key_to_u64(&ScalarValue::Decimal128( + Some(value), + 38, + scale + )), + None + ); + } + } + assert_eq!( + ArrayMap::key_to_u64(&ScalarValue::Decimal128(None, 38, 0)), + None + ); + } + + #[test] + fn decimal_probes_do_not_alias_after_u64_truncation() -> Result<()> { + use arrow::array::Decimal128Array; + + // Exercise both the unique-key and duplicate-chain lookup paths, and + // scaled decimals: the map indexes unscaled values, never SQL casts. + for duplicates in [false, true] { + for scale in [0, 2, 19] { + let values = if duplicates { + vec![Some(-1), Some(0), None, Some(1), Some(1)] + } else { + vec![Some(-1), Some(0), None, Some(1)] + }; + let build: ArrayRef = Arc::new( + Decimal128Array::from(values).with_precision_and_scale(38, scale)?, + ); + let map = ArrayMap::try_new(&build, -1_i64 as u64, 1)?; + let probe = [Arc::new( + Decimal128Array::from(vec![ + Some(-1), + Some(0), + Some(1), + None, + Some((1_i128 << 64) - 1), + Some(1_i128 << 64), + Some((1_i128 << 64) + 1), + Some(-(1_i128 << 64)), + Some(i128::from(i64::MIN) - 1), + Some(i128::from(i64::MAX) + 1), + ]) + .with_precision_and_scale(38, scale)?, + ) as ArrayRef]; + let contains = map.contain_keys(&probe)?; + assert_eq!( + contains.iter().collect::>(), + vec![ + Some(true), + Some(true), + Some(true), + Some(false), + Some(false), + Some(false), + Some(false), + Some(false), + Some(false), + Some(false) + ] + ); + let mut offset = Some((0, None)); + let (mut probes, mut builds, mut matches) = (vec![], vec![], vec![]); + while let Some(next) = offset { + offset = map.get_matched_indices_with_limit_offset( + &probe, + 1, + next, + &mut probes, + &mut builds, + )?; + matches.extend(probes.iter().copied().zip(builds.iter().copied())); + } + let mut expected = vec![(0, 0), (1, 1), (2, 3)]; + if duplicates { + expected.push((2, 4)); + } + assert_eq!(matches, expected); + } + } + Ok(()) + } + + #[test] + fn invalid_decimal_build_cannot_silently_truncate() -> Result<()> { + use arrow::array::Decimal128Array; + let build: ArrayRef = Arc::new( + Decimal128Array::from(vec![(1_i128 << 64) + 1]) + .with_precision_and_scale(38, 0)?, + ); + assert!(ArrayMap::try_new(&build, 0, 2).is_err()); + Ok(()) + } + #[test] fn test_array_map_limit_offset_duplicate_elements() -> Result<()> { let build: ArrayRef = Arc::new(Int32Array::from(vec![1, 1, 2])); diff --git a/datafusion/physical-plan/src/joins/hash_join/exec.rs b/datafusion/physical-plan/src/joins/hash_join/exec.rs index 66a946c0a1cfa..4110faa8d0bee 100644 --- a/datafusion/physical-plan/src/joins/hash_join/exec.rs +++ b/datafusion/physical-plan/src/joins/hash_join/exec.rs @@ -121,6 +121,7 @@ fn try_create_array_map( batches: &[RecordBatch], on_left: &[PhysicalExprRef], reservation: &mut MemoryReservation, + prereserved_table_bytes: usize, perfect_hash_join_small_build_threshold: usize, perfect_hash_join_min_key_density: f64, null_equality: NullEquality, @@ -187,12 +188,23 @@ fn try_create_array_map( } let mem_size = ArrayMap::estimate_memory_size(min_val, max_val, num_row); - reservation.try_grow(mem_size)?; + // The direct map replaces, rather than coexists with, the general hash + // table. Claim only the positive replacement delta. If this optional + // representation cannot fit, the already-budgeted general map remains valid. + if let Err(error) = + reservation.try_grow(mem_size.saturating_sub(prereserved_table_bytes)) + { + return match error { + DataFusionError::ResourcesExhausted(_) => Ok(None), + error => Err(error), + }; + } let batch = concat_batches(schema, batches)?; let left_values = evaluate_expressions_to_arrays(on_left, &batch)?; let array_map = ArrayMap::try_new(&left_values[0], min_val, max_val)?; + reservation.shrink(prereserved_table_bytes.saturating_sub(mem_size)); Ok(Some((array_map, batch, left_values))) } @@ -589,8 +601,9 @@ impl From<&HashJoinExec> for HashJoinExecBuilder { /// (also known as a "perfect hash join") instead of a general-purpose hash map. /// This optimization is used when: /// 1. There is exactly one join key. -/// 2. The join key is an integer type up to 64 bits wide that can be losslessly converted -/// to `u64` (128-bit integer types such as `i128` and `u128` are not supported). +/// 2. The join key is a native integer up to 64 bits, or Decimal128 whose complete +/// build bounds fit i64. Decimal probes outside that domain cannot match and +/// are rejected before any conversion; no Arrow array is narrowed or copied. /// 3. The range of keys is small enough (controlled by `perfect_hash_join_small_build_threshold`) /// OR the keys are sufficiently dense (controlled by `perfect_hash_join_min_key_density`). /// 4. build_side.num_rows() < u32::MAX @@ -2659,16 +2672,13 @@ pub(super) fn build_left_data( batches, on_left, &mut reservation, + prereserved_table_bytes, config.execution.perfect_hash_join_small_build_threshold, config.execution.perfect_hash_join_min_key_density, null_equality, )? { array_map_created_count.add(1); metrics.build_mem_used.add(array_map.size()); - // The perfect-hash map replaces the hash table the caller - // pre-reserved for; release that estimate. - reservation.shrink(prereserved_table_bytes); - (Map::ArrayMap(array_map), batch, left_value) } else { // Estimation of memory size, required for hashtable, prior to allocation. @@ -2837,6 +2847,109 @@ mod tests { } } + #[tokio::test] + async fn decimal_array_map_matches_general_hash_join() -> Result<()> { + use arrow::array::Decimal128Array; + + let make_input = |name: &str, + keys: Vec>, + scale| + -> Result> { + let schema = Arc::new(Schema::new(vec![Field::new( + name, + DataType::Decimal128(38, scale), + true, + )])); + let array = + Decimal128Array::from(keys).with_precision_and_scale(38, scale)?; + let batch = RecordBatch::try_new(Arc::clone(&schema), vec![Arc::new(array)])?; + Ok(TestMemoryExec::try_new_exec(&[vec![batch]], schema, None)?) + }; + + for scale in [0, 2, 19] { + for wide_build in [false, true] { + let mut keys = vec![Some(-2), Some(-1), Some(-1), Some(0), Some(1), None]; + if wide_build { + keys.push(Some((1_i128 << 64) - 1)); + } + let left = make_input("l", keys, scale)?; + let right = make_input( + "r", + vec![ + Some(-2), + Some(-1), + Some(0), + Some(1), + None, + Some((1_i128 << 64) - 1), + Some(1_i128 << 64), + Some(1 - (1_i128 << 64)), + ], + scale, + )?; + let on: JoinOn = + vec![(Arc::new(Column::new("l", 0)), Arc::new(Column::new("r", 0)))]; + for null_equality in [ + NullEquality::NullEqualsNothing, + NullEquality::NullEqualsNull, + ] { + for join_type in [ + JoinType::Inner, + JoinType::Left, + JoinType::Right, + JoinType::Full, + JoinType::LeftSemi, + JoinType::RightSemi, + JoinType::LeftAnti, + JoinType::RightAnti, + JoinType::LeftMark, + JoinType::RightMark, + ] { + for mode in + [PartitionMode::CollectLeft, PartitionMode::Partitioned] + { + let mut outputs = Vec::new(); + for enabled in [false, true] { + let (_, batches, metrics) = + join_collect_with_partition_mode( + Arc::clone(&left), + Arc::clone(&right), + on.clone(), + &join_type, + mode, + null_equality, + prepare_task_ctx(1, enabled), + ) + .await?; + if mode == PartitionMode::CollectLeft { + let used = metrics + .sum_by_name(ARRAY_MAP_CREATED_COUNT_METRIC_NAME) + .map(|v| v.as_usize()) + .unwrap_or(0) + > 0; + assert_eq!( + used, + enabled + && !wide_build + && null_equality + == NullEquality::NullEqualsNothing, + "scale={scale}, wide_build={wide_build}, {join_type:?}, {mode:?}, {null_equality:?}, enabled={enabled}", + ); + } + outputs.push(batches_to_sort_string(&batches)); + } + assert_eq!( + outputs[0], outputs[1], + "scale={scale}, wide_build={wide_build}, {join_type:?}, {mode:?}, {null_equality:?}" + ); + } + } + } + } + } + Ok(()) + } + fn build_schema_and_on() -> Result<(SchemaRef, SchemaRef, JoinOn)> { let left_schema = Arc::new(Schema::new(vec![ Field::new("a1", DataType::Int32, true), @@ -2853,6 +2966,119 @@ mod tests { Ok((left_schema, right_schema, on)) } + fn build_with_tight_map_budget( + keys: Vec, + allow_map_bytes: bool, + ) -> Result<(JoinLeftData, Arc, usize, bool)> { + use datafusion_execution::memory_pool::GreedyMemoryPool; + + let schema = Arc::new(Schema::new(vec![Field::new("k", DataType::Int32, false)])); + let low = *keys.iter().min().unwrap(); + let high = *keys.iter().max().unwrap(); + let num_rows = keys.len(); + let table_bytes = hash_table_estimate(num_rows)?; + let map_bytes = ArrayMap::estimate_memory_size(low as u64, high as u64, num_rows); + let batch = RecordBatch::try_new( + Arc::clone(&schema), + vec![Arc::new(Int32Array::from(keys))], + )?; + let input_bytes = get_record_batch_memory_size(&batch); + let permitted_map = if allow_map_bytes { map_bytes } else { 0 }; + println!( + "MAP_CAPACITY input={input_bytes} general={table_bytes} direct={map_bytes} cap={} allow_direct={allow_map_bytes}", + input_bytes + table_bytes.max(permitted_map) + ); + let pool: Arc = Arc::new(GreedyMemoryPool::new( + input_bytes + table_bytes.max(permitted_map), + )); + let reservation = MemoryConsumer::new("tight-map-test").register(&pool); + reservation.try_grow(input_bytes + table_bytes)?; + let mut config = ConfigOptions::default(); + config.execution.perfect_hash_join_small_build_threshold = usize::MAX; + config.execution.perfect_hash_join_min_key_density = 0.0; + let metrics = ExecutionPlanMetricsSet::new(); + let created = Count::new(); + let data = build_left_data( + &[batch], + num_rows, + &schema, + &[Arc::new(Column::new("k", 0))], + HASH_JOIN_SEED.random_state(), + reservation, + &BuildProbeJoinMetrics::new(0, &metrics), + false, + 1, + Some(PartitionBounds::new(vec![ColumnBounds::new( + ScalarValue::Int32(Some(low)), + ScalarValue::Int32(Some(high)), + )])), + false, + &config, + NullEquality::NullEqualsNothing, + &created, + table_bytes, + None, + )?; + let expected = input_bytes + + if allow_map_bytes { + map_bytes + } else { + table_bytes + }; + Ok((data, pool, expected, created.value() > 0)) + } + + #[test] + fn direct_map_reuses_pre_reserved_table_capacity() -> Result<()> { + // Cover both releasing an overestimate and growing only the positive + // replacement delta. Both fail if old and replacement estimates overlap. + for keys in [vec![0, 1, 2, 3], vec![0, 4095]] { + let (data, pool, expected, used_map) = + build_with_tight_map_budget(keys, true)?; + assert!(used_map); + assert_eq!(pool.reserved(), expected); + drop(data); + assert_eq!(pool.reserved(), 0); + } + Ok(()) + } + + #[test] + fn optional_direct_map_budget_rejection_keeps_general_hash_join() -> Result<()> { + let (data, pool, expected, used_map) = + build_with_tight_map_budget(vec![0, 10_000], false)?; + assert!(!used_map); + let Map::HashMap(map) = data.map.as_ref() else { + panic!("expected the general hash map after direct-map budget rejection"); + }; + // Bucket count is not a row count, especially with force_hash_collisions. + // Probe both keys, a miss, and a duplicate to verify the fallback retains + // the rows and still applies key equality when their hashes collide. + let probe_values: ArrayRef = + Arc::new(Int32Array::from(vec![10_000, 7, 0, 10_000])); + let mut hashes = vec![0; probe_values.len()]; + create_hashes([&probe_values], HASH_JOIN_SEED.random_state(), &mut hashes)?; + let (build_ids, probe_ids, next_offset) = lookup_join_hashmap( + map.as_ref(), + &data.values, + &[probe_values], + NullEquality::NullEqualsNothing, + &hashes, + None, + 8192, + (0, None), + &mut Vec::new(), + &mut Vec::new(), + )?; + assert_eq!(build_ids, UInt64Array::from(vec![1, 0, 1])); + assert_eq!(probe_ids, UInt32Array::from(vec![0, 2, 3])); + assert!(next_offset.is_none()); + assert_eq!(pool.reserved(), expected); + drop(data); + assert_eq!(pool.reserved(), 0); + Ok(()) + } + use crate::coalesce_partitions::CoalescePartitionsExec; use crate::execution_plan::Boundedness; use crate::filter::FilterExecBuilder; @@ -2876,6 +3102,7 @@ mod tests { exec_err, internal_err, }; use datafusion_execution::config::SessionConfig; + use datafusion_execution::memory_pool::MemoryPool; use datafusion_execution::runtime_env::RuntimeEnvBuilder; use datafusion_expr::Operator; use datafusion_physical_expr::expressions::{BinaryExpr, Literal}; diff --git a/datafusion/physical-plan/src/joins/hash_join/shared_bounds.rs b/datafusion/physical-plan/src/joins/hash_join/shared_bounds.rs index ac06c5da301da..c9e890394eccb 100644 --- a/datafusion/physical-plan/src/joins/hash_join/shared_bounds.rs +++ b/datafusion/physical-plan/src/joins/hash_join/shared_bounds.rs @@ -238,9 +238,10 @@ fn combine_membership_and_bounds( /// /// ## Partition Counting /// -/// The `total_partitions` count represents how many times `collect_build_side` will be called: -/// - **CollectLeft**: Number of output partitions (each accesses shared build data) -/// - **Partitioned**: Number of input partitions (each builds independently) +/// Count independent build reports, not the number of probe streams: +/// - **CollectLeft**: One report contains the complete shared build. +/// - **Partitioned**: Every input partition builds independently and must report +/// or be accounted for as canceled/unknown before a complete filter is safe. /// /// ## Thread Safety /// @@ -355,23 +356,20 @@ enum FinalizeInput { impl SharedBuildAccumulator { /// Creates a new SharedBuildAccumulator configured for the given partition mode /// - /// This method calculates how many times `collect_build_side` will be called based on the - /// partition mode's execution pattern. This count is critical for determining when we have - /// complete information from all partitions to build the dynamic filter. + /// Count how many independent build results are needed for a complete filter. /// /// ## Partition Mode Execution Patterns /// /// - **CollectLeft**: Build side is collected ONCE from partition 0 and shared via `OnceFut` - /// across all output partitions. Each output partition calls `collect_build_side` to access the shared build data. - /// Although this results in multiple invocations, the `report_partition_bounds` function contains deduplication logic to handle them safely. - /// Expected calls = number of output partitions. + /// across all output partitions. Every report therefore contains the same + /// complete build. Expected reports = 1; waiting for all probes can deadlock + /// when a parent polls only a subset of streams. Later reports are idempotent. /// /// /// - **Partitioned**: Each partition independently builds its own hash table by calling /// `collect_build_side` once. Expected calls = number of build partitions. /// - /// - **Auto**: Placeholder mode resolved during optimization. Uses 1 as safe default since - /// the actual mode will be determined and a new accumulator created before execution. + /// - **Auto**: Must be resolved during optimization, before execution. /// /// ## Why This Matters /// @@ -392,10 +390,8 @@ impl SharedBuildAccumulator { // Troubleshooting: If partition counts are incorrect, verify this logic matches // the actual execution pattern in collect_build_side() let expected_calls = match partition_mode { - // Each output partition accesses shared build data - PartitionMode::CollectLeft => { - right_child.output_partitioning().partition_count() - } + // Any report already describes the complete OnceFut build. + PartitionMode::CollectLeft => 1, // Each partition builds its own data PartitionMode::Partitioned => { left_child.output_partitioning().partition_count() @@ -450,8 +446,8 @@ impl SharedBuildAccumulator { /// Report build-side data from a partition /// - /// This unified method handles both CollectLeft and Partitioned modes. When all partitions - /// have reported (barrier wait), the leader builds the appropriate filter expression: + /// Once all independent build results are available (one in CollectLeft), + /// the leader builds the appropriate filter expression: /// - CollectLeft: Simple conjunction of bounds and membership check /// - Partitioned: CASE expression routing to per-partition filters /// @@ -962,6 +958,48 @@ mod tests { ) } + // Unlike the old helper (which hard-codes one expected report), use the + // production constructor with several probe partitions. A single report + // already describes the entire shared build, even if other probes are never + // polled by their parent or are canceled before reaching their build report. + #[tokio::test] + async fn collect_left_full_build_does_not_wait_for_unpolled_probes() { + let left = crate::empty::EmptyExec::new(test_probe_schema()); + let right = crate::empty::EmptyExec::new(test_probe_schema()).with_partitions(6); + let on_right = test_on_right(); + let dynamic_filter = test_dynamic_filter(&on_right); + let initial_generation = dynamic_filter.snapshot_generation(); + let acc = SharedBuildAccumulator::new_from_partition_mode( + PartitionMode::CollectLeft, + &left, + &right, + Arc::clone(&dynamic_filter), + on_right, + SeededRandomState::with_seed(1), + NullEquality::NullEqualsNothing, + false, + ); + for _ in 0..2 { + tokio::time::timeout( + std::time::Duration::from_millis(100), + acc.report_build_data(PartitionBuildData::CollectLeft { + pushdown: in_list(&[2, 5]), + bounds: no_bounds(), + keys_have_null: false, + }), + ) + .await + .expect("a complete shared build must not wait for unpolled probe streams") + .unwrap(); + } + assert_eq!( + dynamic_filter.snapshot_generation(), + initial_generation + 1, + "publish the filter only once" + ); + assert_in_list_column_values(¤t_expr(&acc), "probe_key", 0, &[2, 5]); + } + fn make_partitioned_expr_accumulator_for_test( num_partitions: usize, ) -> SharedBuildAccumulator {